mostlyai-engine 2.2.0__tar.gz → 2.3.1__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 (54) hide show
  1. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/PKG-INFO +46 -5
  2. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/README.md +45 -4
  3. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/__init__.py +1 -1
  4. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/common.py +117 -2
  5. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/generation.py +6 -64
  6. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/interface.py +52 -0
  7. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/probability.py +165 -49
  8. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/splitting.py +15 -9
  9. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/pyproject.toml +1 -1
  10. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/.gitignore +0 -0
  11. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/LICENSE +0 -0
  12. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_common.py +0 -0
  13. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_dtypes.py +0 -0
  14. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  15. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  16. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  17. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  18. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  19. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  20. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  21. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  22. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  23. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  24. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  25. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  26. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  27. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/__init__.py +0 -0
  28. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/common.py +0 -0
  29. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/encoding.py +0 -0
  30. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/engine/__init__.py +0 -0
  31. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/engine/base.py +0 -0
  32. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  33. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  34. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/generation.py +0 -0
  35. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/interface.py +0 -0
  36. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/lstm.py +0 -0
  37. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  38. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/training.py +0 -0
  39. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
  40. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_memory.py +0 -0
  41. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/__init__.py +0 -0
  42. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/argn.py +0 -0
  43. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/encoding.py +0 -0
  44. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/fairness.py +0 -0
  45. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_tabular/training.py +0 -0
  46. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_training_utils.py +0 -0
  47. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/_workspace.py +0 -0
  48. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/analysis.py +0 -0
  49. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/domain.py +0 -0
  50. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/encoding.py +0 -0
  51. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/generation.py +0 -0
  52. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/logging.py +0 -0
  53. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/random_state.py +0 -0
  54. {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.1}/mostlyai/engine/training.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mostlyai-engine
3
- Version: 2.2.0
3
+ Version: 2.3.1
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
@@ -88,7 +88,8 @@ Two model classes with these methods are available:
88
88
  * `argn.fit(data)`: Train a TabularARGN model
89
89
  * `argn.sample(n_samples)`: Generate samples
90
90
  * `argn.predict(target, n_draws, agg_fn)`: Predict a feature
91
- * `argn.predict_proba(target, n_draws)`: Estimate probabilities
91
+ * `argn.predict_proba(target)`: Estimate probabilities
92
+ * `argn.log_prob(data)`: Compute log likelihood
92
93
  * `argn.impute(data)`: Fill missing values
93
94
  2. `LanguageModel()`: For semi-structured, flat textual tabular data.
94
95
  * `.fit(data)`: Train a Language model
@@ -191,9 +192,6 @@ from sklearn.metrics import accuracy_score, roc_auc_score
191
192
  # predict class labels for a categorical
192
193
  predictions = argn.predict(data_test, target="income", n_draws=100, agg_fn="mode")
193
194
 
194
- # predict class probabilities for a categorical
195
- probabilities = argn.predict_proba(data_test, target="income", n_draws=100)
196
-
197
195
  # evaluate performance
198
196
  accuracy = accuracy_score(data_test["income"], predictions)
199
197
  auc = roc_auc_score(data_test["income"], probabilities[:, 1])
@@ -215,6 +213,49 @@ mae = mean_absolute_error(data_test["age"], predictions)
215
213
  print(f"MAE: {mae:.1f} years")
216
214
  ```
217
215
 
216
+ ### Conditional Probabilities
217
+
218
+ Assess any marginal conditional probability, for one or more target columns:
219
+
220
+ ```python
221
+ # extract class probabilities for a categorical
222
+ argn.predict_proba(
223
+ X=pd.DataFrame({
224
+ "age": [25, 30, 35],
225
+ "sex": ["Male", "Female", "Male"],
226
+ }),
227
+ target="income"
228
+ )
229
+
230
+ # extract bin probabilities for a numerical
231
+ argn.predict_proba(
232
+ X=pd.DataFrame({
233
+ # "age": [25, 30, 35],
234
+ "sex": ["Male", "Female", "Male"],
235
+ "occupation": ["Craft-repair", "Craft-repair", "Craft-repair"]
236
+ }),
237
+ target="capital_gain"
238
+ )
239
+
240
+ # extract two-way marginals
241
+ argn.predict_proba(
242
+ X=data_test[["age", "race"]],
243
+ target=["sex", "income"]
244
+ )
245
+ ```
246
+
247
+ ### Log Probability
248
+
249
+ Compute log likelihood of observations:
250
+
251
+ ```python
252
+ # compute log probability for each observation
253
+ log_probs = argn.log_prob(data_test)
254
+
255
+ # list top 10 outliers
256
+ data_test.iloc[log_probs.argsort()[:10]]
257
+ ```
258
+
218
259
  ## TabularARGN for Sequential Data
219
260
 
220
261
  For sequential data (e.g., time series or event logs), specify the context key:
@@ -37,7 +37,8 @@ Two model classes with these methods are available:
37
37
  * `argn.fit(data)`: Train a TabularARGN model
38
38
  * `argn.sample(n_samples)`: Generate samples
39
39
  * `argn.predict(target, n_draws, agg_fn)`: Predict a feature
40
- * `argn.predict_proba(target, n_draws)`: Estimate probabilities
40
+ * `argn.predict_proba(target)`: Estimate probabilities
41
+ * `argn.log_prob(data)`: Compute log likelihood
41
42
  * `argn.impute(data)`: Fill missing values
42
43
  2. `LanguageModel()`: For semi-structured, flat textual tabular data.
43
44
  * `.fit(data)`: Train a Language model
@@ -140,9 +141,6 @@ from sklearn.metrics import accuracy_score, roc_auc_score
140
141
  # predict class labels for a categorical
141
142
  predictions = argn.predict(data_test, target="income", n_draws=100, agg_fn="mode")
142
143
 
143
- # predict class probabilities for a categorical
144
- probabilities = argn.predict_proba(data_test, target="income", n_draws=100)
145
-
146
144
  # evaluate performance
147
145
  accuracy = accuracy_score(data_test["income"], predictions)
148
146
  auc = roc_auc_score(data_test["income"], probabilities[:, 1])
@@ -164,6 +162,49 @@ mae = mean_absolute_error(data_test["age"], predictions)
164
162
  print(f"MAE: {mae:.1f} years")
165
163
  ```
