mostlyai-engine 2.1.0__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.
Files changed (55) hide show
  1. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/PKG-INFO +1 -1
  2. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/__init__.py +1 -1
  3. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/argn.py +99 -64
  4. mostlyai_engine-2.2.0/mostlyai/engine/_tabular/common.py +200 -0
  5. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/generation.py +28 -95
  6. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/interface.py +127 -152
  7. mostlyai_engine-2.2.0/mostlyai/engine/_tabular/probability.py +481 -0
  8. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/pyproject.toml +1 -1
  9. mostlyai_engine-2.1.0/mostlyai/engine/_tabular/common.py +0 -37
  10. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/.gitignore +0 -0
  11. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/LICENSE +0 -0
  12. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/README.md +0 -0
  13. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_common.py +0 -0
  14. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_dtypes.py +0 -0
  15. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  16. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  17. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  18. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  19. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  20. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  21. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  22. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  23. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  24. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  25. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  26. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  27. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  28. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/__init__.py +0 -0
  29. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/common.py +0 -0
  30. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/encoding.py +0 -0
  31. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
  32. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/base.py +0 -0
  33. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  34. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  35. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/generation.py +0 -0
  36. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/interface.py +0 -0
  37. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/lstm.py +0 -0
  38. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  39. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/training.py +0 -0
  40. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
  41. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_memory.py +0 -0
  42. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/__init__.py +0 -0
  43. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/encoding.py +0 -0
  44. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/fairness.py +0 -0
  45. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_tabular/training.py +0 -0
  46. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_training_utils.py +0 -0
  47. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/_workspace.py +0 -0
  48. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/analysis.py +0 -0
  49. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/domain.py +0 -0
  50. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/encoding.py +0 -0
  51. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/generation.py +0 -0
  52. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/logging.py +0 -0
  53. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/random_state.py +0 -0
  54. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/splitting.py +0 -0
  55. {mostlyai_engine-2.1.0 → mostlyai_engine-2.2.0}/mostlyai/engine/training.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mostlyai-engine
3
- Version: 2.1.0
3
+ Version: 2.2.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
@@ -34,7 +34,7 @@ __all__ = [
34
34
  "TabularARGN",
35
35
  "LanguageModel",
36
36
  ]
37
- __version__ = "2.1.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[sub_col] for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1076
+ tgt_embeds[sc] for sc in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1024
1077
  ]
1025
1078
 
1026
1079
  # regressor
@@ -1033,84 +1086,68 @@ class FlatModel(nn.Module):
1033
1086
  # update output
1034
1087
  outputs[sub_col] = xs
1035
1088
 
1036
- else: # mode == "gen"
1037
- # forward pass through context compressor
1038
- context = self.context_compressor(x)
1039
- context = self._handle_context(context)
1040
-
1041
- # initialize sub column embeddings
1042
- tgt_embeds = self.embedders.zero_mask(batch_size)
1043
-
1044
- # initialize column embeddings
1045
- tgt_col_embeds = self.column_embedders.zero_mask(batch_size)
1046
-
1047
- # concatenate column embeddings
1048
- col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1089
+ return outputs, {}
1049
1090
 
1050
- # take sub columns in the specified generation order
1051
- column_order = effective_column_order or self.tgt_columns
1052
- sub_column_order = [sub_col for col in column_order for sub_col in self.tgt_column_sub_columns[col]]
1091
+ elif mode == "gen":
1092
+ context, tgt_embeds, tgt_col_embeds, col_embeddings, sub_column_order = self._initialize_generation(
1093
+ x, batch_size, effective_column_order
1094
+ )
1053
1095
 
1054
1096
  for sub_col in sub_column_order:
1055
- lookup = self.tgt_sub_columns_lookup[sub_col]
1056
-
1057
- # if sub column is fixed, skip sampling and use that value
1097
+ # handle fixed values
1058
1098
  if sub_col in fixed_values:
1059
1099
  out = fixed_values[sub_col]
1100
+ else:
1101
+ # compute probabilities and sample
1102
+ logits = self._compute_logits(sub_col, context, col_embeddings, tgt_embeds)
1103
+ probs_tensor = nn.Softmax(dim=-1)(logits)
1060
1104
 
1061
- else: # sample from distribution
1062
- # collect previous sub column embeddings for current column
1063
- prev_sub_col_embeds = [
1064
- tgt_embeds[sub_col]
1065
- for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1066
- ]
1067
-
1068
- # regressor
1069
- regressor_in = context + [col_embeddings] + prev_sub_col_embeds
1070
- xs = self.regressors(regressor_in, sub_col)
1071
-
1072
- # predictor
1073
- xs = self.predictors(xs, sub_col)
1074
-
1075
- # softmax to probs
1076
- xs = nn.Softmax(dim=-1)(xs)
1077
-
1078
- # keep probabilities (used e.g. for fairness)
1105
+ # optionally keep probabilities
1079
1106
  if sub_col in return_probs:
1080
- probs[sub_col] = xs
1107
+ probs[sub_col] = probs_tensor
1081
1108
 
1082
- # apply fairness transform when generating the target sub column
1109
+ # apply fairness transforms
1083
1110
  if fairness_transforms:
