mostlyai-engine 2.0.1__tar.gz → 2.2.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.0.1 → mostlyai_engine-2.2.0}/PKG-INFO +1 -1
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/__init__.py +1 -1
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/argn.py +99 -64
- mostlyai_engine-2.2.0/mostlyai/engine/_tabular/common.py +200 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/generation.py +28 -95
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/interface.py +150 -75
- mostlyai_engine-2.2.0/mostlyai/engine/_tabular/probability.py +481 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/pyproject.toml +1 -1
- mostlyai_engine-2.0.1/mostlyai/engine/_tabular/common.py +0 -37
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/.gitignore +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/LICENSE +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/README.md +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_common.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_dtypes.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/__init__.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/common.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/encoding.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/base.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/generation.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/interface.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/lstm.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/training.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_memory.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/__init__.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/encoding.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/fairness.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/training.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_training_utils.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/_workspace.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/analysis.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/domain.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/encoding.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/generation.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/logging.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/random_state.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/splitting.py +0 -0
- {mostlyai_engine-2.0.1 → mostlyai_engine-2.2.0}/mostlyai/engine/training.py +0 -0
|
@@ -34,7 +34,7 @@ __all__ = [
|
|
|
34
34
|
"TabularARGN",
|
|
35
35
|
"LanguageModel",
|
|
36
36
|
]
|
|
37
|
-
__version__ = "2.0
|
|
37
|
+
__version__ = "2.2.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,200 @@
|
|
|
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 CTXFLT, CTXSEQ
|
|
23
|
+
|
|
24
|
+
_LOG = logging.getLogger(__name__)
|
|
25
|
+
|
|
26
|
+
DPLSTM_SUFFIXES: tuple = ("ih.weight", "ih.bias", "hh.weight", "hh.bias")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def load_model_weights(model: torch.nn.Module, path: Path, device: torch.device) -> None:
|
|
30
|
+
t0 = time.time()
|
|
31
|
+
incompatible_keys = model.load_state_dict(torch.load(f=path, map_location=device, weights_only=True), strict=False)
|
|
32
|
+
missing_keys = incompatible_keys.missing_keys
|
|
33
|
+
unexpected_keys = incompatible_keys.unexpected_keys
|
|
34
|
+
# 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)
|
|
35
|
+
# but if there're any other missing or unexpected keys, an error should be raised
|
|
36
|
+
if len(missing_keys) > 0 or any(not k.endswith(DPLSTM_SUFFIXES) for k in unexpected_keys):
|
|
37
|
+
raise RuntimeError(
|
|
38
|
+
f"failed to load model weights due to incompatibility: {missing_keys = }, {unexpected_keys = }"
|
|
39
|
+
)
|
|
40
|
+
_LOG.info(f"loaded model weights in {time.time() - t0:.2f}s")
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def load_model_artifacts(workspace):
|
|
44
|
+
"""
|
|
45
|
+
Load model configurations and statistics from workspace.
|
|
46
|
+
|
|
47
|
+
Returns:
|
|
48
|
+
Tuple of (model_config, tgt_stats, ctx_stats, is_sequential)
|
|
49
|
+
"""
|
|
50
|
+
model_config = workspace.model_configs.read()
|
|
51
|
+
tgt_stats = workspace.tgt_stats.read()
|
|
52
|
+
ctx_stats = workspace.ctx_stats.read()
|
|
53
|
+
is_sequential = tgt_stats["is_sequential"]
|
|
54
|
+
return model_config, tgt_stats, ctx_stats, is_sequential
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def resolve_device(device: torch.device | str | None) -> torch.device:
|
|
58
|
+
"""
|
|
59
|
+
Resolve device to use for inference.
|
|
60
|
+
|
|
61
|
+
Args:
|
|
62
|
+
device: Device specification ('cuda', 'cpu', or None for auto-detect)
|
|
63
|
+
|
|
64
|
+
Returns:
|
|
65
|
+
torch.device instance
|
|
66
|
+
"""
|
|
67
|
+
if device is None:
|
|
68
|
+
return torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
|
|
69
|
+
return torch.device(device)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def create_and_load_model(
|
|
73
|
+
workspace,
|
|
74
|
+
is_sequential: bool,
|
|
75
|
+
tgt_cardinalities: dict,
|
|
76
|
+
ctx_cardinalities: dict,
|
|
77
|
+
model_units,
|
|
78
|
+
ctx_seq_len_median: int | None,
|
|
79
|
+
column_order: list[str],
|
|
80
|
+
device: torch.device,
|
|
81
|
+
seq_len_median: int | None = None,
|
|
82
|
+
seq_len_max: int | None = None,
|
|
83
|
+
):
|
|
84
|
+
"""
|
|
85
|
+
Create model, load weights, and prepare for inference.
|
|
86
|
+
|
|
87
|
+
Args:
|
|
88
|
+
workspace: Workspace containing model weights
|
|
89
|
+
is_sequential: Whether to create SequentialModel or FlatModel
|
|
90
|
+
tgt_cardinalities: Target column cardinalities
|
|
91
|
+
ctx_cardinalities: Context column cardinalities
|
|
92
|
+
model_units: Model size configuration
|
|
93
|
+
ctx_seq_len_median: Median context sequence length
|
|
94
|
+
column_order: Order of columns for generation
|
|
95
|
+
device: Device to load model on
|
|
96
|
+
seq_len_median: Median sequence length (for sequential models)
|
|
97
|
+
seq_len_max: Maximum sequence length (for sequential models)
|
|
98
|
+
|
|
99
|
+
Returns:
|
|
100
|
+
Initialized model ready for inference
|
|
101
|
+
"""
|
|
102
|
+
from mostlyai.engine._tabular.argn import FlatModel, SequentialModel, get_no_of_model_parameters
|
|
103
|
+
|
|
104
|
+
_LOG.info("Creating generative model")
|
|
105
|
+
|
|
106
|
+
if is_sequential:
|
|
107
|
+
model = SequentialModel(
|
|
108
|
+
tgt_cardinalities=tgt_cardinalities,
|
|
109
|
+
tgt_seq_len_median=seq_len_median,
|
|
110
|
+
tgt_seq_len_max=seq_len_max,
|
|
111
|
+
ctx_cardinalities=ctx_cardinalities,
|
|
112
|
+
ctxseq_len_median=ctx_seq_len_median,
|
|
113
|
+
model_size=model_units,
|
|
114
|
+
column_order=column_order,
|
|
115
|
+
device=device,
|
|
116
|
+
)
|
|
117
|
+
else:
|
|
118
|
+
model = FlatModel(
|
|
119
|
+
tgt_cardinalities=tgt_cardinalities,
|
|
120
|
+
ctx_cardinalities=ctx_cardinalities,
|
|
121
|
+
ctxseq_len_median=ctx_seq_len_median,
|
|
122
|
+
model_size=model_units,
|
|
123
|
+
column_order=column_order,
|
|
124
|
+
device=device,
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
no_of_model_params = get_no_of_model_parameters(model)
|
|
128
|
+
_LOG.info(f"{no_of_model_params=}")
|
|
129
|
+
|
|
130
|
+
if workspace.model_tabular_weights_path.exists():
|
|
131
|
+
load_model_weights(
|
|
132
|
+
model=model,
|
|
133
|
+
path=workspace.model_tabular_weights_path,
|
|
134
|
+
device=device,
|
|
135
|
+
)
|
|
136
|
+
else:
|
|
137
|
+
_LOG.warning("Model weights not found; using untrained model")
|
|
138
|
+
|
|
139
|
+
model.to(device)
|
|
140
|
+
model.eval()
|
|
141
|
+
|
|
142
|
+
return model
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def prepare_context_inputs(
|
|
146
|
+
ctx_data: pd.DataFrame,
|
|
147
|
+
ctx_stats: dict,
|
|
148
|
+
device: torch.device | str,
|
|
149
|
+
ctx_primary_key: str | None = None,
|
|
150
|
+
) -> tuple[dict[str, torch.Tensor], pd.DataFrame, str | None]:
|
|
151
|
+
"""
|
|
152
|
+
Encode context data and prepare tensor inputs for model forward pass.
|
|
153
|
+
|
|
154
|
+
Handles both flat context (CTXFLT) and sequential context (CTXSEQ).
|
|
155
|
+
|
|
156
|
+
Args:
|
|
157
|
+
ctx_data: Context DataFrame to encode
|
|
158
|
+
ctx_stats: Context statistics from training
|
|
159
|
+
device: Device for tensor placement
|
|
160
|
+
ctx_primary_key: Optional primary key column for context
|
|
161
|
+
|
|
162
|
+
Returns:
|
|
163
|
+
Tuple of (context_tensors, encoded_dataframe, encoded_primary_key):
|
|
164
|
+
- context_tensors: Dict of CTXFLT/* and CTXSEQ/* tensors for model.context_compressor()
|
|
165
|
+
- encoded_dataframe: Encoded context DataFrame (for extracting keys if needed)
|
|
166
|
+
- encoded_primary_key: Name of encoded primary key column (None if not provided)
|
|
167
|
+
"""
|
|
168
|
+
from mostlyai.engine._tabular.encoding import encode_df, pad_ctx_sequences
|
|
169
|
+
|
|
170
|
+
# Encode context data
|
|
171
|
+
ctx_encoded, ctx_primary_key_encoded, _ = encode_df(df=ctx_data, stats=ctx_stats, ctx_primary_key=ctx_primary_key)
|
|
172
|
+
|
|
173
|
+
# Pad empty sequences (required for model)
|
|
174
|
+
ctx_encoded = pad_ctx_sequences(ctx_encoded)
|
|
175
|
+
|
|
176
|
+
# Build flat context inputs (CTXFLT/*)
|
|
177
|
+
ctxflt_inputs = {
|
|
178
|
+
col: torch.unsqueeze(
|
|
179
|
+
torch.as_tensor(ctx_encoded[col].to_numpy(), device=device).type(torch.int),
|
|
180
|
+
dim=-1,
|
|
181
|
+
)
|
|
182
|
+
for col in ctx_encoded.columns
|
|
183
|
+
if col.startswith(CTXFLT)
|
|
184
|
+
}
|
|
185
|
+
|
|
186
|
+
# Build sequential context inputs (CTXSEQ/*)
|
|
187
|
+
ctxseq_inputs = {
|
|
188
|
+
col: torch.unsqueeze(
|
|
189
|
+
torch.nested.as_nested_tensor(
|
|
190
|
+
[torch.as_tensor(t, device=device).type(torch.int) for t in ctx_encoded[col]],
|
|
191
|
+
device=device,
|
|
192
|
+
),
|
|
193
|
+
dim=-1,
|
|
194
|
+
)
|
|
195
|
+
for col in ctx_encoded.columns
|
|
196
|
+
if col.startswith(CTXSEQ)
|
|
197
|
+
}
|
|
198
|
+
|
|
199
|
+
# Merge and return with encoded dataframe
|
|
200
|
+
return (ctxflt_inputs | ctxseq_inputs), ctx_encoded, ctx_primary_key_encoded
|
|
@@ -78,10 +78,14 @@ from mostlyai.engine._tabular.argn import (
|
|
|
78
78
|
FlatModel,
|
|
79
79
|
ModelSize,
|
|
80
80
|
SequentialModel,
|
|
81
|
-
get_no_of_model_parameters,
|
|
82
81
|
)
|
|
83
|
-
from mostlyai.engine._tabular.common import
|
|
84
|
-
|
|
82
|
+
from mostlyai.engine._tabular.common import (
|
|
83
|
+
create_and_load_model,
|
|
84
|
+
load_model_artifacts,
|
|
85
|
+
prepare_context_inputs,
|
|
86
|
+
resolve_device,
|
|
87
|
+
)
|
|
88
|
+
from mostlyai.engine._tabular.encoding import encode_df
|
|
85
89
|
from mostlyai.engine._tabular.fairness import FairnessTransforms, get_fairness_transforms
|
|
86
90
|
from mostlyai.engine._workspace import Workspace, ensure_workspace_dir, reset_dir
|
|
87
91
|
from mostlyai.engine.domain import (
|
|
@@ -839,13 +843,10 @@ def generate(
|
|
|
839
843
|
output_path = workspace.generated_data_path
|
|
840
844
|
reset_dir(output_path)
|
|
841
845
|
|
|
842
|
-
model_configs = workspace
|
|
843
|
-
tgt_stats = workspace.tgt_stats.read()
|
|
844
|
-
is_sequential = tgt_stats["is_sequential"]
|
|
846
|
+
model_configs, tgt_stats, ctx_stats, is_sequential = load_model_artifacts(workspace)
|
|
845
847
|
_LOG.info(f"{is_sequential=}")
|
|
846
848
|
has_context = workspace.ctx_stats.path.exists()
|
|
847
849
|
_LOG.info(f"{has_context=}")
|
|
848
|
-
ctx_stats = workspace.ctx_stats.read()
|
|
849
850
|
|
|
850
851
|
# read model config
|
|
851
852
|
model_units = model_configs.get("model_units") or ModelSize.M
|
|
@@ -871,11 +872,7 @@ def generate(
|
|
|
871
872
|
_LOG.info(f"{len(ctx_sub_columns)=}")
|
|
872
873
|
|
|
873
874
|
# resolve device
|
|
874
|
-
device = (
|
|
875
|
-
torch.device(device)
|
|
876
|
-
if device is not None
|
|
877
|
-
else (torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu"))
|
|
878
|
-
)
|
|
875
|
+
device = resolve_device(device)
|
|
879
876
|
_LOG.info(f"{device=}")
|
|
880
877
|
|
|
881
878
|
tgt_primary_key = tgt_stats.get("keys", {}).get("primary_key")
|
|
@@ -912,7 +909,6 @@ def generate(
|
|
|
912
909
|
"The column order for generation does not match the column order from training, due to seed, rebalancing, fairness or imputation configs. "
|
|
913
910
|
"A change in column order is only permitted for models that were trained with `enable_flexible_generation=True`."
|
|
914
911
|
)
|
|
915
|
-
|
|
916
912
|
_LOG.info(f"{rare_category_replacement_method=}")
|
|
917
913
|
rare_token_fixed_probs = _fix_rare_token_probs(tgt_stats, rare_category_replacement_method)
|
|
918
914
|
imputation_fixed_probs = _fix_imputation_probs(tgt_stats, imputation)
|
|
@@ -1016,43 +1012,18 @@ def generate(
|
|
|
1016
1012
|
# init progress with total_count; +1 for the final decoding step
|
|
1017
1013
|
progress.update(completed=0, total=no_of_batches * (seq_len_max + 1))
|
|
1018
1014
|
|
|
1019
|
-
|
|
1020
|
-
|
|
1021
|
-
|
|
1022
|
-
|
|
1023
|
-
|
|
1024
|
-
|
|
1025
|
-
|
|
1026
|
-
|
|
1027
|
-
|
|
1028
|
-
|
|
1029
|
-
|
|
1030
|
-
|
|
1031
|
-
)
|
|
1032
|
-
else:
|
|
1033
|
-
model = FlatModel(
|
|
1034
|
-
tgt_cardinalities=tgt_cardinalities,
|
|
1035
|
-
ctx_cardinalities=ctx_cardinalities,
|
|
1036
|
-
ctxseq_len_median=ctx_seq_len_median,
|
|
1037
|
-
model_size=model_units,
|
|
1038
|
-
column_order=gen_column_order,
|
|
1039
|
-
device=device,
|
|
1040
|
-
)
|
|
1041
|
-
|
|
1042
|
-
no_of_model_params = get_no_of_model_parameters(model)
|
|
1043
|
-
_LOG.info(f"{no_of_model_params=}")
|
|
1044
|
-
|
|
1045
|
-
if workspace.model_tabular_weights_path.exists():
|
|
1046
|
-
load_model_weights(
|
|
1047
|
-
model=model,
|
|
1048
|
-
path=workspace.model_tabular_weights_path,
|
|
1049
|
-
device=device,
|
|
1050
|
-
)
|
|
1051
|
-
else:
|
|
1052
|
-
_LOG.warning("model weights not found; generating data with an untrained model")
|
|
1053
|
-
|
|
1054
|
-
model.to(device)
|
|
1055
|
-
model.eval()
|
|
1015
|
+
model = create_and_load_model(
|
|
1016
|
+
workspace=workspace,
|
|
1017
|
+
is_sequential=is_sequential,
|
|
1018
|
+
tgt_cardinalities=tgt_cardinalities,
|
|
1019
|
+
ctx_cardinalities=ctx_cardinalities,
|
|
1020
|
+
model_units=model_units,
|
|
1021
|
+
ctx_seq_len_median=ctx_seq_len_median,
|
|
1022
|
+
column_order=gen_column_order,
|
|
1023
|
+
device=device,
|
|
1024
|
+
seq_len_median=seq_len_median,
|
|
1025
|
+
seq_len_max=seq_len_max,
|
|
1026
|
+
)
|
|
1056
1027
|
|
|
1057
1028
|
# calculate fairness transforms only once before batch generation
|
|
1058
1029
|
fairness_transforms: FairnessTransforms | None = None
|
|
@@ -1121,13 +1092,11 @@ def generate(
|
|
|
1121
1092
|
ctx_batch = ctx_batch.sort_values(ctx_primary_key).reset_index(drop=True)
|
|
1122
1093
|
seed_batch = seed_batch.sort_values(tgt_context_key).reset_index(drop=True)
|
|
1123
1094
|
|
|
1124
|
-
# encode ctx_batch
|
|
1095
|
+
# encode ctx_batch and prepare tensor inputs
|
|
1125
1096
|
_LOG.info(f"encode context {ctx_batch.shape}")
|
|
1126
|
-
ctx_batch_encoded, ctx_primary_key_encoded
|
|
1127
|
-
|
|
1097
|
+
ctx_inputs, ctx_batch_encoded, ctx_primary_key_encoded = prepare_context_inputs(
|
|
1098
|
+
ctx_data=ctx_batch, ctx_stats=ctx_stats, device=model.device, ctx_primary_key=ctx_primary_key
|
|
1128
1099
|
)
|
|
1129
|
-
# pad left context sequences to ensure non-empty sequences
|
|
1130
|
-
ctx_batch_encoded = pad_ctx_sequences(ctx_batch_encoded)
|
|
1131
1100
|
ctx_keys = ctx_batch_encoded[ctx_primary_key_encoded]
|
|
1132
1101
|
ctx_keys.rename(tgt_context_key, inplace=True)
|
|
1133
1102
|
|
|
@@ -1149,30 +1118,12 @@ def generate(
|
|
|
1149
1118
|
syn = ctx_keys.to_frame().reset_index(drop=True)
|
|
1150
1119
|
buffer.add((syn, seed_batch))
|
|
1151
1120
|
elif isinstance(model, SequentialModel):
|
|
1152
|
-
|
|
1153
|
-
col: torch.unsqueeze(
|
|
1154
|
-
torch.as_tensor(ctx_batch_encoded[col].to_numpy(), device=model.device).type(torch.int),
|
|
1155
|
-
dim=-1,
|
|
1156
|
-
)
|
|
1157
|
-
for col in ctx_batch_encoded.columns
|
|
1158
|
-
if col.startswith(CTXFLT)
|
|
1159
|
-
}
|
|
1160
|
-
ctxseq_inputs = {
|
|
1161
|
-
col: torch.unsqueeze(
|
|
1162
|
-
torch.nested.as_nested_tensor(
|
|
1163
|
-
[torch.as_tensor(t, device=model.device).type(torch.int) for t in ctx_batch_encoded[col]],
|
|
1164
|
-
device=model.device,
|
|
1165
|
-
),
|
|
1166
|
-
dim=-1,
|
|
1167
|
-
)
|
|
1168
|
-
for col in ctx_batch_encoded.columns
|
|
1169
|
-
if col.startswith(CTXSEQ)
|
|
1170
|
-
}
|
|
1121
|
+
# Use context inputs prepared earlier
|
|
1171
1122
|
seq_steps = model.tgt_seq_len_max
|
|
1172
1123
|
history = None
|
|
1173
1124
|
history_state = None
|
|
1174
1125
|
# process context just once for all sequence steps
|
|
1175
|
-
context = model.context_compressor(
|
|
1126
|
+
context = model.context_compressor(ctx_inputs)
|
|
1176
1127
|
# loop over sequence steps, and pass forward history to keep model state-less
|
|
1177
1128
|
out_df: pd.DataFrame | None = None
|
|
1178
1129
|
decode_prev_steps = {}
|
|
@@ -1358,26 +1309,8 @@ def generate(
|
|
|
1358
1309
|
persist_data_part(syn, output_path, f"{buffer.n_clears:06}.{0:06}")
|
|
1359
1310
|
buffer.clear()
|
|
1360
1311
|
else: # isinstance(model, FlatModel)
|
|
1361
|
-
|
|
1362
|
-
|
|
1363
|
-
torch.as_tensor(ctx_batch_encoded[col].to_numpy(), device=model.device).type(torch.int),
|
|
1364
|
-
dim=-1,
|
|
1365
|
-
)
|
|
1366
|
-
for col in ctx_batch_encoded.columns
|
|
1367
|
-
if col.startswith(CTXFLT)
|
|
1368
|
-
}
|
|
1369
|
-
ctxseq_inputs = {
|
|
1370
|
-
col: torch.unsqueeze(
|
|
1371
|
-
torch.nested.as_nested_tensor(
|
|
1372
|
-
[torch.as_tensor(t, device=model.device).type(torch.int) for t in ctx_batch_encoded[col]],
|
|
1373
|
-
device=model.device,
|
|
1374
|
-
),
|
|
1375
|
-
dim=-1,
|
|
1376
|
-
)
|
|
1377
|
-
for col in ctx_batch_encoded.columns
|
|
1378
|
-
if col.startswith(CTXSEQ)
|
|
1379
|
-
}
|
|
1380
|
-
x = ctxflt_inputs | ctxseq_inputs
|
|
1312
|
+
# Use context inputs prepared earlier
|
|
1313
|
+
x = ctx_inputs
|
|
1381
1314
|
fixed_values = {
|
|
1382
1315
|
col: torch.as_tensor(seed_batch_encoded[col].to_numpy(), device=model.device).type(torch.int)
|
|
1383
1316
|
for col in seed_batch_encoded.columns
|