166
164
 
165
+ ### Conditional Probabilities
166
+
167
+ Assess any marginal conditional probability, for one or more target columns:
168
+
169
+ ```python
170
+ # extract class probabilities for a categorical
171
+ argn.predict_proba(
172
+ X=pd.DataFrame({
173
+ "age": [25, 30, 35],
174
+ "sex": ["Male", "Female", "Male"],
175
+ }),
176
+ target="income"
177
+ )
178
+
179
+ # extract bin probabilities for a numerical
180
+ argn.predict_proba(
181
+ X=pd.DataFrame({
182
+ # "age": [25, 30, 35],
183
+ "sex": ["Male", "Female", "Male"],
184
+ "occupation": ["Craft-repair", "Craft-repair", "Craft-repair"]
185
+ }),
186
+ target="capital_gain"
187
+ )
188
+
189
+ # extract two-way marginals
190
+ argn.predict_proba(
191
+ X=data_test[["age", "race"]],
192
+ target=["sex", "income"]
193
+ )
194
+ ```
195
+
196
+ ### Log Probability
197
+
198
+ Compute log likelihood of observations:
199
+
200
+ ```python
201
+ # compute log probability for each observation
202
+ log_probs = argn.log_prob(data_test)
203
+
204
+ # list top 10 outliers
205
+ data_test.iloc[log_probs.argsort()[:10]]
206
+ ```
207
+
167
208
  ## TabularARGN for Sequential Data
168
209
 
169
210
  For sequential data (e.g., time series or event logs), specify the context key:
@@ -34,7 +34,7 @@ __all__ = [
34
34
  "TabularARGN",
35
35
  "LanguageModel",
36
36
  ]
37
- __version__ = "2.2.0"
37
+ __version__ = "2.3.1"
38
38
 
39
39
  # suppress specific warning related to os.fork() in multi-threaded processes
40
40
  warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*multi-threaded.*fork.*")
@@ -19,10 +19,32 @@ from pathlib import Path
19
19
  import pandas as pd
20
20
  import torch
21
21
 
22
- from mostlyai.engine._common import CTXFLT, CTXSEQ
22
+ from mostlyai.engine._common import (
23
+ ARGN_COLUMN,
24
+ ARGN_PROCESSOR,
25
+ ARGN_TABLE,
26
+ CTXFLT,
27
+ CTXSEQ,
28
+ get_argn_name,
29
+ )
30
+ from mostlyai.engine._encoding_types.tabular.categorical import (
31
+ CATEGORICAL_SUB_COL_SUFFIX,
32
+ CATEGORICAL_UNKNOWN_TOKEN,
33
+ )
34
+ from mostlyai.engine._encoding_types.tabular.numeric import (
35
+ NUMERIC_BINNED_SUB_COL_SUFFIX,
36
+ NUMERIC_BINNED_UNKNOWN_TOKEN,
37
+ NUMERIC_DISCRETE_SUB_COL_SUFFIX,
38
+ NUMERIC_DISCRETE_UNKNOWN_TOKEN,
39
+ )
40
+ from mostlyai.engine._tabular.encoding import encode_df, pad_ctx_sequences
41
+ from mostlyai.engine.domain import ModelEncodingType, RareCategoryReplacementMethod
23
42
 
24
43
  _LOG = logging.getLogger(__name__)
25
44
 
45
+ # Type alias for fixed probabilities
46
+ CodeProbabilities = dict[int, float]
47
+
26
48
  DPLSTM_SUFFIXES: tuple = ("ih.weight", "ih.bias", "hh.weight", "hh.bias")
27
49
 
28
50
 