1084
- xs = apply_fairness_transforms(sub_col, xs, outputs, fairness_transforms)
1111
+ probs_tensor = apply_fairness_transforms(sub_col, probs_tensor, outputs, fairness_transforms)
1085
1112
 
1086
1113
  # sample
1087
1114
  out = torch.squeeze(
1088
- _sample(
1089
- probs=xs,
1090
- temperature=temperature,
1091
- top_p=top_p,
1092
- fixed_probs=fixed_probs.get(sub_col),
1093
- ),
1115
+ _sample(probs_tensor, temperature, top_p, fixed_probs.get(sub_col)),
1094
1116
  dim=-1,
1095
1117
  )
1096
1118
 
1097
- # update output
1119
+ # update output and embeddings
1098
1120
  outputs[sub_col] = out
1121
+ col_embeddings = self._update_embeddings(sub_col, out, tgt_embeds, tgt_col_embeds, col_embeddings)
1099
1122
 
1100
- # update current sub column embedding
1101
- tgt_embeds[sub_col] = self.embedders.get(sub_col)(out)
1123
+ # order outputs and return
1124
+ outputs = {sub_col: outputs[sub_col] for sub_col in self.tgt_sub_columns}
1125
+ return outputs, probs
1102
1126
 
1103
- # update current column embedding
1104
- if sub_col in self.last_sub_cols:
1105
- col_sub_cols = self.tgt_column_sub_columns[lookup.col_name]
1106
- col_embed_in = torch.cat([tgt_embeds[sc] for sc in col_sub_cols], dim=-1)
1107
- tgt_col_embeds[lookup.col_name] = self.column_embedders.get(lookup.col_name)(col_embed_in)
1108
- col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1127
+ elif mode == "probs":
1128
+ context, tgt_embeds, tgt_col_embeds, col_embeddings, sub_column_order = self._initialize_generation(
1129
+ x, batch_size, effective_column_order
1130
+ )
1109
1131
 
1110
- # order outputs according to tgt_sub_columns
1111
- outputs = {sub_col: outputs[sub_col] for sub_col in self.tgt_sub_columns}
1132
+ for sub_col in sub_column_order:
1133
+ # handle fixed values
1134
+ if sub_col in fixed_values:
1135
+ out = fixed_values[sub_col]
1136
+ # update embeddings to maintain correct autoregressive context
1137
+ col_embeddings = self._update_embeddings(sub_col, out, tgt_embeds, tgt_col_embeds, col_embeddings)
1138
+ else:
1139
+ # compute probabilities without sampling
1140
+ logits = self._compute_logits(sub_col, context, col_embeddings, tgt_embeds)
1141
+ probs_tensor = nn.Softmax(dim=-1)(logits)
1142
+
1143
+ # apply fixed_probs mask if provided
1144
+ if sub_col in fixed_probs:
1145
+ probs_tensor = _sampling_fixed_probs(probs_tensor, fixed_probs[sub_col])
1146
+
1147
+ # store probabilities (no sampling, no embedding updates)
1148
+ probs[sub_col] = probs_tensor
1112
1149
 
1113
- return outputs, probs
1150
+ return {}, probs
1114
1151
 
1115
1152
 
1116
1153
  class AttentionModule(nn.Module):
@@ -1334,8 +1371,7 @@ class SequentialModel(nn.Module):
1334
1371
 
1335
1372
  # collect previous sub column embeddings for current column
1336
1373
  prev_sub_col_embeds = {
1337
- sub_col: tgt_embeds[sub_col]
1338
- for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1374
+ sc: tgt_embeds[sc] for sc in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1339
1375
  }
1340
1376
  if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1341
1377
  # RIDX sub-columns should not see SLEN sub-columns
@@ -1405,8 +1441,7 @@ class SequentialModel(nn.Module):
1405
1441
  else: # sample from distribution
1406
1442
  # collect previous sub column embeddings for current column
1407
1443
  prev_sub_col_embeds = {
1408
- sub_col: tgt_embeds[sub_col]
1409
- for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1444
+ sc: tgt_embeds[sc] for sc in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1410
1445
  }
1411
1446
  if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1412
1447
  # RIDX sub-columns should not see SLEN sub-columns
@@ -0,0 +1,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 load_model_weights
84
- from mostlyai.engine._tabular.encoding import encode_df, pad_ctx_sequences
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.model_configs.read()
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
- _LOG.info("create generative model")
1020
- model: FlatModel | SequentialModel
1021
- if is_sequential:
1022
- model = SequentialModel(
1023
- tgt_cardinalities=tgt_cardinalities,
1024
- tgt_seq_len_median=seq_len_median,
1025
- tgt_seq_len_max=seq_len_max,
1026
- ctx_cardinalities=ctx_cardinalities,
1027
- ctxseq_len_median=ctx_seq_len_median,
1028
- model_size=model_units,
1029
- column_order=gen_column_order,
1030
- device=device,
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, _ = encode_df(
1127
- df=ctx_batch, stats=ctx_stats, ctx_primary_key=ctx_primary_key
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
- ctxflt_inputs = {
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(ctxflt_inputs | ctxseq_inputs)
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
- ctxflt_inputs = {
1362
- col: torch.unsqueeze(
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