mostlyai-engine 2.1.0__tar.gz → 2.3.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 (55) hide show
  1. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/PKG-INFO +43 -5
  2. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/README.md +42 -4
  3. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/__init__.py +1 -1
  4. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/argn.py +99 -64
  5. mostlyai_engine-2.3.0/mostlyai/engine/_tabular/common.py +315 -0
  6. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/generation.py +34 -159
  7. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/interface.py +178 -151
  8. mostlyai_engine-2.3.0/mostlyai/engine/_tabular/probability.py +597 -0
  9. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/splitting.py +15 -9
  10. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/pyproject.toml +1 -1
  11. mostlyai_engine-2.1.0/mostlyai/engine/_tabular/common.py +0 -37
  12. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/.gitignore +0 -0
  13. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/LICENSE +0 -0
  14. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_common.py +0 -0
  15. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_dtypes.py +0 -0
  16. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  17. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  18. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  19. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  20. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  21. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  22. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  23. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  24. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  25. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  26. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  27. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  28. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  29. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/__init__.py +0 -0
  30. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/common.py +0 -0
  31. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/encoding.py +0 -0
  32. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
  33. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/base.py +0 -0
  34. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  35. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  36. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/generation.py +0 -0
  37. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/interface.py +0 -0
  38. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/lstm.py +0 -0
  39. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  40. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/training.py +0 -0
  41. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
  42. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_memory.py +0 -0
  43. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/__init__.py +0 -0
  44. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/encoding.py +0 -0
  45. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/fairness.py +0 -0
  46. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/training.py +0 -0
  47. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_training_utils.py +0 -0
  48. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_workspace.py +0 -0
  49. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/analysis.py +0 -0
  50. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/domain.py +0 -0
  51. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/encoding.py +0 -0
  52. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/generation.py +0 -0
  53. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/logging.py +0 -0
  54. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/random_state.py +0 -0
  55. {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/training.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mostlyai-engine
3
- Version: 2.1.0
3
+ Version: 2.3.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
@@ -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,46 @@ 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
+
218
256
  ## TabularARGN for Sequential Data
219
257
 
220
258
  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,46 @@ 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
+
167
205
  ## TabularARGN for Sequential Data
168
206
 
169
207
  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.1.0"
37
+ __version__ = "2.3.0"
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.*")
@@ -972,10 +972,63 @@ class FlatModel(nn.Module):
972
972
 
973
973
  return context
974
974
 
975
+ def _initialize_generation(self, x, batch_size, effective_column_order):
976
+ """Initialize context, embeddings, and sub-column order for generation/probs mode."""
977
+ # forward pass through context compressor
978
+ context = self.context_compressor(x)
979
+ context = self._handle_context(context)
980
+
981
+ # initialize embeddings
982
+ tgt_embeds = self.embedders.zero_mask(batch_size)
983
+ tgt_col_embeds = self.column_embedders.zero_mask(batch_size)
984
+ col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
985
+
986
+ # determine sub-column order
987
+ column_order = effective_column_order or self.tgt_columns
988
+ sub_column_order = [sub_col for col in column_order for sub_col in self.tgt_column_sub_columns[col]]
989
+
990
+ return context, tgt_embeds, tgt_col_embeds, col_embeddings, sub_column_order
991
+
992
+ def _update_embeddings(self, sub_col, out, tgt_embeds, tgt_col_embeds, col_embeddings):
993
+ """Update sub-column and column embeddings after setting a value.
994
+
995
+ Returns updated col_embeddings if this sub-column completes a column,
996
+ otherwise returns the unchanged col_embeddings.
997
+ """
998
+ lookup = self.tgt_sub_columns_lookup[sub_col]
999
+
1000
+ # update current sub column embedding
1001
+ tgt_embeds[sub_col] = self.embedders.get(sub_col)(out)
1002
+
1003
+ # update current column embedding if this is the last sub-column
1004
+ if sub_col in self.last_sub_cols:
1005
+ col_sub_cols = self.tgt_column_sub_columns[lookup.col_name]
1006
+ col_embed_in = torch.cat([tgt_embeds[sc] for sc in col_sub_cols], dim=-1)
1007
+ tgt_col_embeds[lookup.col_name] = self.column_embedders.get(lookup.col_name)(col_embed_in)
1008
+ col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1009
+
1010
+ return col_embeddings
1011
+
1012
+ def _compute_logits(self, sub_col, context, col_embeddings, tgt_embeds):
1013
+ """Compute logits for a sub-column given context and previous embeddings."""
1014
+ lookup = self.tgt_sub_columns_lookup[sub_col]
1015
+
1016
+ # collect previous sub column embeddings for current column
1017
+ prev_sub_col_embeds = [
1018
+ tgt_embeds[sc] for sc in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1019
+ ]
1020
+
1021
+ # regressor + predictor
1022
+ regressor_in = context + [col_embeddings] + prev_sub_col_embeds
1023
+ xs = self.regressors(regressor_in, sub_col)
1024
+ xs = self.predictors(xs, sub_col)
1025
+
1026
+ return xs
1027
+
975
1028
  def forward(
976
1029
  self,
977
1030
  x,
978
- mode: Literal["trn", "gen"],
1031
+ mode: Literal["trn", "gen", "probs"],
979
1032
  batch_size: int | None = None,
980
1033
  fixed_probs=None,
981
1034
  fixed_values=None,
@@ -1020,7 +1073,7 @@ class FlatModel(nn.Module):
1020
1073
 
1021
1074
  # collect previous sub column embeddings for current column
1022
1075
  prev_sub_col_embeds = [
1023
- tgt_embeds[sub_col] for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1076
+ tgt_embeds[sc] for sc in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1024
1077
  ]
1025
1078
 
1026
1079
  # regressor
@@ -1033,84 +1086,68 @@ class FlatModel(nn.Module):
1033
1086
  # update output
1034
1087
  outputs[sub_col] = xs
1035
1088
 
1036
- else: # mode == "gen"
1037
- # forward pass through context compressor
1038
- context = self.context_compressor(x)
1039
- context = self._handle_context(context)
1040
-
1041
- # initialize sub column embeddings
1042
- tgt_embeds = self.embedders.zero_mask(batch_size)
1043
-
1044
- # initialize column embeddings
1045
- tgt_col_embeds = self.column_embedders.zero_mask(batch_size)
1046
-
1047
- # concatenate column embeddings
1048
- col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1089
+ return outputs, {}
1049
1090
 
1050
- # take sub columns in the specified generation order
1051
- column_order = effective_column_order or self.tgt_columns
1052
- sub_column_order = [sub_col for col in column_order for sub_col in self.tgt_column_sub_columns[col]]
1091
+ elif mode == "gen":
1092
+ context, tgt_embeds, tgt_col_embeds, col_embeddings, sub_column_order = self._initialize_generation(
1093
+ x, batch_size, effective_column_order
1094
+ )
1053
1095
 
1054
1096
  for sub_col in sub_column_order:
1055
- lookup = self.tgt_sub_columns_lookup[sub_col]
1056
-
1057
- # if sub column is fixed, skip sampling and use that value
1097
+ # handle fixed values
1058
1098
  if sub_col in fixed_values:
1059
1099
  out = fixed_values[sub_col]
1100
+ else:
1101
+ # compute probabilities and sample
1102
+ logits = self._compute_logits(sub_col, context, col_embeddings, tgt_embeds)
1103
+ probs_tensor = nn.Softmax(dim=-1)(logits)
1060
1104
 
1061
- else: # sample from distribution
1062
- # collect previous sub column embeddings for current column
1063
- prev_sub_col_embeds = [
1064
- tgt_embeds[sub_col]
1065
- for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1066
- ]
1067
-
1068
- # regressor
1069
- regressor_in = context + [col_embeddings] + prev_sub_col_embeds
1070
- xs = self.regressors(regressor_in, sub_col)
1071
-
1072
- # predictor
1073
- xs = self.predictors(xs, sub_col)
1074
-
1075
- # softmax to probs
1076
- xs = nn.Softmax(dim=-1)(xs)
1077
-
1078
- # keep probabilities (used e.g. for fairness)
1105
+ # optionally keep probabilities
1079
1106
  if sub_col in return_probs:
1080
- probs[sub_col] = xs
1107
+ probs[sub_col] = probs_tensor
1081
1108
 
1082
- # apply fairness transform when generating the target sub column
1109
+ # apply fairness transforms
1083
1110
  if fairness_transforms:
1084
- xs = apply_fairness_transforms(sub_col, xs, outputs, fairness_transforms)
1111
+ probs_tensor = apply_fairness_transforms(sub_col, probs_tensor, outputs, fairness_transforms)
1085
1112
 
1086
1113
  # sample
1087
1114
  out = torch.squeeze(
1088
- _sample(
1089
- probs=xs,
1090
- temperature=temperature,
1091
- top_p=top_p,
1092
- fixed_probs=fixed_probs.get(sub_col),
1093
- ),
1115
+ _sample(probs_tensor, temperature, top_p, fixed_probs.get(sub_col)),
1094
1116
  dim=-1,
1095
1117
  )
1096
1118
 
1097
- # update output
1119
+ # update output and embeddings
1098
1120
  outputs[sub_col] = out
1121
+ col_embeddings = self._update_embeddings(sub_col, out, tgt_embeds, tgt_col_embeds, col_embeddings)
1099
1122
 
1100
- # update current sub column embedding
1101
- tgt_embeds[sub_col] = self.embedders.get(sub_col)(out)
1123
+ # order outputs and return
1124
+ outputs = {sub_col: outputs[sub_col] for sub_col in self.tgt_sub_columns}
1125
+ return outputs, probs
1102
1126
 
1103
- # update current column embedding
1104
- if sub_col in self.last_sub_cols:
1105
- col_sub_cols = self.tgt_column_sub_columns[lookup.col_name]
1106
- col_embed_in = torch.cat([tgt_embeds[sc] for sc in col_sub_cols], dim=-1)
1107
- tgt_col_embeds[lookup.col_name] = self.column_embedders.get(lookup.col_name)(col_embed_in)
1108
- col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1127
+ elif mode == "probs":
1128
+ context, tgt_embeds, tgt_col_embeds, col_embeddings, sub_column_order = self._initialize_generation(
1129
+ x, batch_size, effective_column_order
1130
+ )
1109
1131
 
1110
- # order outputs according to tgt_sub_columns
1111
- outputs = {sub_col: outputs[sub_col] for sub_col in self.tgt_sub_columns}
1132
+ for sub_col in sub_column_order:
1133
+ # handle fixed values
1134
+ if sub_col in fixed_values:
1135
+ out = fixed_values[sub_col]
1136
+ # update embeddings to maintain correct autoregressive context
1137
+ col_embeddings = self._update_embeddings(sub_col, out, tgt_embeds, tgt_col_embeds, col_embeddings)
1138
+ else:
1139
+ # compute probabilities without sampling
1140
+ logits = self._compute_logits(sub_col, context, col_embeddings, tgt_embeds)
1141
+ probs_tensor = nn.Softmax(dim=-1)(logits)
1142
+
1143
+ # apply fixed_probs mask if provided
1144
+ if sub_col in fixed_probs:
1145
+ probs_tensor = _sampling_fixed_probs(probs_tensor, fixed_probs[sub_col])
1146
+
1147
+ # store probabilities (no sampling, no embedding updates)
1148
+ probs[sub_col] = probs_tensor
1112
1149
 
1113
- return outputs, probs
1150
+ return {}, probs
1114
1151
 
1115
1152
 
1116
1153
  class AttentionModule(nn.Module):
@@ -1334,8 +1371,7 @@ class SequentialModel(nn.Module):
1334
1371
 
1335
1372
  # collect previous sub column embeddings for current column
1336
1373
  prev_sub_col_embeds = {
1337
- sub_col: tgt_embeds[sub_col]
1338
- for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1374
+ sc: tgt_embeds[sc] for sc in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1339
1375
  }
1340
1376
  if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1341
1377
  # RIDX sub-columns should not see SLEN sub-columns
@@ -1405,8 +1441,7 @@ class SequentialModel(nn.Module):
1405
1441
  else: # sample from distribution
1406
1442
  # collect previous sub column embeddings for current column
1407
1443
  prev_sub_col_embeds = {
1408
- sub_col: tgt_embeds[sub_col]
1409
- for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1444
+ sc: tgt_embeds[sc] for sc in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1410
1445
  }
1411
1446
  if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1412
1447
  # RIDX sub-columns should not see SLEN sub-columns
@@ -0,0 +1,315 @@
1
+ # Copyright 2025 MOSTLY AI
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+
15
+ import logging
16
+ import time
17
+ from pathlib import Path
18
+
19
+ import pandas as pd
20
+ import torch
21
+
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
42
+
43
+ _LOG = logging.getLogger(__name__)
44
+
45
+ # Type alias for fixed probabilities
46
+ CodeProbabilities = dict[int, float]
47
+
48
+ DPLSTM_SUFFIXES: tuple = ("ih.weight", "ih.bias", "hh.weight", "hh.bias")
49
+
50
+
51
+ def load_model_weights(model: torch.nn.Module, path: Path, device: torch.device) -> None:
52
+ t0 = time.time()
53
+ incompatible_keys = model.load_state_dict(torch.load(f=path, map_location=device, weights_only=True), strict=False)
54
+ missing_keys = incompatible_keys.missing_keys
55
+ unexpected_keys = incompatible_keys.unexpected_keys
56
+ # for DP-trained models, we expect extra keys from the DPLSTM layers (which is fine to ignore because we use standard LSTM layers during generation)
57
+ # but if there're any other missing or unexpected keys, an error should be raised
58
+ if len(missing_keys) > 0 or any(not k.endswith(DPLSTM_SUFFIXES) for k in unexpected_keys):
59
+ raise RuntimeError(
60
+ f"failed to load model weights due to incompatibility: {missing_keys = }, {unexpected_keys = }"
61
+ )
62
+ _LOG.info(f"loaded model weights in {time.time() - t0:.2f}s")
63
+
64
+
65
+ def load_model_artifacts(workspace):
66
+ """
67
+ Load model configurations and statistics from workspace.
68
+
69
+ Returns:
70
+ Tuple of (model_config, tgt_stats, ctx_stats, is_sequential)
71
+ """
72
+ model_config = workspace.model_configs.read()
73
+ tgt_stats = workspace.tgt_stats.read()
74
+ ctx_stats = workspace.ctx_stats.read()
75
+ is_sequential = tgt_stats["is_sequential"]
76
+ return model_config, tgt_stats, ctx_stats, is_sequential
77
+
78
+
79
+ def resolve_device(device: torch.device | str | None) -> torch.device:
80
+ """
81
+ Resolve device to use for inference.
82
+
83
+ Args:
84
+ device: Device specification ('cuda', 'cpu', or None for auto-detect)
85
+
86
+ Returns:
87
+ torch.device instance
88
+ """
89
+ if device is None:
90
+ return torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
91
+ return torch.device(device)
92
+
93
+
94
+ def create_and_load_model(
95
+ workspace,
96
+ is_sequential: bool,
97
+ tgt_cardinalities: dict,
98
+ ctx_cardinalities: dict,
99
+ model_units,
100
+ ctx_seq_len_median: int | None,
101
+ column_order: list[str],
102
+ device: torch.device,
103
+ seq_len_median: int | None = None,
104
+ seq_len_max: int | None = None,
105
+ ):
106
+ """
107
+ Create model, load weights, and prepare for inference.
108
+
109
+ Args:
110
+ workspace: Workspace containing model weights
111
+ is_sequential: Whether to create SequentialModel or FlatModel
112
+ tgt_cardinalities: Target column cardinalities
113
+ ctx_cardinalities: Context column cardinalities
114
+ model_units: Model size configuration
115
+ ctx_seq_len_median: Median context sequence length
116
+ column_order: Order of columns for generation
117
+ device: Device to load model on
118
+ seq_len_median: Median sequence length (for sequential models)
119
+ seq_len_max: Maximum sequence length (for sequential models)
120
+
121
+ Returns:
122
+ Initialized model ready for inference
123
+ """
124
+ from mostlyai.engine._tabular.argn import FlatModel, SequentialModel, get_no_of_model_parameters
125
+
126
+ _LOG.info("Creating generative model")
127
+
128
+ if is_sequential:
129
+ model = SequentialModel(
130
+ tgt_cardinalities=tgt_cardinalities,
131
+ tgt_seq_len_median=seq_len_median,
132
+ tgt_seq_len_max=seq_len_max,
133
+ ctx_cardinalities=ctx_cardinalities,
134
+ ctxseq_len_median=ctx_seq_len_median,
135
+ model_size=model_units,
136
+ column_order=column_order,
137
+ device=device,
138
+ )
139
+ else:
140
+ model = FlatModel(
141
+ tgt_cardinalities=tgt_cardinalities,
142
+ ctx_cardinalities=ctx_cardinalities,
143
+ ctxseq_len_median=ctx_seq_len_median,
144
+ model_size=model_units,
145
+ column_order=column_order,
146
+ device=device,
147
+ )
148
+
149
+ no_of_model_params = get_no_of_model_parameters(model)
150
+ _LOG.info(f"{no_of_model_params=}")
151
+
152
+ if workspace.model_tabular_weights_path.exists():
153
+ load_model_weights(
154
+ model=model,
155
+ path=workspace.model_tabular_weights_path,
156
+ device=device,
157
+ )
158
+ else:
159
+ _LOG.warning("Model weights not found; using untrained model")
160
+
161
+ model.to(device)
162
+ model.eval()
163
+
164
+ return model
165
+
166
+
167
+ def prepare_context_inputs(
168
+ ctx_data: pd.DataFrame,
169
+ ctx_stats: dict,
170
+ device: torch.device | str,
171
+ ctx_primary_key: str | None = None,
172
+ ) -> tuple[dict[str, torch.Tensor], pd.DataFrame, str | None]:
173
+ """
174
+ Encode context data and prepare tensor inputs for model forward pass.
175
+
176
+ Handles both flat context (CTXFLT) and sequential context (CTXSEQ).
177
+
178
+ Args:
179
+ ctx_data: Context DataFrame to encode
180
+ ctx_stats: Context statistics from training
181
+ device: Device for tensor placement
182
+ ctx_primary_key: Optional primary key column for context
183
+
184
+ Returns:
185
+ Tuple of (context_tensors, encoded_dataframe, encoded_primary_key):
186
+ - context_tensors: Dict of CTXFLT/* and CTXSEQ/* tensors for model.context_compressor()
187
+ - encoded_dataframe: Encoded context DataFrame (for extracting keys if needed)
188
+ - encoded_primary_key: Name of encoded primary key column (None if not provided)
189
+ """
190
+
191
+ # Encode context data
192
+ ctx_encoded, ctx_primary_key_encoded, _ = encode_df(df=ctx_data, stats=ctx_stats, ctx_primary_key=ctx_primary_key)
193
+
194
+ # Pad empty sequences (required for model)
195
+ ctx_encoded = pad_ctx_sequences(ctx_encoded)
196
+
197
+ # Build flat context inputs (CTXFLT/*)
198
+ ctxflt_inputs = {
199
+ col: torch.unsqueeze(
200
+ torch.as_tensor(ctx_encoded[col].to_numpy(), device=device).type(torch.int),
201
+ dim=-1,
202
+ )
203
+ for col in ctx_encoded.columns
204
+ if col.startswith(CTXFLT)
205
+ }
206
+
207
+ # Build sequential context inputs (CTXSEQ/*)
208
+ ctxseq_inputs = {
209
+ col: torch.unsqueeze(
210
+ torch.nested.as_nested_tensor(
211
+ [torch.as_tensor(t, device=device).type(torch.int) for t in ctx_encoded[col]],
212
+ device=device,
213
+ ),
214
+ dim=-1,
215
+ )
216
+ for col in ctx_encoded.columns
217
+ if col.startswith(CTXSEQ)
218
+ }
219
+
220
+ # Merge and return with encoded dataframe
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