mostlyai-engine 2.2.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.2.0 → mostlyai_engine-2.3.0}/PKG-INFO +43 -5
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/README.md +42 -4
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/__init__.py +1 -1
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/common.py +117 -2
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/generation.py +6 -64
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/interface.py +52 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/probability.py +165 -49
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/splitting.py +15 -9
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/pyproject.toml +1 -1
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/.gitignore +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/LICENSE +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_common.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_dtypes.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/__init__.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/common.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/encoding.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/base.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/generation.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/interface.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/lstm.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/training.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_memory.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/__init__.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/argn.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/encoding.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/fairness.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_tabular/training.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_training_utils.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_workspace.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/analysis.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/domain.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/encoding.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/generation.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/logging.py +0 -0
- {mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/random_state.py +0 -0
- {mostlyai_engine-2.2.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.*")
|
|
@@ -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
|
|
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
|
-
|
|
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 =
|
|
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 =
|
|
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
|
+
) -> pd.DataFrame:
|
|
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
|
+
pd.DataFrame of shape (n_samples, 1) with log probability per row.
|
|
758
|
+
Column name is "log_prob". 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.
|
|
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
|
-
|
|
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
|
-
#
|
|
296
|
-
|
|
297
|
-
|
|
298
|
-
|
|
299
|
-
|
|
300
|
-
|
|
301
|
-
|
|
302
|
-
|
|
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
|
-
|
|
356
|
-
|
|
357
|
-
|
|
358
|
-
|
|
359
|
-
|
|
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
|
+
) -> pd.DataFrame:
|
|
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
|
+
pd.DataFrame of shape (n_samples, 1) with log probability per row.
|
|
560
|
+
Column name is "log_prob". 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 pd.DataFrame({"log_prob": 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
|
|
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
|
-
#
|
|
193
|
-
|
|
194
|
-
|
|
195
|
-
|
|
196
|
-
|
|
197
|
-
|
|
198
|
-
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/numeric.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/language/text.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/character.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/itt.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_encoding_types/tabular/numeric.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/hf_engine.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/engine/vllm_engine.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.2.0 → mostlyai_engine-2.3.0}/mostlyai/engine/_language/tokenizer_utils.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|