@@ -165,7 +187,6 @@ def prepare_context_inputs(
165
187
  - encoded_dataframe: Encoded context DataFrame (for extracting keys if needed)
166
188
  - encoded_primary_key: Name of encoded primary key column (None if not provided)
167
189
  """
168
- from mostlyai.engine._tabular.encoding import encode_df, pad_ctx_sequences
169
190
 
170
191
  # Encode context data
171
192
  ctx_encoded, ctx_primary_key_encoded, _ = encode_df(df=ctx_data, stats=ctx_stats, ctx_primary_key=ctx_primary_key)
@@ -198,3 +219,97 @@ def prepare_context_inputs(
198
219
 
199
220
  # Merge and return with encoded dataframe
200
221
  return (ctxflt_inputs | ctxseq_inputs), ctx_encoded, ctx_primary_key_encoded
222
+
223
+
224
+ def check_column_order(
225
+ gen_column_order: list[str],
226
+ trn_column_order: list[str],
227
+ ) -> None:
228
+ """
229
+ Check if column order matches training order.
230
+
231
+ Args:
232
+ gen_column_order: Column order for the current operation
233
+ trn_column_order: Column order from training
234
+
235
+ Raises:
236
+ ValueError: If column order doesn't match training order
237
+ """
238
+ if gen_column_order != trn_column_order:
239
+ raise ValueError(
240
+ "Column order does not match training order. "
241
+ "A change in column order is only permitted for models that were trained with `enable_flexible_generation=True`."
242
+ )
243
+
244
+
245
+ def fix_rare_token_probs(
246
+ stats: dict,
247
+ rare_category_replacement_method: RareCategoryReplacementMethod | None = None,
248
+ ) -> dict[str, dict[str, CodeProbabilities]]:
249
+ """
250
+ Create fixed probabilities to suppress rare tokens.
251
+
252
+ Args:
253
+ stats: Target statistics dict
254
+ rare_category_replacement_method: How to handle rare categories
255
+
256
+ Returns:
257
+ Dict of column -> sub_column -> code -> probability
258
+ """
259
+ # suppress rare token for categorical when no_of_rare_categories == 0
260
+ mask = {
261
+ col: {CATEGORICAL_SUB_COL_SUFFIX: {col_stats["codes"][CATEGORICAL_UNKNOWN_TOKEN]: 0.0}}
262
+ for col, col_stats in stats["columns"].items()
263
+ if col_stats["encoding_type"] == ModelEncodingType.tabular_categorical
264
+ if "codes" in col_stats
265
+ if col_stats.get("no_of_rare_categories", 0) == 0
266
+ }
267
+ # suppress rare token for categorical if RareCategoryReplacementMethod is sample
268
+ if rare_category_replacement_method == RareCategoryReplacementMethod.sample:
269
+ mask |= {
270
+ col: {CATEGORICAL_SUB_COL_SUFFIX: {col_stats["codes"][CATEGORICAL_UNKNOWN_TOKEN]: 0.0}}
271
+ for col, col_stats in stats["columns"].items()
272
+ if col_stats["encoding_type"] == ModelEncodingType.tabular_categorical
273
+ if "codes" in col_stats
274
+ }
275
+ # always suppress rare token for numeric_binned
276
+ mask |= {
277
+ col: {NUMERIC_BINNED_SUB_COL_SUFFIX: {col_stats["codes"][NUMERIC_BINNED_UNKNOWN_TOKEN]: 0.0}}
278
+ for col, col_stats in stats["columns"].items()
279
+ if col_stats["encoding_type"] == ModelEncodingType.tabular_numeric_binned
280
+ if "codes" in col_stats
281
+ }
282
+ # always suppress rare token for numeric_discrete
283
+ mask |= {
284
+ col: {NUMERIC_DISCRETE_SUB_COL_SUFFIX: {col_stats["codes"][NUMERIC_DISCRETE_UNKNOWN_TOKEN]: 0.0}}
285
+ for col, col_stats in stats["columns"].items()
286
+ if col_stats["encoding_type"] == ModelEncodingType.tabular_numeric_discrete
287
+ if "codes" in col_stats
288
+ }
289
+ return mask
290
+
291
+
292
+ def translate_fixed_probs(
293
+ fixed_probs: dict[str, dict[str, CodeProbabilities]], stats: dict
294
+ ) -> dict[str, CodeProbabilities]:
295
+ """
296
+ Translate fixed probs to ARGN naming conventions.
297
+
298
+ Args:
299
+ fixed_probs: Dict of column -> sub_column -> code -> probability
300
+ stats: Target statistics dict
301
+
302
+ Returns:
303
+ Dict of ARGN sub_column name -> code -> probability
304
+ """
305
+ mask = {
306
+ get_argn_name(
307
+ argn_processor=stats["columns"][col][ARGN_PROCESSOR],
308
+ argn_table=stats["columns"][col][ARGN_TABLE],
309
+ argn_column=stats["columns"][col][ARGN_COLUMN],
310
+ argn_sub_column=sub_col,
311
+ ): sub_col_mask
312
+ for col, col_mask in fixed_probs.items()
313
+ for sub_col, sub_col_mask in col_mask.items()
314
+ }
315
+ return mask
@@ -67,10 +67,8 @@ from mostlyai.engine._encoding_types.tabular.lat_long import decode_latlong
67
67
  from mostlyai.engine._encoding_types.tabular.numeric import (
68
68
  NUMERIC_BINNED_NULL_TOKEN,
69
69
  NUMERIC_BINNED_SUB_COL_SUFFIX,
70
- NUMERIC_BINNED_UNKNOWN_TOKEN,
71
70
  NUMERIC_DISCRETE_NULL_TOKEN,
72
71
  NUMERIC_DISCRETE_SUB_COL_SUFFIX,
73
- NUMERIC_DISCRETE_UNKNOWN_TOKEN,
74
72
  decode_numeric,
75
73
  )
76
74
  from mostlyai.engine._memory import get_available_ram_for_heuristics, get_available_vram_for_heuristics
@@ -80,10 +78,13 @@ from mostlyai.engine._tabular.argn import (
80
78
  SequentialModel,
81
79
  )
82
80
  from mostlyai.engine._tabular.common import (
81
+ check_column_order,
83
82
  create_and_load_model,
83
+ fix_rare_token_probs,
84
84
  load_model_artifacts,
85
85
  prepare_context_inputs,
86
86
  resolve_device,
87
+ translate_fixed_probs,
87
88
  )
88
89
  from mostlyai.engine._tabular.encoding import encode_df
89
90
  from mostlyai.engine._tabular.fairness import FairnessTransforms, get_fairness_transforms
@@ -454,43 +455,6 @@ def _generation_batch_size_heuristic(mem_available_gb: float, ctx_stats: dict, t
454
455
  #########################
455
456
 
456
457
 
457
- def _fix_rare_token_probs(
458
- stats: dict,
459
- rare_category_replacement_method: RareCategoryReplacementMethod | None = None,
460
- ) -> dict[str, dict[str, CodeProbabilities]]:
461
- # suppress rare token for categorical when no_of_rare_categories == 0
462
- mask = {
463
- col: {CATEGORICAL_SUB_COL_SUFFIX: {col_stats["codes"][CATEGORICAL_UNKNOWN_TOKEN]: 0.0}}
464
- for col, col_stats in stats["columns"].items()
465
- if col_stats["encoding_type"] == ModelEncodingType.tabular_categorical
466
- if "codes" in col_stats
467
- if col_stats.get("no_of_rare_categories", 0) == 0
468
- }
469
- # suppress rare token for categorical if RareCategoryReplacementMethod is sample
470
- if rare_category_replacement_method == RareCategoryReplacementMethod.sample:
471
- mask |= {
472
- col: {CATEGORICAL_SUB_COL_SUFFIX: {col_stats["codes"][CATEGORICAL_UNKNOWN_TOKEN]: 0.0}}
473
- for col, col_stats in stats["columns"].items()
474
- if col_stats["encoding_type"] == ModelEncodingType.tabular_categorical
475
- if "codes" in col_stats
476
- }
477
- # always suppress rare token for numeric_binned
478
- mask |= {
479
- col: {NUMERIC_BINNED_SUB_COL_SUFFIX: {col_stats["codes"][NUMERIC_BINNED_UNKNOWN_TOKEN]: 0.0}}
480
- for col, col_stats in stats["columns"].items()
481
- if col_stats["encoding_type"] == ModelEncodingType.tabular_numeric_binned
482
- if "codes" in col_stats
483
- }
484
- # always suppress rare token for numeric_discrete
485
- mask |= {
486
- col: {NUMERIC_DISCRETE_SUB_COL_SUFFIX: {col_stats["codes"][NUMERIC_DISCRETE_UNKNOWN_TOKEN]: 0.0}}
487
- for col, col_stats in stats["columns"].items()
488
- if col_stats["encoding_type"] == ModelEncodingType.tabular_numeric_discrete
489
- if "codes" in col_stats
490
- }
491
- return mask
492
-
493
-
494
458
  def _fix_imputation_probs(
495
459
  stats: dict,
496
460
  imputation: ImputationConfig | None = None,
@@ -573,23 +537,6 @@ def _fix_rebalancing_probs(
573
537
  return mask
574
538
 
575
539
 
576
- def _translate_fixed_probs(
577
- fixed_probs: dict[str, dict[str, CodeProbabilities]], stats: dict
578
- ) -> dict[str, CodeProbabilities]:
579
- # translate fixed probs to ARGN conventions
580
- mask = {
581
- get_argn_name(
582
- argn_processor=stats["columns"][col][ARGN_PROCESSOR],
583
- argn_table=stats["columns"][col][ARGN_TABLE],
584
- argn_column=stats["columns"][col][ARGN_COLUMN],
585
- argn_sub_column=sub_col,
586
- ): sub_col_mask
587
- for col, col_mask in fixed_probs.items()
588
- for sub_col, sub_col_mask in col_mask.items()
589
- }
590
- return mask
591
-
592
-
593
540
  def _deepmerge(*dictionaries: dict, merged: dict | None = None) -> dict:
594
541
  merged = merged or {}
595
542
  for dictionary in dictionaries:
@@ -903,17 +850,12 @@ def generate(
903
850
  _LOG.info(f"{trn_column_order=}")
904
851
 
905
852
  if not enable_flexible_generation:
906
- # check if resolved column order is the same as the one from training
907
- if gen_column_order != trn_column_order:
908
- raise ValueError(
909
- "The column order for generation does not match the column order from training, due to seed, rebalancing, fairness or imputation configs. "
910
- "A change in column order is only permitted for models that were trained with `enable_flexible_generation=True`."
911
- )
853
+ check_column_order(gen_column_order, trn_column_order)
912
854
  _LOG.info(f"{rare_category_replacement_method=}")
913
- rare_token_fixed_probs = _fix_rare_token_probs(tgt_stats, rare_category_replacement_method)
855
+ rare_token_fixed_probs = fix_rare_token_probs(tgt_stats, rare_category_replacement_method)
914
856
  imputation_fixed_probs = _fix_imputation_probs(tgt_stats, imputation)
915
857
  rebalancing_fixed_probs = _fix_rebalancing_probs(tgt_stats, rebalancing)
916
- fixed_probs = _translate_fixed_probs(
858
+ fixed_probs = translate_fixed_probs(
917
859
  fixed_probs=_deepmerge(
918
860
  rare_token_fixed_probs,
919
861
  imputation_fixed_probs,
@@ -38,6 +38,7 @@ from mostlyai.engine._common import (
38
38
  median_fn,
39
39
  mode_fn,
40
40
  )
41
+ from mostlyai.engine._tabular.probability import log_prob as _log_prob
41
42
  from mostlyai.engine._tabular.probability import predict_proba as _predict_proba
42
43
  from mostlyai.engine._workspace import Workspace
43
44
  from mostlyai.engine.analysis import analyze
@@ -730,3 +731,54 @@ class TabularARGN(BaseEstimator):
730
731
  )
731
732
 
732
733
  return probs_df
734
+
735
+ def log_prob(
736
+ self,
737
+ X: pd.DataFrame,
738
+ ctx_data: pd.DataFrame | None = None,
739
+ **kwargs,
740
+ ) -> np.ndarray:
741
+ """
742
+ Compute log probability of observations.
743
+
744
+ This method computes P(observation | model) - the likelihood of the observation
745
+ under the trained model. For autoregressive models:
746
+
747
+ log P(x1, x2, ..., xn) = log P(x1) + log P(x2|x1) + ... + log P(xn|x1,...,xn-1)
748
+
749
+ Each term is the probability the model assigns to the actual observed value.
750
+
751
+ Args:
752
+ X: DataFrame with all columns containing observed values.
753
+ ctx_data: Context data for generation. If None, uses the context data from training.
754
+ **kwargs: Additional generation parameters (device, etc.).
755
+
756
+ Returns:
757
+ np.ndarray of shape (n_samples,) with log probability per row.
758
+ Values are <= 0 (log probabilities).
759
+ More negative values indicate less likely samples.
760
+
761
+ Raises:
762
+ ValueError: If model is not fitted.
763
+ """
764
+ if not self._fitted:
765
+ raise ValueError("Model must be fitted before computing log_prob. Call fit() first.")
766
+
767
+ X_df = ensure_dataframe(X, columns=self._feature_names)
768
+
769
+ workspace = Workspace(self.workspace_dir)
770
+ device = kwargs.get("device", self.device)
771
+
772
+ # Prepare ctx_data
773
+ if ctx_data is None:
774
+ ctx_data = self.ctx_data
775
+ ctx_data_df = ensure_dataframe(ctx_data) if ctx_data is not None else None
776
+
777
+ log_probs = _log_prob(
778
+ workspace=workspace,
779
+ data=X_df,
780
+ ctx_data=ctx_data_df,
781
+ device=device,
782
+ )
783
+
784
+ return log_probs
@@ -43,22 +43,90 @@ from mostlyai.engine._encoding_types.tabular.numeric import (
43
43
  )
44
44
  from mostlyai.engine._tabular.argn import ModelSize
45
45
  from mostlyai.engine._tabular.common import (
46
+ check_column_order,
46
47
  create_and_load_model,
48
+ fix_rare_token_probs,
47
49
  load_model_artifacts,
48
50
  prepare_context_inputs,
49
51
  resolve_device,
52
+ translate_fixed_probs,
50
53
  )
51
54
  from mostlyai.engine._tabular.encoding import encode_df
52
- from mostlyai.engine._tabular.generation import _fix_rare_token_probs, _translate_fixed_probs
55
+ from mostlyai.engine._tabular.training import _calculate_sample_losses
53
56
  from mostlyai.engine._workspace import Workspace
54
57
  from mostlyai.engine.domain import ModelEncodingType, RareCategoryReplacementMethod
55
58
 
56
59
  _LOG = logging.getLogger(__name__)
57
60
 
58
61
 
59
- ##########################
60
- ### PROBABILITY UTILS ###
61
- ##########################
62
+ def _initialize_model(
63
+ *,
64
+ workspace: Workspace,
65
+ rare_category_replacement_method: RareCategoryReplacementMethod | str | None = None,
66
+ device: torch.device | str | None = None,
67
+ allow_sequential: bool = True,
68
+ ) -> tuple[torch.nn.Module, dict, dict, dict, list[str], dict, torch.device, bool]:
69
+ """
70
+ Initialize model and artifacts for probability computation.
71
+
72
+ Args:
73
+ workspace: Workspace containing model and stats
74
+ rare_category_replacement_method: How to handle rare categories (None to skip fixed_probs)
75
+ device: Device for computation
76
+ allow_sequential: Whether to allow sequential models
77
+
78
+ Returns:
79
+ Tuple of (model, tgt_stats, ctx_stats, tgt_cardinalities, all_columns, fixed_probs, device, enable_flexible_generation)
80
+
81
+ Raises:
82
+ ValueError: If model is sequential and allow_sequential is False
83
+ """
84
+ # Load model artifacts
85
+ model_config, tgt_stats, ctx_stats, is_sequential = load_model_artifacts(workspace)
86
+
87
+ # Check model type
88
+ if is_sequential and not allow_sequential:
89
+ raise ValueError("Sequential models are not supported for this operation")
90
+
91
+ # Get cardinalities and config
92
+ tgt_cardinalities = get_cardinalities(tgt_stats)
93
+ ctx_cardinalities = get_cardinalities(ctx_stats)
94
+ enable_flexible_generation = model_config.get("enable_flexible_generation", True)
95
+
96
+ # Resolve device
97
+ device = resolve_device(device)
98
+ _LOG.info(f"Using device: {device}")
99
+
100
+ # Get all columns in training order
101
+ all_columns = get_columns_from_cardinalities(tgt_cardinalities)
102
+
103
+ # Create model
104
+ model_units = model_config.get("model_units") or ModelSize.M
105
+ ctx_seq_len_median = ctx_stats.get("sequence_len_median")
106
+
107
+ model = create_and_load_model(
108
+ workspace=workspace,
109
+ is_sequential=is_sequential,
110
+ tgt_cardinalities=tgt_cardinalities,
111
+ ctx_cardinalities=ctx_cardinalities,
112
+ model_units=model_units,
113
+ ctx_seq_len_median=ctx_seq_len_median,
114
+ column_order=all_columns,
115
+ device=device,
116
+ )
117
+
118
+ # Prepare fixed_probs to suppress rare tokens (if replacement method provided)
119
+ if rare_category_replacement_method is not None:
120
+ _LOG.info(f"{rare_category_replacement_method=}")
121
+ rare_token_fixed_probs = fix_rare_token_probs(tgt_stats, rare_category_replacement_method)
122
+ fixed_probs = translate_fixed_probs(
123
+ fixed_probs=rare_token_fixed_probs,
124
+ stats=tgt_stats,
125
+ )
126
+ else:
127
+ fixed_probs = {}
128
+
129
+ return model, tgt_stats, ctx_stats, tgt_cardinalities, all_columns, fixed_probs, device, enable_flexible_generation
62
130
 
63
131
 
64
132
  def _get_column_metadata(target_column: str, target_stats: dict) -> list[dict]:
@@ -181,6 +249,7 @@ def _generate_marginal_probs(
181
249
  ctx_data: pd.DataFrame | None = None,
182
250
  ctx_stats: dict | None = None,
183
251
  fixed_probs: dict | None = None,
252
+ enable_flexible_generation: bool = True,
184
253
  ) -> np.ndarray:
185
254
  """
186
255
  Generate P(target | seed_features, context).
@@ -195,6 +264,8 @@ def _generate_marginal_probs(
195
264
  device: Device for computation
196
265
  ctx_data: Optional context data
197
266
  ctx_stats: Optional context statistics (required if ctx_data provided)
267
+ fixed_probs: Optional fixed probabilities for rare token handling
268
+ enable_flexible_generation: Whether flexible generation is enabled
198
269
 
199
270
  Returns:
200
271
  DataFrame of shape (n_samples, cardinality) with probabilities and column names
@@ -229,6 +300,10 @@ def _generate_marginal_probs(
229
300
  # Determine column order: seed columns + target (in training order)
230
301
  gen_column_order = [col for col in all_columns if col in seed_columns or col == target_column]
231
302
 
303
+ # Check column order when flexible generation is disabled
304
+ if not enable_flexible_generation:
305
+ check_column_order(gen_column_order, all_columns)
306
+
232
307
  # Prepare context inputs if provided
233
308
  if ctx_data is not None and ctx_stats is not None:
234
309
  x, _, _ = prepare_context_inputs(ctx_data=ctx_data, ctx_stats=ctx_stats, device=device)
@@ -292,45 +367,14 @@ def predict_proba(
292
367
  """
293
368
  _LOG.info(f"PREDICT_PROBA started for targets: {target_columns}")
294
369
 
295
- # Load model artifacts
296
- model_config, tgt_stats, ctx_stats, is_sequential = load_model_artifacts(workspace)
297
-
298
- # Check model type
299
- if is_sequential:
300
- raise ValueError("predict_proba is not yet supported for sequential models")
301
-
302
- # Get cardinalities
303
- tgt_cardinalities = get_cardinalities(tgt_stats)
304
- ctx_cardinalities = get_cardinalities(ctx_stats)
305
-
306
- # Resolve device
307
- device = resolve_device(device)
308
- _LOG.info(f"Using device: {device}")
309
-
310
- # Get all columns in training order
311
- all_columns = get_columns_from_cardinalities(tgt_cardinalities)
312
-
313
- # Create model (will override column_order in forward pass for optimization)
314
- model_units = model_config.get("model_units") or ModelSize.M
315
- ctx_seq_len_median = ctx_stats.get("sequence_len_median")
316
-
317
- model = create_and_load_model(
318
- workspace=workspace,
319
- is_sequential=is_sequential,
320
- tgt_cardinalities=tgt_cardinalities,
321
- ctx_cardinalities=ctx_cardinalities,
322
- model_units=model_units,
323
- ctx_seq_len_median=ctx_seq_len_median,
324
- column_order=all_columns, # Will override in forward pass
325
- device=device,
326
- )
327
-
328
- # Prepare fixed_probs to suppress rare tokens
329
- _LOG.info(f"{rare_category_replacement_method=}")
330
- rare_token_fixed_probs = _fix_rare_token_probs(tgt_stats, rare_category_replacement_method)
331
- fixed_probs = _translate_fixed_probs(
332
- fixed_probs=rare_token_fixed_probs,
333
- stats=tgt_stats,
370
+ # Initialize model
371
+ model, tgt_stats, ctx_stats, tgt_cardinalities, all_columns, fixed_probs, device, enable_flexible_generation = (
372
+ _initialize_model(
373
+ workspace=workspace,
374
+ rare_category_replacement_method=rare_category_replacement_method,
375
+ device=device,
376
+ allow_sequential=False,
377
+ )
334
378
  )
335
379
 
336
380
  # Encode seed data (features to condition on) - common for both single and multi-target
@@ -352,12 +396,13 @@ def predict_proba(
352
396
  target_cardinality = sum(col_stats["cardinalities"].values())
353
397
  total_cardinality *= target_cardinality
354
398
 
355
- _LOG.warning(
356
- f"Computing joint probabilities for {len(target_columns)} targets "
357
- f"results in {total_cardinality:,} total probability values per sample. "
358
- f"Computation complexity grows exponentially with the number of targets. "
359
- f"Consider computing probabilities for targets separately if this takes too long."
360
- )
399
+ if total_cardinality > 100:
400
+ _LOG.warning(
401
+ f"Computing joint probabilities for {len(target_columns)} targets "
402
+ f"results in {total_cardinality:,} total probability values per sample. "
403
+ f"Computation complexity grows exponentially with the number of targets. "
404
+ f"Consider computing probabilities for targets separately if this takes too long."
405
+ )
361
406
 
362
407
  # Initialize with first target: P(col1)
363
408
  first_target_df = _generate_marginal_probs(
@@ -371,6 +416,7 @@ def predict_proba(
371
416
  ctx_data=ctx_data,
372
417
  ctx_stats=ctx_stats,
373
418
  fixed_probs=fixed_probs,
419
+ enable_flexible_generation=enable_flexible_generation,
374
420
  )
375
421
  _LOG.info(f"Generated P({target_columns[0]}) with shape {first_target_df.shape}")
376
422
 
@@ -452,6 +498,7 @@ def predict_proba(
452
498
  ctx_data=batched_ctx_data,
453
499
  ctx_stats=ctx_stats,
454
500
  fixed_probs=fixed_probs,
501
+ enable_flexible_generation=enable_flexible_generation,
455
502
  ) # DataFrame: (n_samples * num_combos, current_card)
456
503
 
457
504
  # Extract probabilities for each combo and compute joint probabilities
@@ -479,3 +526,72 @@ def predict_proba(
479
526
 
480
527
  _LOG.info(f"PREDICT_PROBA finished: returned probabilities for {len(probs_df)} samples")
481
528
  return probs_df
529
+
530
+
531
+ @torch.no_grad()
532
+ def log_prob(
533
+ *,
534
+ workspace: Workspace,
535
+ data: pd.DataFrame,
536
+ ctx_data: pd.DataFrame | None = None,
537
+ device: torch.device | str | None = None,
538
+ ) -> np.ndarray:
539
+ """
540
+ Compute log probability of full observations.
541
+
542
+ This function computes P(observation | model) - the likelihood of the observation
543
+ under the trained model. For autoregressive models:
544
+
545
+ log P(x1, x2, ..., xn) = log P(x1) + log P(x2|x1) + ... + log P(xn|x1,...,xn-1)
546
+
547
+ Each term is the probability the model assigns to the actual observed value.
548
+
549
+ Supports all encoding types including multi-sub-column encodings like numeric_digit,
550
+ and both flat and sequential models.
551
+
552
+ Args:
553
+ workspace: Workspace object containing model and stats
554
+ data: DataFrame with ALL columns containing observed values
555
+ ctx_data: Optional context data (for models with context)
556
+ device: Device to run inference on ('cuda' or 'cpu'). Defaults to 'cuda' if available.
557
+
558
+ Returns:
559
+ np.ndarray of shape (n_samples,) with log probability per row.
560
+ Values are <= 0 (log probabilities).
561
+ """
562
+ _LOG.info("LOG_PROB started")
563
+
564
+ # Initialize model
565
+ model, tgt_stats, ctx_stats, _, all_columns, _, device, enable_flexible_generation = _initialize_model(
566
+ workspace=workspace,
567
+ device=device,
568
+ )
569
+
570
+ # Check column order of input data when flexible generation is disabled
571
+ if not enable_flexible_generation:
572
+ check_column_order(list(data.columns), all_columns)
573
+
574
+ # Encode full data to get observed codes for all columns
575
+ full_encoded, _, _ = encode_df(df=data, stats=tgt_stats)
576
+
577
+ n_samples = len(data)
578
+ _LOG.info(f"Computing log probabilities for {n_samples} samples")
579
+
580
+ # Build batch dict with ALL encoded values as tensors
581
+ batch_dict: dict[str, torch.Tensor] = {}
582
+ for col in full_encoded.columns:
583
+ batch_dict[col] = torch.tensor(full_encoded[col].values, dtype=torch.long, device=device).unsqueeze(-1)
584
+
585
+ # Add context data if provided
586
+ if ctx_data is not None and ctx_stats:
587
+ ctx_inputs, _, _ = prepare_context_inputs(ctx_data=ctx_data, ctx_stats=ctx_stats, device=device)
588
+ batch_dict.update(ctx_inputs)
589
+
590
+ # Use the training loss calculation directly - handles both flat and sequential models
591
+ losses = _calculate_sample_losses(model, batch_dict)
592
+
593
+ # Negate loss to get log probability (loss = -log_prob)
594
+ log_probs = -losses.cpu().numpy()
595
+
596
+ _LOG.info(f"LOG_PROB finished: computed log probs for {n_samples} samples")
597
+ return log_probs
@@ -19,6 +19,7 @@ Split original data for training and validation.
19
19
  import logging
20
20
  import time
21
21
  import warnings
22
+ from collections.abc import Callable
22
23
  from pathlib import Path
23
24
 
24
25
  import numpy as np
@@ -73,7 +74,7 @@ def split(
73
74
  tgt_encoding_types: dict[str, str | ModelEncodingType] | None = None,
74
75
  ctx_encoding_types: dict[str, str | ModelEncodingType] | None = None,
75
76
  n_partitions: int = 1,
76
- trn_val_split: float = 0.8,
77
+ trn_val_split: float | Callable[[pd.Series], tuple[pd.Series, pd.Series]] = 0.8,
77
78
  workspace_dir: str | Path = "engine-ws",
78
79
  update_progress: ProgressCallback | None = None,
79
80
  ) -> None:
@@ -99,7 +100,8 @@ def split(
99
100
  tgt_encoding_types: Encoding types for columns in the target data (excluding key columns).
100
101
  ctx_encoding_types: Encoding types for columns in the context data (excluding key columns).
101
102
  n_partitions: Number of partitions to split the data into.
102
- trn_val_split: Fraction of data to use for training, with the remaining data used for validation.
103
+ trn_val_split: Fraction of data to use for training (0 < value < 1), or a callable
104
+ that takes keys as input and returns (trn_keys, val_keys) tuple.
103
105
  workspace_dir: Path to the workspace directory where files will be created.
104
106
  update_progress: A custom progress callback.
105
107
  """
@@ -189,13 +191,17 @@ def split(
189
191
  keys = tgt_data[tgt_context_key].drop_duplicates()
190
192
  else:
191
193
  keys = ctx_data[ctx_primary_key]
192
- # shuffle keys
193
- keys = keys.sample(frac=1)
194
- # split randomly into trn and val
195
- assert 0 < trn_val_split < 1, f"invalid trn_val_split: {trn_val_split}"
196
- trn_cnt = round(trn_val_split * len(keys))
197
- trn_keys = keys[:trn_cnt]
198
- val_keys = keys[trn_cnt:]
194
+ # split into trn and val
195
+ if callable(trn_val_split):
196
+ trn_keys, val_keys = trn_val_split(keys)
197
+ else:
198
+ # shuffle keys
199
+ keys = keys.sample(frac=1)
200
+ # split randomly into trn and val
201
+ assert 0 < trn_val_split < 1, f"invalid trn_val_split: {trn_val_split}"
202
+ trn_cnt = round(trn_val_split * len(keys))
203
+ trn_keys = keys[:trn_cnt]
204
+ val_keys = keys[trn_cnt:]
199
205
 
200
206
  def save_partition(
201
207
  df: pd.DataFrame, path: Path, key: str, sel_trn_keys: np.ndarray, sel_val_keys: np.ndarray, idx: int
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "mostlyai-engine"
3
- version = "2.2.0"
3
+ version = "2.3.1"
4
4
  description = "Synthetic Data Engine"
5
5
  authors = [{ name = "MOSTLY AI", email = "dev@mostly.ai" }]
6
6
  requires-python = ">=3.10"
File without changes