mostlyai-engine 1.6.0__tar.gz → 1.7.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/PKG-INFO +1 -1
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/__init__.py +1 -1
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_common.py +101 -122
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/argn.py +9 -4
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/generation.py +215 -10
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/pyproject.toml +2 -1
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/.gitignore +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/LICENSE +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/README.md +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_dtypes.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/__init__.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/common.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/encoding.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/base.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/generation.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/lstm.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/training.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_memory.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/__init__.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/common.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/encoding.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/fairness.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/training.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_training_utils.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_workspace.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/analysis.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/domain.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/encoding.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/generation.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/logging.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/random_state.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/splitting.py +0 -0
- {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/training.py +0 -0
|
@@ -22,7 +22,7 @@ from mostlyai.engine.splitting import split
|
|
|
22
22
|
from mostlyai.engine.training import train
|
|
23
23
|
|
|
24
24
|
__all__ = ["split", "analyze", "encode", "train", "generate", "init_logging", "set_random_state"]
|
|
25
|
-
__version__ = "1.
|
|
25
|
+
__version__ = "1.7.0"
|
|
26
26
|
|
|
27
27
|
# suppress specific warning related to os.fork() in multi-threaded processes
|
|
28
28
|
warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*multi-threaded.*fork.*")
|
|
@@ -618,10 +618,23 @@ class FixedSizeSampleBuffer:
|
|
|
618
618
|
self.n_clears += 1
|
|
619
619
|
|
|
620
620
|
|
|
621
|
-
def
|
|
621
|
+
def _get_log_histogram_bin_bounds(idx: int, bins: int = 64) -> tuple[float, float]:
|
|
622
622
|
"""
|
|
623
|
+
Compute the lower and upper boundaries for a logarithmically-spaced histogram bin.
|
|
624
|
+
|
|
625
|
+
Creates symmetric logarithmic bins around zero that efficiently represent values
|
|
626
|
+
across many orders of magnitude. With bins=64, creates 128 total bins covering
|
|
627
|
+
negative powers of 2, the range [-1, 1], and positive powers of 2.
|
|
628
|
+
|
|
623
629
|
Modified from OpenDP's SmartNoise SDK (MIT License)
|
|
624
630
|
Source: https://github.com/opendp/smartnoise-sdk/blob/main/sql/snsql/sql/_mechanisms/approx_bounds.py
|
|
631
|
+
|
|
632
|
+
Args:
|
|
633
|
+
idx: The bin index (0 to bins*2-1)
|
|
634
|
+
bins: Number of bins per side (default 64, creating 128 total bins)
|
|
635
|
+
|
|
636
|
+
Returns:
|
|
637
|
+
Tuple of (lower_edge, upper_edge) for the bin
|
|
625
638
|
"""
|
|
626
639
|
if idx == bins:
|
|
627
640
|
return (0.0, 1.0)
|
|
@@ -635,50 +648,59 @@ def _get_log_histogram_edges(idx: int, bins: int = 64) -> tuple[float, float]:
|
|
|
635
648
|
|
|
636
649
|
def compute_log_histogram(values: np.ndarray, bins: int = 64) -> list[int]:
|
|
637
650
|
"""
|
|
651
|
+
Compute a histogram using logarithmically-spaced bins for efficient distribution analysis.
|
|
652
|
+
|
|
653
|
+
This creates a histogram that can efficiently represent values spanning many orders of
|
|
654
|
+
magnitude (e.g., 0.001 to 1,000,000) using relatively few bins. The bins are symmetric
|
|
655
|
+
around zero with exponentially increasing widths away from zero.
|
|
656
|
+
|
|
638
657
|
Modified from OpenDP's SmartNoise SDK (MIT License)
|
|
639
658
|
Source: https://github.com/opendp/smartnoise-sdk/blob/main/sql/snsql/sql/_mechanisms/approx_bounds.py
|
|
640
|
-
"""
|
|
641
|
-
hist = [0.0] * bins * 2
|
|
642
659
|
|
|
660
|
+
Args:
|
|
661
|
+
values: Array of numeric values to histogram
|
|
662
|
+
bins: Number of bins per side (default 64, creating 128 total bins)
|
|
663
|
+
|
|
664
|
+
Returns:
|
|
665
|
+
List of counts for each bin. Invalid values (NaN, inf) are filtered out.
|
|
666
|
+
"""
|
|
667
|
+
# filter out invalid values
|
|
643
668
|
values = np.array(values, dtype=np.float64)
|
|
644
|
-
values = values[values
|
|
645
|
-
|
|
646
|
-
|
|
647
|
-
edge_list = [
|
|
648
|
-
|
|
649
|
-
|
|
669
|
+
values = values[~np.isinf(values) & ~np.isnan(values)]
|
|
670
|
+
|
|
671
|
+
# generate all bin edges efficiently
|
|
672
|
+
edge_list = [_get_log_histogram_bin_bounds(idx, bins) for idx in range(bins * 2)]
|
|
673
|
+
bin_edges = np.array([lower for lower, _ in edge_list] + [edge_list[-1][1]])
|
|
674
|
+
|
|
675
|
+
# clip values to be within the bin edges to ensure all values are counted
|
|
676
|
+
min_val = bin_edges[0]
|
|
677
|
+
max_val = bin_edges[-1]
|
|
650
678
|
values = np.clip(values, min_val, max_val)
|
|
651
679
|
|
|
652
|
-
#
|
|
653
|
-
|
|
654
|
-
bin = None
|
|
655
|
-
for idx, (lower, upper) in enumerate(edge_list):
|
|
656
|
-
if lower <= v < upper:
|
|
657
|
-
bin = idx
|
|
658
|
-
break
|
|
659
|
-
if bin is None:
|
|
660
|
-
bin = idx
|
|
661
|
-
hist[bin] += 1
|
|
680
|
+
# use numpy's histogram for efficient binning (O(n log bins) vs O(n * bins))
|
|
681
|
+
hist, _ = np.histogram(values, bins=bin_edges)
|
|
662
682
|
|
|
663
|
-
|
|
664
|
-
lower, upper = _get_log_histogram_edges(bin)
|
|
665
|
-
return hist
|
|
683
|
+
return hist.tolist()
|
|
666
684
|
|
|
667
685
|
|
|
668
686
|
def dp_approx_bounds(hist: list[int], epsilon: float) -> tuple[float | None, float | None]:
|
|
669
687
|
"""
|
|
688
|
+
Estimate the minimum and maximum values using a differentially private histogram.
|
|
689
|
+
|
|
690
|
+
Uses Laplace noise on histogram bin counts, then finds the lowest and highest bins
|
|
691
|
+
that exceed a threshold (based on failure probability). Returns None if insufficient
|
|
692
|
+
data or privacy budget makes reliable estimation impossible.
|
|
693
|
+
|
|
694
|
+
Reference: https://desfontain.es/thesis/Usability.html#usability-u-ding-
|
|
670
695
|
Modified from OpenDP's SmartNoise SDK (MIT License)
|
|
671
696
|
Source: https://github.com/opendp/smartnoise-sdk/blob/main/sql/snsql/sql/_mechanisms/approx_bounds.py
|
|
672
697
|
|
|
673
|
-
Estimate the minimium and maximum values of a list of values.
|
|
674
|
-
from: https://desfontain.es/thesis/Usability.html#usability-u-ding-
|
|
675
|
-
|
|
676
698
|
Args:
|
|
677
|
-
hist
|
|
678
|
-
epsilon
|
|
699
|
+
hist: A list of log histogram counts (typically from compute_log_histogram).
|
|
700
|
+
epsilon: The privacy budget to spend estimating the bounds.
|
|
679
701
|
|
|
680
702
|
Returns:
|
|
681
|
-
|
|
703
|
+
Tuple of (min, max) estimates, or (None, None) if bounds cannot be reliably estimated.
|
|
682
704
|
"""
|
|
683
705
|
|
|
684
706
|
n_bins = len(hist)
|
|
@@ -700,8 +722,8 @@ def dp_approx_bounds(hist: list[int], epsilon: float) -> tuple[float | None, flo
|
|
|
700
722
|
return (None, None)
|
|
701
723
|
|
|
702
724
|
lower_bin, upper_bin = min(exceeds), max(exceeds)
|
|
703
|
-
lower, _ =
|
|
704
|
-
_, upper =
|
|
725
|
+
lower, _ = _get_log_histogram_bin_bounds(lower_bin)
|
|
726
|
+
_, upper = _get_log_histogram_bin_bounds(upper_bin)
|
|
705
727
|
return (float(lower), float(upper))
|
|
706
728
|
|
|
707
729
|
|
|
@@ -709,18 +731,23 @@ def _dp_bounded_quantiles(
|
|
|
709
731
|
values: np.ndarray, quantiles: list[float], epsilon: float, lower: float, upper: float
|
|
710
732
|
) -> list[float]:
|
|
711
733
|
"""
|
|
712
|
-
Estimate the
|
|
713
|
-
|
|
734
|
+
Estimate quantiles with differential privacy using the Smith (2011) smooth sensitivity method.
|
|
735
|
+
|
|
736
|
+
Assumes values are bounded within [lower, upper]. Uses exponential mechanism to sample
|
|
737
|
+
quantile estimates with noise proportional to local sensitivity. Privacy budget is split
|
|
738
|
+
evenly across all requested quantiles. Results are post-processed to ensure monotonicity.
|
|
739
|
+
|
|
740
|
+
Reference: http://cs-people.bu.edu/ads22/pubs/2011/stoc194-smith.pdf
|
|
714
741
|
|
|
715
742
|
Args:
|
|
716
|
-
values
|
|
717
|
-
quantiles
|
|
718
|
-
epsilon
|
|
719
|
-
lower
|
|
720
|
-
upper
|
|
743
|
+
values: A 1D array of numeric values.
|
|
744
|
+
quantiles: List of quantile probabilities to estimate (e.g., [0.05, 0.5, 0.95]).
|
|
745
|
+
epsilon: Privacy budget (split evenly across quantiles).
|
|
746
|
+
lower: Lower bound for clipping values.
|
|
747
|
+
upper: Upper bound for clipping values.
|
|
721
748
|
|
|
722
749
|
Returns:
|
|
723
|
-
|
|
750
|
+
List of differentially private quantile estimates (monotonically ordered).
|
|
724
751
|
"""
|
|
725
752
|
|
|
726
753
|
_LOG.info(f"compute DP bounded quantiles within [{lower}, {upper}]")
|
|
@@ -747,89 +774,23 @@ def _dp_bounded_quantiles(
|
|
|
747
774
|
return results
|
|
748
775
|
|
|
749
776
|
|
|
750
|
-
# NOTE: the unbounded method is not used in the current implementation
|
|
751
|
-
# def _dp_unbounded_quantiles(
|
|
752
|
-
# values: np.ndarray, quantiles: list[float], epsilon: float, beta: float = 1.01
|
|
753
|
-
# ) -> tuple[list[float], float]:
|
|
754
|
-
# """
|
|
755
|
-
# Fully unbounded differentially private quantile estimation using two AboveThreshold calls
|
|
756
|
-
# with Exponential noise (one-sided Laplace).
|
|
757
|
-
|
|
758
|
-
# Implements Algorithm 4 from Durfee (2023):
|
|
759
|
-
# 1) AboveThreshold on positives: T1 = q*n, f_i = |{x_j + 1 < beta^i}|
|
|
760
|
-
# 2) AboveThreshold on negatives: T2 = (1-q)*n, f_i = |{x_j - 1 > -beta^i}|
|
|
761
|
-
# 3) If first halts at k>0: return beta^k - 1
|
|
762
|
-
# 4) If second halts at k>0: return -beta^k + 1
|
|
763
|
-
# 5) Otherwise return 0
|
|
764
|
-
|
|
765
|
-
# Args:
|
|
766
|
-
# values (np.ndarray): A 1D array of numeric values.
|
|
767
|
-
# quantiles (list[float]): List of probabilities of the quantiles to estimate.
|
|
768
|
-
# epsilon (float): Privacy budget.
|
|
769
|
-
# beta (float): Multiplicative step size (default 1.01). Section 6.4 from Durfee (2023) suggests the range [1.01, 1.001] and 1.01 for general use, especially for more significant
|
|
770
|
-
# decreases in epsilon or in the data size.
|
|
771
|
-
|
|
772
|
-
# Returns:
|
|
773
|
-
# list[float]: Differentially private estimates of the quantiles.
|
|
774
|
-
# """
|
|
775
|
-
|
|
776
|
-
# def above_threshold(
|
|
777
|
-
# values: np.ndarray, q: float, eps: float, beta: float, is_positive_side: bool
|
|
778
|
-
# ) -> tuple[int, float]:
|
|
779
|
-
# n = len(values)
|
|
780
|
-
# eps1 = eps2 = eps / 2.0
|
|
781
|
-
# T = q * n if is_positive_side else (1 - q) * n
|
|
782
|
-
# noisy_T = T + np.random.exponential(scale=1 / eps1)
|
|
783
|
-
# i = 0
|
|
784
|
-
# while True:
|
|
785
|
-
# candidate = beta**i - 1 if is_positive_side else -(beta**i - 1)
|
|
786
|
-
# f_i = (values < candidate).sum() if is_positive_side else (values > candidate).sum()
|
|
787
|
-
# noisy_f_i = f_i + np.random.exponential(scale=1 / eps2)
|
|
788
|
-
# if noisy_f_i >= noisy_T:
|
|
789
|
-
# return i, candidate
|
|
790
|
-
# i += 1
|
|
791
|
-
|
|
792
|
-
# _LOG.info("compute DP unbounded quantiles")
|
|
793
|
-
# # Split epsilon across quantiles and the two AboveThreshold calls per quantile
|
|
794
|
-
# eps_pass = epsilon / len(quantiles) / 2.0
|
|
795
|
-
|
|
796
|
-
# results = []
|
|
797
|
-
# for q in quantiles:
|
|
798
|
-
# # 1) Positive-side AboveThreshold
|
|
799
|
-
# k, candidate = above_threshold(values, q, eps_pass, beta, is_positive_side=True)
|
|
800
|
-
# if k > 0:
|
|
801
|
-
# results.append(candidate)
|
|
802
|
-
# else:
|
|
803
|
-
# # 2) Continue with negative-side AboveThreshold only if the first one did not halt at k > 0
|
|
804
|
-
# k, candidate = above_threshold(values, q, eps_pass, beta, is_positive_side=False)
|
|
805
|
-
# if k > 0:
|
|
806
|
-
# results.append(candidate)
|
|
807
|
-
# else:
|
|
808
|
-
# # 3) Return 0 if both AboveThreshold calls did not halt at k > 0
|
|
809
|
-
# results.append(0.0)
|
|
810
|
-
|
|
811
|
-
# # ensure monotonicity of results with respect to quantiles
|
|
812
|
-
# sorted_indices = [t[0] for t in sorted(enumerate(quantiles), key=lambda x: x[1])]
|
|
813
|
-
# sorted_results = sorted(results)
|
|
814
|
-
# results = [sorted_results[sorted_indices.index(i)] for i in range(len(quantiles))]
|
|
815
|
-
|
|
816
|
-
# # NOTE: consider returning the actual epsilon spent in the future, so that the unused budget can be used for training later
|
|
817
|
-
# return results
|
|
818
|
-
|
|
819
|
-
|
|
820
777
|
def dp_quantiles(values: list | np.ndarray, quantiles: list[float], epsilon: float) -> list[float]:
|
|
821
778
|
"""
|
|
822
|
-
|
|
823
|
-
|
|
824
|
-
|
|
779
|
+
Estimate quantiles with differential privacy using a two-phase approach.
|
|
780
|
+
|
|
781
|
+
Phase 1: Estimate data bounds using dp_approx_bounds on a log histogram.
|
|
782
|
+
Phase 2: Estimate quantiles within those bounds using _dp_bounded_quantiles.
|
|
783
|
+
|
|
784
|
+
Privacy budget is split as epsilon/(m+1) for bounds and m*epsilon/(m+1) for m quantiles.
|
|
785
|
+
Returns None values if bounds cannot be reliably estimated (insufficient data/privacy budget).
|
|
825
786
|
|
|
826
787
|
Args:
|
|
827
|
-
values
|
|
828
|
-
quantiles
|
|
829
|
-
epsilon
|
|
788
|
+
values: A list or array of numeric values.
|
|
789
|
+
quantiles: List of quantile probabilities to estimate (e.g., [0.05, 0.95]).
|
|
790
|
+
epsilon: Total privacy budget to allocate.
|
|
830
791
|
|
|
831
792
|
Returns:
|
|
832
|
-
list
|
|
793
|
+
List of differentially private quantile estimates, or list of None if estimation fails.
|
|
833
794
|
"""
|
|
834
795
|
values = np.array(values)
|
|
835
796
|
|
|
@@ -850,17 +811,21 @@ def dp_quantiles(values: list | np.ndarray, quantiles: list[float], epsilon: flo
|
|
|
850
811
|
|
|
851
812
|
def dp_non_rare(value_counts: dict[str, int], epsilon: float, threshold: int = 5) -> tuple[list[str], float]:
|
|
852
813
|
"""
|
|
853
|
-
|
|
854
|
-
|
|
814
|
+
Select non-rare categories (count >= threshold) with differential privacy.
|
|
815
|
+
|
|
816
|
+
Uses the Laplace vector mechanism: adds independent Laplace(1/ε) noise to each count,
|
|
817
|
+
then selects categories where noisy_count >= threshold. Also computes the non-rare ratio
|
|
818
|
+
(fraction of total counts in selected categories).
|
|
819
|
+
|
|
820
|
+
Provides ε-differential privacy via the Laplace mechanism with L1 sensitivity = 1.
|
|
855
821
|
|
|
856
822
|
Args:
|
|
857
|
-
value_counts
|
|
858
|
-
epsilon
|
|
859
|
-
threshold
|
|
823
|
+
value_counts: Mapping from category name to its count.
|
|
824
|
+
epsilon: Privacy budget.
|
|
825
|
+
threshold: Minimum count threshold for non-rare categories (default: 5).
|
|
860
826
|
|
|
861
827
|
Returns:
|
|
862
|
-
|
|
863
|
-
float: Non-rare ratio (DP guarantee: ε-DP).
|
|
828
|
+
Tuple of (selected_categories, non_rare_ratio), both with ε-DP guarantees.
|
|
864
829
|
"""
|
|
865
830
|
|
|
866
831
|
# 1. Add independent Laplace(1/ε) noise to each count (vector Laplace mechanism)
|
|
@@ -883,6 +848,20 @@ def dp_non_rare(value_counts: dict[str, int], epsilon: float, threshold: int = 5
|
|
|
883
848
|
|
|
884
849
|
|
|
885
850
|
def get_stochastic_rare_threshold(min_threshold: int = 5, noise_multiplier: float = 3) -> int:
|
|
851
|
+
"""
|
|
852
|
+
Generate a randomized threshold for rare category detection.
|
|
853
|
+
|
|
854
|
+
Adds uniform random noise to the base threshold to prevent adversaries from
|
|
855
|
+
exploiting knowledge of exact threshold values. The threshold is sampled from
|
|
856
|
+
[min_threshold, min_threshold + noise_multiplier).
|
|
857
|
+
|
|
858
|
+
Args:
|
|
859
|
+
min_threshold: Base threshold value (default: 5).
|
|
860
|
+
noise_multiplier: Maximum noise to add (default: 3).
|
|
861
|
+
|
|
862
|
+
Returns:
|
|
863
|
+
Integer threshold in range [min_threshold, min_threshold + noise_multiplier).
|
|
864
|
+
"""
|
|
886
865
|
return min_threshold + int(noise_multiplier * np.random.uniform())
|
|
887
866
|
|
|
888
867
|
|
|
@@ -983,6 +983,7 @@ class FlatModel(nn.Module):
|
|
|
983
983
|
top_p: float | None = None,
|
|
984
984
|
return_probs: list[str] | None = None,
|
|
985
985
|
fairness_transforms: dict[str, Any] | None = None,
|
|
986
|
+
column_order: list[str] | None = None,
|
|
986
987
|
) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]:
|
|
987
988
|
fixed_probs = fixed_probs or {}
|
|
988
989
|
fixed_values = fixed_values or {}
|
|
@@ -990,6 +991,7 @@ class FlatModel(nn.Module):
|
|
|
990
991
|
outputs = {}
|
|
991
992
|
probs = {}
|
|
992
993
|
fairness_transforms = fairness_transforms or {}
|
|
994
|
+
effective_column_order = column_order or self.column_order
|
|
993
995
|
|
|
994
996
|
if mode == "trn":
|
|
995
997
|
# forward pass through context compressor
|
|
@@ -1007,7 +1009,7 @@ class FlatModel(nn.Module):
|
|
|
1007
1009
|
col_mask = _make_permutation_mask(
|
|
1008
1010
|
col_embedding_dims=self.column_embedders.dims,
|
|
1009
1011
|
columns=self.tgt_columns,
|
|
1010
|
-
column_order=
|
|
1012
|
+
column_order=effective_column_order,
|
|
1011
1013
|
is_sequential=False,
|
|
1012
1014
|
device=self.device,
|
|
1013
1015
|
)
|
|
@@ -1046,7 +1048,7 @@ class FlatModel(nn.Module):
|
|
|
1046
1048
|
col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
|
|
1047
1049
|
|
|
1048
1050
|
# take sub columns in the specified generation order
|
|
1049
|
-
column_order =
|
|
1051
|
+
column_order = effective_column_order or self.tgt_columns
|
|
1050
1052
|
sub_column_order = [sub_col for col in column_order for sub_col in self.tgt_column_sub_columns[col]]
|
|
1051
1053
|
|
|
1052
1054
|
for sub_col in sub_column_order:
|
|
@@ -1267,12 +1269,15 @@ class SequentialModel(nn.Module):
|
|
|
1267
1269
|
history=None,
|
|
1268
1270
|
history_state=None,
|
|
1269
1271
|
context=None,
|
|
1272
|
+
column_order: list[str] | None = None,
|
|
1270
1273
|
) -> tuple[dict[str, torch.Tensor], torch.Tensor, torch.Tensor]:
|
|
1271
1274
|
fixed_probs = fixed_probs or {}
|
|
1272
1275
|
fixed_values = fixed_values or {}
|
|
1273
1276
|
if context is None:
|
|
1274
1277
|
context = self.context_compressor(x)
|
|
1275
1278
|
|
|
1279
|
+
effective_column_order = column_order or self.column_order
|
|
1280
|
+
|
|
1276
1281
|
has_ridx = any(sub_col.startswith(RIDX_SUB_COLUMN_PREFIX) for sub_col in self.tgt_cardinalities)
|
|
1277
1282
|
|
|
1278
1283
|
# SLEN and RIDX are masked for history
|
|
@@ -1318,7 +1323,7 @@ class SequentialModel(nn.Module):
|
|
|
1318
1323
|
col_mask = _make_permutation_mask(
|
|
1319
1324
|
col_embedding_dims=self.column_embedders.dims,
|
|
1320
1325
|
columns=self.tgt_columns,
|
|
1321
|
-
column_order=
|
|
1326
|
+
column_order=effective_column_order,
|
|
1322
1327
|
is_sequential=True,
|
|
1323
1328
|
device=self.device,
|
|
1324
1329
|
)
|
|
@@ -1387,7 +1392,7 @@ class SequentialModel(nn.Module):
|
|
|
1387
1392
|
col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
|
|
1388
1393
|
|
|
1389
1394
|
# take sub columns in the specified generation order
|
|
1390
|
-
column_order =
|
|
1395
|
+
column_order = effective_column_order or self.tgt_columns
|
|
1391
1396
|
sub_column_order = [sub_col for col in column_order for sub_col in self.tgt_column_sub_columns[col]]
|
|
1392
1397
|
|
|
1393
1398
|
for sub_col in sub_column_order:
|
|
@@ -16,6 +16,7 @@ import logging
|
|
|
16
16
|
import random
|
|
17
17
|
import time
|
|
18
18
|
import uuid
|
|
19
|
+
from collections.abc import Callable
|
|
19
20
|
from functools import partial
|
|
20
21
|
from pathlib import Path
|
|
21
22
|
from typing import Literal
|
|
@@ -231,6 +232,113 @@ def _regroup_partial_sequences_by_length(
|
|
|
231
232
|
return ctx_data, len(new_batches)
|
|
232
233
|
|
|
233
234
|
|
|
235
|
+
def _flat_null_pattern(group: pd.DataFrame, relevant_columns: list[str]) -> tuple:
|
|
236
|
+
"""returns tuple of bools: True if column is fully NULL"""
|
|
237
|
+
return tuple(group[col].isna().all() for col in relevant_columns)
|
|
238
|
+
|
|
239
|
+
|
|
240
|
+
def _trailing_null_pattern(group: pd.DataFrame, relevant_columns: list[str]) -> tuple:
|
|
241
|
+
"""returns tuple of bools: True if column has trailing NULLs"""
|
|
242
|
+
pattern = []
|
|
243
|
+
for col in relevant_columns:
|
|
244
|
+
col_values = group[col].reset_index(drop=True)
|
|
245
|
+
non_null_mask = col_values.notna()
|
|
246
|
+
if non_null_mask.any():
|
|
247
|
+
last_non_null_idx = non_null_mask[::-1].idxmax()
|
|
248
|
+
has_trailing_nulls = last_non_null_idx < len(col_values) - 1
|
|
249
|
+
else:
|
|
250
|
+
has_trailing_nulls = True
|
|
251
|
+
pattern.append(has_trailing_nulls)
|
|
252
|
+
return tuple(pattern)
|
|
253
|
+
|
|
254
|
+
|
|
255
|
+
def _regroup_by_pattern(
|
|
256
|
+
ctx_data: pd.DataFrame,
|
|
257
|
+
seed_data: pd.DataFrame,
|
|
258
|
+
ctx_primary_key: str,
|
|
259
|
+
imputation_columns: list[str],
|
|
260
|
+
pattern_fn: Callable[[pd.DataFrame, list[str]], tuple],
|
|
261
|
+
*,
|
|
262
|
+
groupby_key: str | None = None,
|
|
263
|
+
use_vectorized_regroup: bool = True,
|
|
264
|
+
) -> tuple[pd.DataFrame, int]:
|
|
265
|
+
"""regroup batches so that rows/sequences with the same NULL pattern are together
|
|
266
|
+
|
|
267
|
+
Args:
|
|
268
|
+
pattern_fn: function that computes NULL pattern for a grouped dataframe
|
|
269
|
+
groupby_key: key to group seed_data by (defaults to ctx_primary_key)
|
|
270
|
+
use_vectorized_regroup: if True, use fast categorical factorization;
|
|
271
|
+
if False, use nested loop approach
|
|
272
|
+
"""
|
|
273
|
+
# only consider columns that are BOTH in imputation_columns AND in seed_data
|
|
274
|
+
relevant_columns = [col for col in imputation_columns if col in seed_data.columns]
|
|
275
|
+
|
|
276
|
+
# early exit: no relevant columns
|
|
277
|
+
if not relevant_columns or seed_data[relevant_columns].isna().all().all():
|
|
278
|
+
return ctx_data, ctx_data["__BATCH"].nunique()
|
|
279
|
+
|
|
280
|
+
# compute NULL pattern for each group
|
|
281
|
+
groupby_key = groupby_key or ctx_primary_key
|
|
282
|
+
seed_data_grouped = seed_data.groupby(groupby_key, sort=False)
|
|
283
|
+
null_patterns = seed_data_grouped.apply(
|
|
284
|
+
lambda group: pattern_fn(group, relevant_columns),
|
|
285
|
+
include_groups=False,
|
|
286
|
+
).rename("__NULL_PATTERN")
|
|
287
|
+
|
|
288
|
+
# early exit: all NULL patterns are the same
|
|
289
|
+
if null_patterns.nunique() == 1:
|
|
290
|
+
return ctx_data, ctx_data["__BATCH"].nunique()
|
|
291
|
+
|
|
292
|
+
# add __NULL_PATTERN to ctx_data
|
|
293
|
+
ctx_data = ctx_data.assign(__NULL_PATTERN=ctx_data[ctx_primary_key].map(null_patterns))
|
|
294
|
+
|
|
295
|
+
# regroup batches
|
|
296
|
+
if use_vectorized_regroup:
|
|
297
|
+
# vectorized approach for flat data
|
|
298
|
+
ctx_data = ctx_data.assign(
|
|
299
|
+
__COMPOSITE_KEY=ctx_data["__BATCH"].astype(str) + "_" + ctx_data["__NULL_PATTERN"].astype(str)
|
|
300
|
+
)
|
|
301
|
+
composite_cat = pd.Categorical(ctx_data["__COMPOSITE_KEY"], categories=ctx_data["__COMPOSITE_KEY"].unique())
|
|
302
|
+
ctx_data = ctx_data.assign(__BATCH=pd.factorize(composite_cat)[0] + 1)
|
|
303
|
+
num_batches = ctx_data["__BATCH"].max()
|
|
304
|
+
ctx_data = ctx_data.drop(columns=["__NULL_PATTERN", "__COMPOSITE_KEY"]).reset_index(drop=True)
|
|
305
|
+
else:
|
|
306
|
+
# nested loop approach for sequential data
|
|
307
|
+
new_batches = []
|
|
308
|
+
for _, old_batch_df in ctx_data.groupby("__BATCH", sort=False):
|
|
309
|
+
for _, new_batch_df in old_batch_df.groupby("__NULL_PATTERN", sort=False):
|
|
310
|
+
new_batch_df = new_batch_df.assign(__BATCH=len(new_batches) + 1)
|
|
311
|
+
new_batches.append(new_batch_df)
|
|
312
|
+
ctx_data = pd.concat(new_batches, axis=0).drop(columns=["__NULL_PATTERN"]).reset_index(drop=True)
|
|
313
|
+
num_batches = len(new_batches)
|
|
314
|
+
|
|
315
|
+
return ctx_data, num_batches
|
|
316
|
+
|
|
317
|
+
|
|
318
|
+
def _regroup_by_null_pattern(
|
|
319
|
+
ctx_data: pd.DataFrame,
|
|
320
|
+
seed_data: pd.DataFrame,
|
|
321
|
+
ctx_primary_key: str,
|
|
322
|
+
tgt_context_key: str,
|
|
323
|
+
imputation_columns: list[str],
|
|
324
|
+
is_sequential: bool,
|
|
325
|
+
) -> tuple[pd.DataFrame, int]:
|
|
326
|
+
"""regroup batches by NULL pattern (flat) or trailing NULL pattern (sequential)"""
|
|
327
|
+
pattern_fn = _trailing_null_pattern if is_sequential else _flat_null_pattern
|
|
328
|
+
groupby_key = tgt_context_key if is_sequential else None
|
|
329
|
+
use_vectorized_regroup = not is_sequential
|
|
330
|
+
|
|
331
|
+
return _regroup_by_pattern(
|
|
332
|
+
ctx_data,
|
|
333
|
+
seed_data,
|
|
334
|
+
ctx_primary_key,
|
|
335
|
+
imputation_columns,
|
|
336
|
+
pattern_fn=pattern_fn,
|
|
337
|
+
groupby_key=groupby_key,
|
|
338
|
+
use_vectorized_regroup=use_vectorized_regroup,
|
|
339
|
+
)
|
|
340
|
+
|
|
341
|
+
|
|
234
342
|
def _reshape_pt_to_pandas(
|
|
235
343
|
data: list[torch.Tensor], sub_cols: list[str], keys: list[pd.Series], key_name: str
|
|
236
344
|
) -> pd.DataFrame:
|
|
@@ -262,18 +370,34 @@ def _reshape_pt_to_pandas(
|
|
|
262
370
|
return pd.concat([keys, df], axis=1)
|
|
263
371
|
|
|
264
372
|
|
|
373
|
+
def _drop_fully_null_imputed_columns(seed_batch: pd.DataFrame, imputation_columns: list[str]) -> pd.DataFrame:
|
|
374
|
+
"""drop columns from seed_batch that are in imputation_columns and are fully NULL
|
|
375
|
+
|
|
376
|
+
this allows the model to freely generate these columns rather than conditioning on NULL values.
|
|
377
|
+
"""
|
|
378
|
+
if not imputation_columns or seed_batch.empty:
|
|
379
|
+
return seed_batch
|
|
380
|
+
|
|
381
|
+
fully_null_cols = [col for col in imputation_columns if col in seed_batch.columns and seed_batch[col].isna().all()]
|
|
382
|
+
return seed_batch.drop(columns=fully_null_cols) if fully_null_cols else seed_batch
|
|
383
|
+
|
|
384
|
+
|
|
265
385
|
def _post_process_decoding(
|
|
266
386
|
syn: pd.DataFrame,
|
|
267
387
|
tgt_primary_key: str | None = None,
|
|
268
388
|
) -> pd.DataFrame:
|
|
269
|
-
#
|
|
389
|
+
# sort by dummy context key to restore original order (if exists)
|
|
270
390
|
if DUMMY_CONTEXT_KEY in syn:
|
|
391
|
+
syn = syn.sort_values(DUMMY_CONTEXT_KEY).reset_index(drop=True)
|
|
271
392
|
syn = syn.drop(columns=DUMMY_CONTEXT_KEY)
|
|
272
393
|
|
|
273
394
|
# generate primary keys, if they are not present
|
|
274
395
|
if tgt_primary_key and tgt_primary_key not in syn:
|
|
275
396
|
syn[tgt_primary_key] = _generate_primary_keys(len(syn), type="uuid")
|
|
276
397
|
|
|
398
|
+
# reset index to ensure sequential indices for consistent test assertions
|
|
399
|
+
syn = syn.reset_index(drop=True)
|
|
400
|
+
|
|
277
401
|
return syn
|
|
278
402
|
|
|
279
403
|
|
|
@@ -580,6 +704,7 @@ def decode_buffered_samples(
|
|
|
580
704
|
tgt_primary_key: str,
|
|
581
705
|
tgt_context_key: str,
|
|
582
706
|
decode_prev_steps: dict | None = None,
|
|
707
|
+
impute_columns: list[str] | None = None,
|
|
583
708
|
) -> pd.DataFrame:
|
|
584
709
|
is_sequential = tgt_stats["is_sequential"]
|
|
585
710
|
seq_len_stats = get_sequence_length_stats(tgt_stats)
|
|
@@ -614,10 +739,12 @@ def decode_buffered_samples(
|
|
|
614
739
|
|
|
615
740
|
# preserve all seed values
|
|
616
741
|
df_seed = pd.concat(seed_data, axis=0).reset_index(drop=True) if seed_data else pd.DataFrame()
|
|
742
|
+
|
|
617
743
|
if not df_seed.empty:
|
|
618
744
|
seed_columns = [col for col in df_seed.columns]
|
|
619
745
|
if is_sequential:
|
|
620
746
|
# overwrite first steps of each sequence in synthetic data with values from seed data
|
|
747
|
+
impute_columns = impute_columns or []
|
|
621
748
|
df_syn["__SEQ_IDX"] = df_syn.groupby(tgt_context_key).cumcount()
|
|
622
749
|
df_seed["__SEQ_IDX"] = df_seed.groupby(tgt_context_key).cumcount()
|
|
623
750
|
# df_overwrite is a dataframe with the same shape as df_syn, but with the seed values for the first steps of each sequence
|
|
@@ -630,12 +757,47 @@ def decode_buffered_samples(
|
|
|
630
757
|
)
|
|
631
758
|
# project df_overwrite onto df_syn
|
|
632
759
|
seed_rows = df_overwrite["__INDICATOR"] == "both"
|
|
633
|
-
|
|
760
|
+
# overwrite columns based on imputation logic
|
|
761
|
+
for col in seed_columns:
|
|
762
|
+
if col in [tgt_context_key, "__SEQ_IDX"]:
|
|
763
|
+
continue # skip the key columns
|
|
764
|
+
if col not in impute_columns:
|
|
765
|
+
# non-impute columns: override all values
|
|
766
|
+
df_syn.loc[seed_rows, col] = df_overwrite.loc[seed_rows, col]
|
|
767
|
+
else:
|
|
768
|
+
# impute columns: override only non-NULL seed values
|
|
769
|
+
mask = seed_rows & df_overwrite[col].notna()
|
|
770
|
+
df_syn.loc[mask, col] = df_overwrite.loc[mask, col]
|
|
634
771
|
df_syn.drop(columns=["__SEQ_IDX"], inplace=True)
|
|
635
772
|
df_seed.drop(columns=["__SEQ_IDX"], inplace=True)
|
|
636
773
|
else:
|
|
637
|
-
# for flat data,
|
|
638
|
-
|
|
774
|
+
# for flat data, overwrite seed columns using merge to handle reordered rows
|
|
775
|
+
# for non-impute columns: override all values
|
|
776
|
+
# for impute columns: override only non-NULL seed values (let model impute NULL values)
|
|
777
|
+
impute_columns = impute_columns or []
|
|
778
|
+
|
|
779
|
+
# use merge on context key to properly align seed values with synthetic data
|
|
780
|
+
df_overwrite = pd.merge(
|
|
781
|
+
df_syn[[tgt_context_key]].copy(),
|
|
782
|
+
df_seed,
|
|
783
|
+
on=tgt_context_key,
|
|
784
|
+
how="left",
|
|
785
|
+
suffixes=("", "_seed"),
|
|
786
|
+
)
|
|
787
|
+
|
|
788
|
+
# overwrite columns based on imputation logic
|
|
789
|
+
for col in seed_columns:
|
|
790
|
+
if col == tgt_context_key:
|
|
791
|
+
continue # skip the key column itself
|
|
792
|
+
seed_col_name = col if col in df_overwrite.columns else f"{col}_seed"
|
|
793
|
+
if seed_col_name in df_overwrite.columns:
|
|
794
|
+
if col not in impute_columns:
|
|
795
|
+
# non-impute columns: override all values
|
|
796
|
+
df_syn[col] = df_overwrite[seed_col_name]
|
|
797
|
+
else:
|
|
798
|
+
# impute columns: override only non-NULL seed values
|
|
799
|
+
mask = df_overwrite[seed_col_name].notna()
|
|
800
|
+
df_syn.loc[mask, col] = df_overwrite.loc[mask, seed_col_name]
|
|
639
801
|
|
|
640
802
|
# postprocess generated data
|
|
641
803
|
_LOG.info(f"post-process generated data {df_syn.shape}")
|
|
@@ -740,10 +902,11 @@ def generate(
|
|
|
740
902
|
fairness=fairness,
|
|
741
903
|
)
|
|
742
904
|
_LOG.info(f"{gen_column_order=}")
|
|
905
|
+
trn_column_order = get_columns_from_cardinalities(tgt_cardinalities)
|
|
906
|
+
_LOG.info(f"{trn_column_order=}")
|
|
907
|
+
|
|
743
908
|
if not enable_flexible_generation:
|
|
744
909
|
# check if resolved column order is the same as the one from training
|
|
745
|
-
trn_column_order = get_columns_from_cardinalities(tgt_cardinalities)
|
|
746
|
-
_LOG.info(f"{trn_column_order=}")
|
|
747
910
|
if gen_column_order != trn_column_order:
|
|
748
911
|
raise ValueError(
|
|
749
912
|
"The column order for generation does not match the column order from training, due to seed, rebalancing, fairness or imputation configs. "
|
|
@@ -925,6 +1088,12 @@ def generate(
|
|
|
925
1088
|
ctx_data, seed_data, ctx_primary_key, tgt_context_key
|
|
926
1089
|
)
|
|
927
1090
|
|
|
1091
|
+
# regroup by NULL pattern if imputation is enabled
|
|
1092
|
+
if imputation and seed_data is not None and len(seed_data) > 0:
|
|
1093
|
+
ctx_data, no_of_batches = _regroup_by_null_pattern(
|
|
1094
|
+
ctx_data, seed_data, ctx_primary_key, tgt_context_key, imputation.columns, is_sequential
|
|
1095
|
+
)
|
|
1096
|
+
|
|
928
1097
|
# keep at most 500k samples in memory before decoding and writing to disk
|
|
929
1098
|
buffer = FixedSizeSampleBuffer(capacity=500_000)
|
|
930
1099
|
|
|
@@ -937,6 +1106,9 @@ def generate(
|
|
|
937
1106
|
batch_size = len(ctx_batch)
|
|
938
1107
|
|
|
939
1108
|
seed_batch = seed_data[seed_data[tgt_context_key].isin(ctx_batch[ctx_primary_key])]
|
|
1109
|
+
# drop fully-NULL imputation columns from seed_batch to allow conditional generation
|
|
1110
|
+
if imputation:
|
|
1111
|
+
seed_batch = _drop_fully_null_imputed_columns(seed_batch, imputation.columns)
|
|
940
1112
|
seed_batch = apply_encoding_type_dtypes(seed_batch, seed_encoding_types)
|
|
941
1113
|
|
|
942
1114
|
if ctx_primary_key not in ctx_batch.columns:
|
|
@@ -1014,9 +1186,18 @@ def generate(
|
|
|
1014
1186
|
|
|
1015
1187
|
# get seed data for current step
|
|
1016
1188
|
seed_step = seed_batch_grouped.nth(seq_step) if seq_step < n_seed_steps else pd.DataFrame()
|
|
1017
|
-
|
|
1018
|
-
|
|
1019
|
-
)
|
|
1189
|
+
|
|
1190
|
+
# drop NULL imputation columns from seed_step to allow conditional generation
|
|
1191
|
+
if imputation and len(seed_step) > 0:
|
|
1192
|
+
seed_step = _drop_fully_null_imputed_columns(seed_step, imputation.columns)
|
|
1193
|
+
|
|
1194
|
+
# encode seed_step (after dropping NULL imputation columns)
|
|
1195
|
+
if len(seed_step) > 0:
|
|
1196
|
+
seed_step_encoded, _, _ = encode_df(
|
|
1197
|
+
df=seed_step, stats=tgt_stats, tgt_context_key=tgt_context_key
|
|
1198
|
+
)
|
|
1199
|
+
else:
|
|
1200
|
+
seed_step_encoded = pd.DataFrame()
|
|
1020
1201
|
|
|
1021
1202
|
# fix SIDX by incrementing ourselves instead of sampling
|
|
1022
1203
|
sidx = pd.Series([seq_step] * step_size)
|
|
@@ -1079,11 +1260,19 @@ def generate(
|
|
|
1079
1260
|
torch.as_tensor(seed_step_encoded[col].to_numpy(), device=model.device).type(torch.int),
|
|
1080
1261
|
dim=-1,
|
|
1081
1262
|
)
|
|
1082
|
-
for col in
|
|
1263
|
+
for col in seed_step_encoded.columns
|
|
1083
1264
|
if col in tgt_sub_columns
|
|
1084
1265
|
}
|
|
1085
1266
|
|
|
1086
1267
|
fixed_values = sidx_vals | slen_vals | ridx_vals | sdec_vals | seed_vals
|
|
1268
|
+
column_order = _resolve_gen_column_order(
|
|
1269
|
+
column_stats=tgt_stats["columns"],
|
|
1270
|
+
cardinalities=tgt_cardinalities,
|
|
1271
|
+
rebalancing=rebalancing,
|
|
1272
|
+
imputation=imputation,
|
|
1273
|
+
seed_data=seed_step,
|
|
1274
|
+
fairness=fairness,
|
|
1275
|
+
)
|
|
1087
1276
|
out_dct, history, history_state = model(
|
|
1088
1277
|
x=None, # not used in generation forward pass
|
|
1089
1278
|
mode="gen",
|
|
@@ -1095,6 +1284,7 @@ def generate(
|
|
|
1095
1284
|
history=history,
|
|
1096
1285
|
history_state=history_state,
|
|
1097
1286
|
context=context,
|
|
1287
|
+
column_order=column_order,
|
|
1098
1288
|
)
|
|
1099
1289
|
|
|
1100
1290
|
# transform output dict to tensor for memory efficiency
|
|
@@ -1147,6 +1337,9 @@ def generate(
|
|
|
1147
1337
|
]
|
|
1148
1338
|
history = history[include_mask, ...]
|
|
1149
1339
|
history_state = tuple(h[:, include_mask, ...] for h in history_state)
|
|
1340
|
+
# filter seed_step to match step_ctx_keys (always, not just when filtering above)
|
|
1341
|
+
if len(seed_step) > 0:
|
|
1342
|
+
seed_step = seed_step[seed_step[tgt_context_key].isin(step_ctx_keys)].reset_index(drop=True)
|
|
1150
1343
|
# accumulate outputs in memory
|
|
1151
1344
|
buffer.add((out_pt, step_ctx_keys, seed_step))
|
|
1152
1345
|
# increment progress by 1 for each step
|
|
@@ -1160,6 +1353,7 @@ def generate(
|
|
|
1160
1353
|
tgt_primary_key=tgt_primary_key,
|
|
1161
1354
|
tgt_context_key=tgt_context_key,
|
|
1162
1355
|
decode_prev_steps=decode_prev_steps,
|
|
1356
|
+
impute_columns=imputation.columns if imputation else None,
|
|
1163
1357
|
)
|
|
1164
1358
|
persist_data_part(syn, output_path, f"{buffer.n_clears:06}.{0:06}")
|
|
1165
1359
|
buffer.clear()
|
|
@@ -1190,6 +1384,14 @@ def generate(
|
|
|
1190
1384
|
if col in tgt_sub_columns
|
|
1191
1385
|
}
|
|
1192
1386
|
|
|
1387
|
+
column_order = _resolve_gen_column_order(
|
|
1388
|
+
column_stats=tgt_stats["columns"],
|
|
1389
|
+
cardinalities=tgt_cardinalities,
|
|
1390
|
+
rebalancing=rebalancing,
|
|
1391
|
+
imputation=imputation,
|
|
1392
|
+
seed_data=seed_batch,
|
|
1393
|
+
fairness=fairness,
|
|
1394
|
+
)
|
|
1193
1395
|
out_dct, _ = model(
|
|
1194
1396
|
x,
|
|
1195
1397
|
mode="gen",
|
|
@@ -1199,6 +1401,7 @@ def generate(
|
|
|
1199
1401
|
temperature=sampling_temperature,
|
|
1200
1402
|
top_p=sampling_top_p,
|
|
1201
1403
|
fairness_transforms=fairness_transforms,
|
|
1404
|
+
column_order=column_order,
|
|
1202
1405
|
)
|
|
1203
1406
|
|
|
1204
1407
|
syn = pd.concat(
|
|
@@ -1224,6 +1427,7 @@ def generate(
|
|
|
1224
1427
|
tgt_primary_key=tgt_primary_key,
|
|
1225
1428
|
tgt_context_key=tgt_context_key,
|
|
1226
1429
|
decode_prev_steps=decode_prev_steps,
|
|
1430
|
+
impute_columns=imputation.columns if imputation else None,
|
|
1227
1431
|
)
|
|
1228
1432
|
persist_data_part(syn, output_path, f"{buffer.n_clears:06}.{0:06}")
|
|
1229
1433
|
buffer.clear()
|
|
@@ -1239,6 +1443,7 @@ def generate(
|
|
|
1239
1443
|
tgt_primary_key=tgt_primary_key,
|
|
1240
1444
|
tgt_context_key=tgt_context_key,
|
|
1241
1445
|
decode_prev_steps=decode_prev_steps,
|
|
1446
|
+
impute_columns=imputation.columns if imputation else None,
|
|
1242
1447
|
)
|
|
1243
1448
|
persist_data_part(syn, output_path, f"{buffer.n_clears:06}.{0:06}")
|
|
1244
1449
|
buffer.clear()
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "mostlyai-engine"
|
|
3
|
-
version = "1.
|
|
3
|
+
version = "1.7.0"
|
|
4
4
|
description = "Synthetic Data Engine"
|
|
5
5
|
authors = [{ name = "MOSTLY AI", email = "dev@mostly.ai" }]
|
|
6
6
|
requires-python = ">=3.10"
|
|
@@ -55,6 +55,7 @@ gpu = [
|
|
|
55
55
|
[dependency-groups]
|
|
56
56
|
dev = [
|
|
57
57
|
"pytest>=8.0",
|
|
58
|
+
"pytest-rerunfailures>=15.0",
|
|
58
59
|
"ruff>=0.11", # sync'ed with .pre-commit-config
|
|
59
60
|
"pre-commit>=4.0",
|
|
60
61
|
"twine>=6.1",
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/numeric.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/text.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/character.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/itt.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/numeric.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/hf_engine.py
RENAMED
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/vllm_engine.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/tokenizer_utils.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|