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.
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/PKG-INFO +43 -5
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/README.md +42 -4
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/__init__.py +1 -1
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/argn.py +99 -64
- mostlyai_engine-2.3.0/mostlyai/engine/_tabular/common.py +315 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/generation.py +34 -159
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/interface.py +178 -151
- mostlyai_engine-2.3.0/mostlyai/engine/_tabular/probability.py +597 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/splitting.py +15 -9
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/pyproject.toml +1 -1
- mostlyai_engine-2.1.0/mostlyai/engine/_tabular/common.py +0 -37
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/.gitignore +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/LICENSE +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_common.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_dtypes.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/__init__.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/common.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/encoding.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/base.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/generation.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/interface.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/lstm.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/training.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_memory.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/__init__.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/encoding.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/fairness.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/training.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_training_utils.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_workspace.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/analysis.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/domain.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/encoding.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/generation.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/logging.py +0 -0
- {mostlyai_engine-2.1.0 → mostlyai_engine-2.3.0}/mostlyai/engine/random_state.py +0 -0
- {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.
|
|
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
|
|
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
|
|
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.
|
|
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[
|
|
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
|
-
|
|
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
|
-
|
|
1051
|
-
|
|
1052
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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] =
|
|
1107
|
+
probs[sub_col] = probs_tensor
|
|
1081
1108
|
|
|
1082
|
-
# apply fairness
|
|
1109
|
+
# apply fairness transforms
|
|
1083
1110
|
if fairness_transforms:
|
|
1084
|
-
|
|
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
|
-
|
|
1101
|
-
|
|
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
|
-
|
|
1104
|
-
|
|
1105
|
-
|
|
1106
|
-
|
|
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
|
-
|
|
1111
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|