mostlyai-engine 2.3.0__tar.gz → 2.3.2__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.3.0 → mostlyai_engine-2.3.2}/PKG-INFO +4 -1
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/README.md +3 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/__init__.py +1 -1
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/common.py +22 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/generation.py +6 -37
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/interface.py +3 -3
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/probability.py +31 -37
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/pyproject.toml +1 -1
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/.gitignore +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/LICENSE +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_common.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_dtypes.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/__init__.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/text.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/__init__.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/common.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/encoding.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/engine/__init__.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/engine/base.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/generation.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/interface.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/lstm.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/training.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_memory.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/__init__.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/argn.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/encoding.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/fairness.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_tabular/training.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_training_utils.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_workspace.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/analysis.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/domain.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/encoding.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/generation.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/logging.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/random_state.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/splitting.py +0 -0
- {mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/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.
|
|
3
|
+
Version: 2.3.2
|
|
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
|
|
@@ -251,6 +251,9 @@ Compute log likelihood of observations:
|
|
|
251
251
|
```python
|
|
252
252
|
# compute log probability for each observation
|
|
253
253
|
log_probs = argn.log_prob(data_test)
|
|
254
|
+
|
|
255
|
+
# list top 10 outliers
|
|
256
|
+
data_test.iloc[log_probs.argsort()[:10]]
|
|
254
257
|
```
|
|
255
258
|
|
|
256
259
|
## TabularARGN for Sequential Data
|
|
@@ -200,6 +200,9 @@ Compute log likelihood of observations:
|
|
|
200
200
|
```python
|
|
201
201
|
# compute log probability for each observation
|
|
202
202
|
log_probs = argn.log_prob(data_test)
|
|
203
|
+
|
|
204
|
+
# list top 10 outliers
|
|
205
|
+
data_test.iloc[log_probs.argsort()[:10]]
|
|
203
206
|
```
|
|
204
207
|
|
|
205
208
|
## TabularARGN for Sequential Data
|
|
@@ -34,7 +34,7 @@ __all__ = [
|
|
|
34
34
|
"TabularARGN",
|
|
35
35
|
"LanguageModel",
|
|
36
36
|
]
|
|
37
|
-
__version__ = "2.3.
|
|
37
|
+
__version__ = "2.3.2"
|
|
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.*")
|
|
@@ -242,6 +242,28 @@ def check_column_order(
|
|
|
242
242
|
)
|
|
243
243
|
|
|
244
244
|
|
|
245
|
+
def get_argn_column_names(column_stats: dict, columns: list[str]) -> list[str]:
|
|
246
|
+
"""
|
|
247
|
+
Convert original column names to internal ARGN column names.
|
|
248
|
+
|
|
249
|
+
Args:
|
|
250
|
+
column_stats: Column statistics dict (e.g., tgt_stats["columns"])
|
|
251
|
+
columns: List of original column names
|
|
252
|
+
|
|
253
|
+
Returns:
|
|
254
|
+
List of ARGN column names (e.g., ['tgt:t0/c0', 'tgt:t1/c1'])
|
|
255
|
+
"""
|
|
256
|
+
return [
|
|
257
|
+
get_argn_name(
|
|
258
|
+
argn_processor=column_stats[col][ARGN_PROCESSOR],
|
|
259
|
+
argn_table=column_stats[col][ARGN_TABLE],
|
|
260
|
+
argn_column=column_stats[col][ARGN_COLUMN],
|
|
261
|
+
)
|
|
262
|
+
for col in columns
|
|
263
|
+
if col in column_stats
|
|
264
|
+
]
|
|
265
|
+
|
|
266
|
+
|
|
245
267
|
def fix_rare_token_probs(
|
|
246
268
|
stats: dict,
|
|
247
269
|
rare_category_replacement_method: RareCategoryReplacementMethod | None = None,
|
|
@@ -81,6 +81,7 @@ from mostlyai.engine._tabular.common import (
|
|
|
81
81
|
check_column_order,
|
|
82
82
|
create_and_load_model,
|
|
83
83
|
fix_rare_token_probs,
|
|
84
|
+
get_argn_column_names,
|
|
84
85
|
load_model_artifacts,
|
|
85
86
|
prepare_context_inputs,
|
|
86
87
|
resolve_device,
|
|
@@ -127,40 +128,20 @@ def _resolve_gen_column_order(
|
|
|
127
128
|
|
|
128
129
|
if imputation:
|
|
129
130
|
# imputed columns should be at the end in the generation model
|
|
130
|
-
imputation_argn =
|
|
131
|
-
get_argn_name(
|
|
132
|
-
argn_processor=column_stats[col][ARGN_PROCESSOR],
|
|
133
|
-
argn_table=column_stats[col][ARGN_TABLE],
|
|
134
|
-
argn_column=column_stats[col][ARGN_COLUMN],
|
|
135
|
-
)
|
|
136
|
-
for col in imputation.columns
|
|
137
|
-
if col in column_stats
|
|
138
|
-
]
|
|
131
|
+
imputation_argn = get_argn_column_names(column_stats, imputation.columns)
|
|
139
132
|
column_order = [c for c in column_order if c not in imputation_argn] + imputation_argn
|
|
140
133
|
else:
|
|
141
134
|
imputation_argn = []
|
|
142
135
|
|
|
143
136
|
if fairness:
|
|
144
137
|
# bring sensitive columns to the front and target column to the back
|
|
145
|
-
sensitive_columns_argn =
|
|
146
|
-
get_argn_name(
|
|
147
|
-
argn_processor=column_stats[col][ARGN_PROCESSOR],
|
|
148
|
-
argn_table=column_stats[col][ARGN_TABLE],
|
|
149
|
-
argn_column=column_stats[col][ARGN_COLUMN],
|
|
150
|
-
)
|
|
151
|
-
for col in fairness.sensitive_columns
|
|
152
|
-
if col in column_stats
|
|
153
|
-
]
|
|
138
|
+
sensitive_columns_argn = get_argn_column_names(column_stats, fairness.sensitive_columns)
|
|
154
139
|
# imputed sensitive columns should be after other usual sensitive columns
|
|
155
140
|
sensitive_columns_argn = [c for c in sensitive_columns_argn if c not in imputation_argn] + [
|
|
156
141
|
c for c in sensitive_columns_argn if c in imputation_argn
|
|
157
142
|
]
|
|
158
143
|
|
|
159
|
-
target_column_argn =
|
|
160
|
-
argn_processor=column_stats[fairness.target_column][ARGN_PROCESSOR],
|
|
161
|
-
argn_table=column_stats[fairness.target_column][ARGN_TABLE],
|
|
162
|
-
argn_column=column_stats[fairness.target_column][ARGN_COLUMN],
|
|
163
|
-
)
|
|
144
|
+
target_column_argn = get_argn_column_names(column_stats, [fairness.target_column])[0]
|
|
164
145
|
column_order = (
|
|
165
146
|
sensitive_columns_argn
|
|
166
147
|
+ [c for c in column_order if c not in sensitive_columns_argn + [target_column_argn]]
|
|
@@ -171,25 +152,13 @@ def _resolve_gen_column_order(
|
|
|
171
152
|
# rebalance column should be at the beginning in the generation model
|
|
172
153
|
# rebalancing has higher priority than imputation
|
|
173
154
|
if rebalancing.column in column_stats:
|
|
174
|
-
rebalance_column_argn =
|
|
175
|
-
argn_processor=column_stats[rebalancing.column][ARGN_PROCESSOR],
|
|
176
|
-
argn_table=column_stats[rebalancing.column][ARGN_TABLE],
|
|
177
|
-
argn_column=column_stats[rebalancing.column][ARGN_COLUMN],
|
|
178
|
-
)
|
|
155
|
+
rebalance_column_argn = get_argn_column_names(column_stats, [rebalancing.column])[0]
|
|
179
156
|
column_order = [rebalance_column_argn] + [c for c in column_order if c != rebalance_column_argn]
|
|
180
157
|
|
|
181
158
|
if seed_data is not None:
|
|
182
159
|
# seed_data columns should be at the beginning in the generation model
|
|
183
160
|
# seed_data has higher priority than rebalancing and imputation
|
|
184
|
-
seed_columns_argn =
|
|
185
|
-
get_argn_name(
|
|
186
|
-
argn_processor=column_stats[col][ARGN_PROCESSOR],
|
|
187
|
-
argn_table=column_stats[col][ARGN_TABLE],
|
|
188
|
-
argn_column=column_stats[col][ARGN_COLUMN],
|
|
189
|
-
)
|
|
190
|
-
for col in seed_data.columns
|
|
191
|
-
if col in column_stats
|
|
192
|
-
]
|
|
161
|
+
seed_columns_argn = get_argn_column_names(column_stats, list(seed_data.columns))
|
|
193
162
|
column_order = seed_columns_argn + [c for c in column_order if c not in seed_columns_argn]
|
|
194
163
|
|
|
195
164
|
if POSITIONAL_COLUMN in column_order:
|
|
@@ -737,7 +737,7 @@ class TabularARGN(BaseEstimator):
|
|
|
737
737
|
X: pd.DataFrame,
|
|
738
738
|
ctx_data: pd.DataFrame | None = None,
|
|
739
739
|
**kwargs,
|
|
740
|
-
) ->
|
|
740
|
+
) -> np.ndarray:
|
|
741
741
|
"""
|
|
742
742
|
Compute log probability of observations.
|
|
743
743
|
|
|
@@ -754,8 +754,8 @@ class TabularARGN(BaseEstimator):
|
|
|
754
754
|
**kwargs: Additional generation parameters (device, etc.).
|
|
755
755
|
|
|
756
756
|
Returns:
|
|
757
|
-
|
|
758
|
-
|
|
757
|
+
np.ndarray of shape (n_samples,) with log probability per row.
|
|
758
|
+
Values are <= 0 (log probabilities).
|
|
759
759
|
More negative values indicate less likely samples.
|
|
760
760
|
|
|
761
761
|
Raises:
|
|
@@ -34,7 +34,6 @@ from mostlyai.engine._common import (
|
|
|
34
34
|
get_argn_name,
|
|
35
35
|
get_cardinalities,
|
|
36
36
|
get_columns_from_cardinalities,
|
|
37
|
-
get_sub_columns_from_cardinalities,
|
|
38
37
|
)
|
|
39
38
|
from mostlyai.engine._encoding_types.tabular.numeric import (
|
|
40
39
|
NUMERIC_BINNED_MAX_TOKEN,
|
|
@@ -46,6 +45,7 @@ from mostlyai.engine._tabular.common import (
|
|
|
46
45
|
check_column_order,
|
|
47
46
|
create_and_load_model,
|
|
48
47
|
fix_rare_token_probs,
|
|
48
|
+
get_argn_column_names,
|
|
49
49
|
load_model_artifacts,
|
|
50
50
|
prepare_context_inputs,
|
|
51
51
|
resolve_device,
|
|
@@ -243,29 +243,25 @@ def _generate_marginal_probs(
|
|
|
243
243
|
seed_encoded: pd.DataFrame,
|
|
244
244
|
target_column: str,
|
|
245
245
|
tgt_stats: dict,
|
|
246
|
-
|
|
247
|
-
all_columns: list[str],
|
|
246
|
+
seed_columns: list[str],
|
|
248
247
|
device: torch.device,
|
|
249
248
|
ctx_data: pd.DataFrame | None = None,
|
|
250
249
|
ctx_stats: dict | None = None,
|
|
251
250
|
fixed_probs: dict | None = None,
|
|
252
|
-
|
|
253
|
-
) -> np.ndarray:
|
|
251
|
+
) -> pd.DataFrame:
|
|
254
252
|
"""
|
|
255
253
|
Generate P(target | seed_features, context).
|
|
256
254
|
|
|
257
255
|
Args:
|
|
258
256
|
model: The generative model
|
|
259
257
|
seed_encoded: Encoded seed features (may include previous targets)
|
|
260
|
-
target_column: Target column to predict
|
|
258
|
+
target_column: Target column to predict (original column name)
|
|
261
259
|
tgt_stats: Target statistics
|
|
262
|
-
|
|
263
|
-
all_columns: All columns in training order
|
|
260
|
+
seed_columns: Seed column names in original format, in correct order
|
|
264
261
|
device: Device for computation
|
|
265
262
|
ctx_data: Optional context data
|
|
266
263
|
ctx_stats: Optional context statistics (required if ctx_data provided)
|
|
267
264
|
fixed_probs: Optional fixed probabilities for rare token handling
|
|
268
|
-
enable_flexible_generation: Whether flexible generation is enabled
|
|
269
265
|
|
|
270
266
|
Returns:
|
|
271
267
|
DataFrame of shape (n_samples, cardinality) with probabilities and column names
|
|
@@ -275,16 +271,10 @@ def _generate_marginal_probs(
|
|
|
275
271
|
|
|
276
272
|
# Build fixed_values dict from seed_encoded
|
|
277
273
|
seed_batch_dict = {}
|
|
278
|
-
seed_columns = set()
|
|
279
274
|
|
|
280
275
|
# Add all columns from seed_encoded
|
|
281
|
-
for sub_col in
|
|
282
|
-
|
|
283
|
-
seed_batch_dict[sub_col] = torch.tensor(seed_encoded[sub_col].values, dtype=torch.long, device=device)
|
|
284
|
-
parts = sub_col.split("/")
|
|
285
|
-
if len(parts) >= 2:
|
|
286
|
-
col_name = parts[-2]
|
|
287
|
-
seed_columns.add(col_name)
|
|
276
|
+
for sub_col in seed_encoded.columns:
|
|
277
|
+
seed_batch_dict[sub_col] = torch.tensor(seed_encoded[sub_col].values, dtype=torch.long, device=device)
|
|
288
278
|
|
|
289
279
|
# Get target sub-columns
|
|
290
280
|
target_sub_cols = [
|
|
@@ -297,12 +287,12 @@ def _generate_marginal_probs(
|
|
|
297
287
|
for sub_col in target_stats["cardinalities"].keys()
|
|
298
288
|
]
|
|
299
289
|
|
|
300
|
-
#
|
|
301
|
-
|
|
290
|
+
# Convert column names to ARGN format for model
|
|
291
|
+
seed_columns_argn = get_argn_column_names(tgt_stats["columns"], seed_columns)
|
|
292
|
+
target_argn_name = get_argn_column_names(tgt_stats["columns"], [target_column])[0]
|
|
302
293
|
|
|
303
|
-
#
|
|
304
|
-
|
|
305
|
-
check_column_order(gen_column_order, all_columns)
|
|
294
|
+
# Determine column order: seed columns + target (in training order)
|
|
295
|
+
gen_column_order = seed_columns_argn + [target_argn_name]
|
|
306
296
|
|
|
307
297
|
# Prepare context inputs if provided
|
|
308
298
|
if ctx_data is not None and ctx_stats is not None:
|
|
@@ -320,11 +310,6 @@ def _generate_marginal_probs(
|
|
|
320
310
|
column_order=gen_column_order,
|
|
321
311
|
)
|
|
322
312
|
|
|
323
|
-
# Extract probabilities (assuming single sub-column for simplicity)
|
|
324
|
-
# For multi-sub-column cases, we'd need to handle differently
|
|
325
|
-
if len(target_sub_cols) != 1:
|
|
326
|
-
raise NotImplementedError("Multi-target joint probabilities currently only support single sub-column targets")
|
|
327
|
-
|
|
328
313
|
probs_array = probs_dct[target_sub_cols[0]].cpu().numpy()
|
|
329
314
|
|
|
330
315
|
# Create DataFrame with filtered and formatted columns
|
|
@@ -377,6 +362,16 @@ def predict_proba(
|
|
|
377
362
|
)
|
|
378
363
|
)
|
|
379
364
|
|
|
365
|
+
# Get seed column names (needed for column order check and _generate_marginal_probs)
|
|
366
|
+
seed_columns = list(seed_data.columns)
|
|
367
|
+
|
|
368
|
+
# Check column order when flexible generation is disabled
|
|
369
|
+
if not enable_flexible_generation:
|
|
370
|
+
seed_columns_argn = get_argn_column_names(tgt_stats["columns"], seed_columns)
|
|
371
|
+
target_columns_argn = get_argn_column_names(tgt_stats["columns"], target_columns)
|
|
372
|
+
gen_column_order = seed_columns_argn + target_columns_argn
|
|
373
|
+
check_column_order(gen_column_order, all_columns)
|
|
374
|
+
|
|
380
375
|
# Encode seed data (features to condition on) - common for both single and multi-target
|
|
381
376
|
# seed_data should NOT include any target columns
|
|
382
377
|
seed_encoded, _, _ = encode_df(
|
|
@@ -410,13 +405,11 @@ def predict_proba(
|
|
|
410
405
|
seed_encoded=seed_encoded,
|
|
411
406
|
target_column=target_columns[0],
|
|
412
407
|
tgt_stats=tgt_stats,
|
|
413
|
-
|
|
414
|
-
all_columns=all_columns,
|
|
408
|
+
seed_columns=seed_columns,
|
|
415
409
|
device=device,
|
|
416
410
|
ctx_data=ctx_data,
|
|
417
411
|
ctx_stats=ctx_stats,
|
|
418
412
|
fixed_probs=fixed_probs,
|
|
419
|
-
enable_flexible_generation=enable_flexible_generation,
|
|
420
413
|
)
|
|
421
414
|
_LOG.info(f"Generated P({target_columns[0]}) with shape {first_target_df.shape}")
|
|
422
415
|
|
|
@@ -486,19 +479,20 @@ def predict_proba(
|
|
|
486
479
|
if ctx_data is not None:
|
|
487
480
|
batched_ctx_data = pd.concat([ctx_data] * num_prev_combos, ignore_index=True)
|
|
488
481
|
|
|
482
|
+
# Compute extended seed_columns including previous targets
|
|
483
|
+
extended_seed_columns = seed_columns + target_columns[:target_idx]
|
|
484
|
+
|
|
489
485
|
# Single batched forward pass for all combinations
|
|
490
486
|
all_conditional_df = _generate_marginal_probs(
|
|
491
487
|
model=model,
|
|
492
488
|
seed_encoded=batched_seed,
|
|
493
489
|
target_column=target_col,
|
|
494
490
|
tgt_stats=tgt_stats,
|
|
495
|
-
|
|
496
|
-
all_columns=all_columns,
|
|
491
|
+
seed_columns=extended_seed_columns,
|
|
497
492
|
device=device,
|
|
498
493
|
ctx_data=batched_ctx_data,
|
|
499
494
|
ctx_stats=ctx_stats,
|
|
500
495
|
fixed_probs=fixed_probs,
|
|
501
|
-
enable_flexible_generation=enable_flexible_generation,
|
|
502
496
|
) # DataFrame: (n_samples * num_combos, current_card)
|
|
503
497
|
|
|
504
498
|
# Extract probabilities for each combo and compute joint probabilities
|
|
@@ -535,7 +529,7 @@ def log_prob(
|
|
|
535
529
|
data: pd.DataFrame,
|
|
536
530
|
ctx_data: pd.DataFrame | None = None,
|
|
537
531
|
device: torch.device | str | None = None,
|
|
538
|
-
) ->
|
|
532
|
+
) -> np.ndarray:
|
|
539
533
|
"""
|
|
540
534
|
Compute log probability of full observations.
|
|
541
535
|
|
|
@@ -556,8 +550,8 @@ def log_prob(
|
|
|
556
550
|
device: Device to run inference on ('cuda' or 'cpu'). Defaults to 'cuda' if available.
|
|
557
551
|
|
|
558
552
|
Returns:
|
|
559
|
-
|
|
560
|
-
|
|
553
|
+
np.ndarray of shape (n_samples,) with log probability per row.
|
|
554
|
+
Values are <= 0 (log probabilities).
|
|
561
555
|
"""
|
|
562
556
|
_LOG.info("LOG_PROB started")
|
|
563
557
|
|
|
@@ -594,4 +588,4 @@ def log_prob(
|
|
|
594
588
|
log_probs = -losses.cpu().numpy()
|
|
595
589
|
|
|
596
590
|
_LOG.info(f"LOG_PROB finished: computed log probs for {n_samples} samples")
|
|
597
|
-
return
|
|
591
|
+
return log_probs
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/numeric.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/language/text.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/character.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/datetime.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/itt.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/lat_long.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_encoding_types/tabular/numeric.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/engine/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/engine/hf_engine.py
RENAMED
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/mostlyai/engine/_language/engine/vllm_engine.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{mostlyai_engine-2.3.0 → mostlyai_engine-2.3.2}/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
|
|
File without changes
|