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.
Files changed (51) hide show
  1. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/PKG-INFO +1 -1
  2. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/__init__.py +1 -1
  3. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_common.py +101 -122
  4. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/argn.py +9 -4
  5. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/generation.py +215 -10
  6. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/pyproject.toml +2 -1
  7. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/.gitignore +0 -0
  8. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/LICENSE +0 -0
  9. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/README.md +0 -0
  10. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_dtypes.py +0 -0
  11. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  12. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  13. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  14. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  15. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  16. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  17. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  18. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  19. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  20. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  21. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  22. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  23. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  24. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/__init__.py +0 -0
  25. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/common.py +0 -0
  26. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/encoding.py +0 -0
  27. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
  28. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/base.py +0 -0
  29. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  30. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  31. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/generation.py +0 -0
  32. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/lstm.py +0 -0
  33. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  34. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/training.py +0 -0
  35. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
  36. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_memory.py +0 -0
  37. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/__init__.py +0 -0
  38. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/common.py +0 -0
  39. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/encoding.py +0 -0
  40. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/fairness.py +0 -0
  41. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_tabular/training.py +0 -0
  42. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_training_utils.py +0 -0
  43. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/_workspace.py +0 -0
  44. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/analysis.py +0 -0
  45. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/domain.py +0 -0
  46. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/encoding.py +0 -0
  47. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/generation.py +0 -0
  48. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/logging.py +0 -0
  49. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/random_state.py +0 -0
  50. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/splitting.py +0 -0
  51. {mostlyai_engine-1.6.0 → mostlyai_engine-1.7.0}/mostlyai/engine/training.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mostlyai-engine
3
- Version: 1.6.0
3
+ Version: 1.7.0
4
4
  Summary: Synthetic Data Engine
5
5
  Project-URL: homepage, https://github.com/mostly-ai/mostlyai-engine
6
6
  Project-URL: repository, https://github.com/mostly-ai/mostlyai-engine
@@ -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.6.0"
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 _get_log_histogram_edges(idx: int, bins: int = 64) -> tuple[float, float]:
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 != np.inf]
645
- values = values[values != -np.inf]
646
- values = values[~np.isnan(values)]
647
- edge_list = [_get_log_histogram_edges(idx) for idx in range(len(hist))]
648
- min_val = min([lower for lower, _ in edge_list])
649
- max_val = max([upper for _, upper in edge_list]) - 1
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
- # compute histograms
653
- for v in values:
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
- # for testing
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 (list[int]): A list of log histogram counts.
678
- epsilon (float): The privacy budget to spend estimating the bounds.
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
- tuple[float | None, float | None]: A tuple of the estimated minimum and maximum values.
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, _ = _get_log_histogram_edges(lower_bin)
704
- _, upper = _get_log_histogram_edges(upper_bin)
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 quantile.
713
- from: http://cs-people.bu.edu/ads22/pubs/2011/stoc194-smith.pdf
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 (np.ndarray): A 1D array of numeric values.
717
- quantiles (list[float]): List of probabilities of the quantiles to estimate.
718
- epsilon (float): Privacy budget.
719
- lower (float): A bounding parameter. The quantile will be estimated only for values greater than or equal to this bound.
720
- upper (float): A bounding parameter. The quantile will be estimated only for values less than or equal to this bound.
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
- list[float]: The estimated quantile.
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
- Differentially private quantile estimation.
823
- First, estimate the bounds of the values, then use the bounds to estimate the quantiles.
824
- If the bounds are not available, estimate the quantiles using the unbounded method.
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 (list | np.ndarray): A list of numeric values.
828
- quantiles (list[float]): List of probabilities of the quantiles to estimate.
829
- epsilon (float): Privacy budget.
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[float]: The estimated quantiles.
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
- Differentially private selection of all categories whose true count >= threshold,
854
- via the Laplace vector mechanism + post-processing.
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 (dict): Mapping from category to its count.
858
- epsilon (float): Privacy budget.
859
- threshold (int): Threshold for non-rare values.
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
- list[str]: Categories whose noisy counts are above the threshold (DP guarantee: ε-DP).
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=self.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 = self.column_order or self.tgt_columns
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=self.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 = self.column_order or self.tgt_columns
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
- # remove dummy context key (if exists)
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
- df_syn.loc[seed_rows, seed_columns] = df_overwrite.loc[seed_rows, seed_columns]
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, just overwrite all seed columns
638
- df_syn[seed_columns] = df_seed[seed_columns].copy()
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
- seed_step_encoded = (
1018
- seed_batch_encoded_grouped.nth(seq_step) if seq_step < n_seed_steps else pd.DataFrame()
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 seed_batch_encoded.columns
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.6.0"
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