mostlyai-engine 1.5.4__tar.gz → 1.5.6__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 (52) hide show
  1. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/PKG-INFO +1 -1
  2. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/__init__.py +1 -1
  3. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_common.py +10 -5
  4. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/argn.py +36 -18
  5. mostlyai_engine-1.5.6/mostlyai/engine/_tabular/common.py +37 -0
  6. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/generation.py +21 -15
  7. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/training.py +1 -0
  8. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/pyproject.toml +1 -1
  9. mostlyai_engine-1.5.4/mostlyai/engine/_tabular/common.py +0 -30
  10. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/.gitignore +0 -0
  11. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/LICENSE +0 -0
  12. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/README.md +0 -0
  13. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_dtypes.py +0 -0
  14. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  15. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  16. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  17. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  18. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  19. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  20. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  21. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  22. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  23. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  24. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  25. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  26. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  27. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/__init__.py +0 -0
  28. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/common.py +0 -0
  29. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/encoding.py +0 -0
  30. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/__init__.py +0 -0
  31. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/base.py +0 -0
  32. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  33. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  34. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/generation.py +0 -0
  35. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/lstm.py +0 -0
  36. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  37. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/training.py +0 -0
  38. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
  39. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_memory.py +0 -0
  40. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/__init__.py +0 -0
  41. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/encoding.py +0 -0
  42. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_tabular/fairness.py +0 -0
  43. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_training_utils.py +0 -0
  44. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/_workspace.py +0 -0
  45. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/analysis.py +0 -0
  46. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/domain.py +0 -0
  47. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/encoding.py +0 -0
  48. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/generation.py +0 -0
  49. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/logging.py +0 -0
  50. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/random_state.py +0 -0
  51. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/splitting.py +0 -0
  52. {mostlyai_engine-1.5.4 → mostlyai_engine-1.5.6}/mostlyai/engine/training.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mostlyai-engine
3
- Version: 1.5.4
3
+ Version: 1.5.6
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
@@ -22,7 +22,7 @@ from mostlyai.engine.splitting import split
22
22
  from mostlyai.engine.training import train
23
23
 
24
24
  __all__ = ["split", "analyze", "encode", "train", "generate", "init_logging", "set_random_state"]
25
- __version__ = "1.5.4"
25
+ __version__ = "1.5.6"
26
26
 
27
27
  # suppress specific warning related to os.fork() in multi-threaded processes
28
28
  warnings.filterwarnings("ignore", category=DeprecationWarning, message=".*multi-threaded.*fork.*")
@@ -56,6 +56,11 @@ SLEN_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}slen_" # seque
56
56
  SDEC_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}sdec_" # sequence index decile
57
57
  TABLE_COLUMN_INFIX = "::" # this should be consistent as in mostly-data and mostlyai-qa
58
58
 
59
+ # the latest version of the model uses SIDX/SLEN/RIDX positional column
60
+ DEFAULT_HAS_SLEN = True
61
+ DEFAULT_HAS_RIDX = True
62
+ DEFAULT_HAS_SDEC = False
63
+
59
64
  ANALYZE_MIN_MAX_TOP_N = 1000 # the number of min/max values to be kept from each partition
60
65
 
61
66
  # the minimal number of min/max values to trigger the reduction; if less, the min/max will be reduced to None
@@ -318,6 +323,7 @@ def get_argn_name(
318
323
  def get_cardinalities(
319
324
  stats: dict, has_slen: bool | None = None, has_ridx: bool | None = None, has_sdec: bool | None = None
320
325
  ) -> dict[str, int]:
326
+ # the latest version of the model uses SIDX/SLEN/RIDX positional column (applies to sequential model only)
321
327
  cardinalities: dict[str, int] = {}
322
328
 
323
329
  if stats.get("is_sequential", False):
@@ -538,12 +544,11 @@ def decode_positional_column(df_encoded: pd.DataFrame, max_seq_len: int, prefix:
538
544
 
539
545
 
540
546
  def get_positional_cardinalities(
541
- max_seq_len: int, has_slen: bool | None, has_ridx: bool | None, has_sdec: bool | None
547
+ max_seq_len: int, has_slen: bool | None = None, has_ridx: bool | None = None, has_sdec: bool | None = None
542
548
  ) -> dict[str, int]:
543
- # the latest version of the model uses SIDX/SLEN/RIDX positional column
544
- has_slen = has_slen if has_slen is not None else True
545
- has_ridx = has_ridx if has_ridx is not None else True
546
- has_sdec = has_sdec if has_sdec is not None else False
549
+ has_slen = has_slen if has_slen is not None else DEFAULT_HAS_SLEN
550
+ has_ridx = has_ridx if has_ridx is not None else DEFAULT_HAS_RIDX
551
+ has_sdec = has_sdec if has_sdec is not None else DEFAULT_HAS_SDEC
547
552
 
548
553
  if max_seq_len < SIDX_RIDX_DIGIT_ENCODING_THRESHOLD:
549
554
  # encode positional columns as numeric_discrete
@@ -226,15 +226,12 @@ class Embedders(nn.Module):
226
226
  ),
227
227
  None,
228
228
  )
229
- last_ridx_sub_col = next(
230
- (sub_col for sub_col in reversed(self.cardinalities) if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX)), None
231
- )
232
229
  for sub_col, dim_input in self.cardinalities.items():
233
230
  dim_output = _embedding_heuristic(id=self.id(sub_col), model_size=model_size, dim_input=dim_input)
234
231
  embedder = nn.Embedding(num_embeddings=dim_input, embedding_dim=dim_output, device=device)
235
- # the embeddings of the last slen and ridx sub columns are never used
232
+ # the embeddings of the last SLEN sub column are never used
236
233
  # so we explicitly freeze them to make opacus not complain about "per sample gradient is not initialized"
237
- if sub_col in [last_ridx_sub_col, last_slen_sub_col]:
234
+ if sub_col == last_slen_sub_col:
238
235
  embedder.weight.requires_grad = False
239
236
  self.add(sub_column=sub_col, embedder=embedder)
240
237
  self.dims.append(dim_output)
@@ -874,6 +871,7 @@ class FlatModel(nn.Module):
874
871
  model_size: ModelSizeOrUnits,
875
872
  column_order: list[str] | None,
876
873
  device: torch.device,
874
+ with_dp: bool = False,
877
875
  ):
878
876
  super().__init__()
879
877
 
@@ -895,6 +893,7 @@ class FlatModel(nn.Module):
895
893
  ctx_cardinalities=self.ctx_cardinalities,
896
894
  ctxseq_len_median=self.ctxseq_len_median,
897
895
  device=device,
896
+ with_dp=with_dp,
898
897
  )
899
898
 
900
899
  # sub column embeddings
@@ -1256,26 +1255,28 @@ class SequentialModel(nn.Module):
1256
1255
  if context is None:
1257
1256
  context = self.context_compressor(x)
1258
1257
 
1259
- # SLEN is only masked for models with RIDX
1260
1258
  has_ridx = any(sub_col.startswith(RIDX_SUB_COLUMN_PREFIX) for sub_col in self.tgt_cardinalities)
1261
- masked_positional_columns = (
1259
+
1260
+ # SLEN and RIDX are masked for history
1261
+ # NOTE: SLEN is not masked for models without RIDX (backwards compatibility)
1262
+ history_masked_sub_cols = (
1262
1263
  (SLEN_SUB_COLUMN_PREFIX, RIDX_SUB_COLUMN_PREFIX) if has_ridx else (RIDX_SUB_COLUMN_PREFIX,)
1263
1264
  )
1265
+ # SLEN is masked for column embeddings
1266
+ # NOTE: SLEN is not masked for models without RIDX (backwards compatibility)
1267
+ col_embeddings_masked_sub_cols = (SLEN_SUB_COLUMN_PREFIX,) if has_ridx else ()
1264
1268
 
1265
1269
  outputs = {}
1266
1270
  if mode == "trn":
1267
1271
  # forward pass through sub column embedders
1268
1272
  tgt_embeds = self.embedders(x)
1269
- tgt_embeds_positional_masked = {
1270
- k: torch.zeros_like(v) if k.startswith(masked_positional_columns) else v for k, v in tgt_embeds.items()
1271
- }
1272
-
1273
- # forward pass through column embedders
1274
- tgt_col_embeds = self.column_embedders(tgt_embeds_positional_masked)
1275
1273
 
1276
1274
  # history
1277
1275
  # time shift: remove last time step; add zeros for first time step; add randoms for all others
1278
- embeddings = torch.cat(list(tgt_embeds_positional_masked.values()), dim=-1)
1276
+ tgt_embeds_for_history = {
1277
+ k: torch.zeros_like(v) if k.startswith(history_masked_sub_cols) else v for k, v in tgt_embeds.items()
1278
+ }
1279
+ embeddings = torch.cat(list(tgt_embeds_for_history.values()), dim=-1)
1279
1280
  history_in = embeddings[:, :-1, :]
1280
1281
  history_in = nn.ConstantPad2d((0, 0, 1, 0), 0)(history_in)
1281
1282
  history, _ = self.history_compressor(history_in)
@@ -1287,6 +1288,13 @@ class SequentialModel(nn.Module):
1287
1288
  flat_ctx = self._repeat_flat_context(flat_ctx, history.size(1))
1288
1289
  context_history = [torch.cat(flat_ctx + seq_ctx + [history], -1)]
1289
1290
 
1291
+ # forward pass through column embedders
1292
+ tgt_embeds_for_col_embeddings = {
1293
+ k: torch.zeros_like(v) if k.startswith(col_embeddings_masked_sub_cols) else v
1294
+ for k, v in tgt_embeds.items()
1295
+ }
1296
+ tgt_col_embeds = self.column_embedders(tgt_embeds_for_col_embeddings)
1297
+
1290
1298
  # create batch-wise permutation mask
1291
1299
  col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1292
1300
  col_mask = _make_permutation_mask(
@@ -1327,7 +1335,9 @@ class SequentialModel(nn.Module):
1327
1335
  return outputs, {}
1328
1336
 
1329
1337
  else: # mode == "gen"
1338
+ is_0th_step = False
1330
1339
  if history is None or history_state is None:
1340
+ is_0th_step = True
1331
1341
  # initialize history
1332
1342
  history_in = torch.cat(
1333
1343
  [
@@ -1401,24 +1411,32 @@ class SequentialModel(nn.Module):
1401
1411
  )
1402
1412
 
1403
1413
  # update timestep output
1414
+ if is_0th_step and sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1415
+ # overwrite output for RIDX sub-columns on 0th step with SLEN sub-columns
1416
+ slen_sub_col = sub_col.replace(RIDX_SUB_COLUMN_PREFIX, SLEN_SUB_COLUMN_PREFIX)
1417
+ out = outputs[slen_sub_col] if slen_sub_col in outputs else out
1404
1418
  outputs[sub_col] = out
1405
1419
 
1406
1420
  # update current sub column embedding
1407
1421
  tgt_embeds[sub_col] = self.embedders.get(sub_col)(out)
1408
- tgt_embeds_positional_masked = {
1409
- k: torch.zeros_like(v) if k.startswith(masked_positional_columns) else v
1422
+ tgt_embeds_for_history = {
1423
+ k: torch.zeros_like(v) if k.startswith(history_masked_sub_cols) else v
1424
+ for k, v in tgt_embeds.items()
1425
+ }
1426
+ tgt_embeds_for_col_embeddings = {
1427
+ k: torch.zeros_like(v) if k.startswith(col_embeddings_masked_sub_cols) else v
1410
1428
  for k, v in tgt_embeds.items()
1411
1429
  }
1412
1430
 
1413
1431
  # update current column embedding
1414
1432
  if sub_col in self.tgt_last_sub_cols:
1415
1433
  col_sub_cols = self.tgt_column_sub_columns[lookup.col_name]
1416
- col_embed_in = torch.cat([tgt_embeds_positional_masked[sc] for sc in col_sub_cols], dim=-1)
1434
+ col_embed_in = torch.cat([tgt_embeds_for_col_embeddings[sc] for sc in col_sub_cols], dim=-1)
1417
1435
  tgt_col_embeds[lookup.col_name] = self.column_embedders.get(lookup.col_name)(col_embed_in)
1418
1436
  col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1419
1437
 
1420
1438
  # update history and hidden state
1421
- history_in = torch.cat([v for v in tgt_embeds_positional_masked.values()], dim=-1)
1439
+ history_in = torch.cat(list(tgt_embeds_for_history.values()), dim=-1)
1422
1440
  history, history_state = self.history_compressor(history_in, history_state=history_state)
1423
1441
 
1424
1442
  # order outputs according to tgt_sub_columns
@@ -0,0 +1,37 @@
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 torch
20
+
21
+ _LOG = logging.getLogger(__name__)
22
+
23
+ DPLSTM_SUFFIXES: tuple = ("ih.weight", "ih.bias", "hh.weight", "hh.bias")
24
+
25
+
26
+ def load_model_weights(model: torch.nn.Module, path: Path, device: torch.device) -> None:
27
+ t0 = time.time()
28
+ incompatible_keys = model.load_state_dict(torch.load(f=path, map_location=device, weights_only=True), strict=False)
29
+ missing_keys = incompatible_keys.missing_keys
30
+ unexpected_keys = incompatible_keys.unexpected_keys
31
+ # 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)
32
+ # but if there're any other missing or unexpected keys, an error should be raised
33
+ if len(missing_keys) > 0 or any(not k.endswith(DPLSTM_SUFFIXES) for k in unexpected_keys):
34
+ raise RuntimeError(
35
+ f"failed to load model weights due to incompatibility: {missing_keys = }, {unexpected_keys = }"
36
+ )
37
+ _LOG.info(f"loaded model weights in {time.time() - t0:.2f}s")
@@ -30,6 +30,9 @@ from mostlyai.engine._common import (
30
30
  ARGN_TABLE,
31
31
  CTXFLT,
32
32
  CTXSEQ,
33
+ DEFAULT_HAS_RIDX,
34
+ DEFAULT_HAS_SDEC,
35
+ DEFAULT_HAS_SLEN,
33
36
  POSITIONAL_COLUMN,
34
37
  RIDX_SUB_COLUMN_PREFIX,
35
38
  SDEC_SUB_COLUMN_PREFIX,
@@ -695,6 +698,8 @@ def generate(
695
698
  has_slen = any(SLEN_SUB_COLUMN_PREFIX in k for k in model_units.keys())
696
699
  has_ridx = any(RIDX_SUB_COLUMN_PREFIX in k for k in model_units.keys())
697
700
  has_sdec = any(SDEC_SUB_COLUMN_PREFIX in k for k in model_units.keys())
701
+ else:
702
+ has_slen, has_ridx, has_sdec = DEFAULT_HAS_SLEN, DEFAULT_HAS_RIDX, DEFAULT_HAS_SDEC
698
703
 
699
704
  tgt_cardinalities = get_cardinalities(tgt_stats, has_slen, has_ridx, has_sdec)
700
705
  ctx_cardinalities = get_cardinalities(ctx_stats)
@@ -874,11 +879,14 @@ def generate(
874
879
  no_of_model_params = get_no_of_model_parameters(model)
875
880
  _LOG.info(f"{no_of_model_params=}")
876
881
 
877
- load_model_weights(
878
- model=model,
879
- path=workspace.model_tabular_weights_path,
880
- device=device,
881
- )
882
+ if workspace.model_tabular_weights_path.exists():
883
+ load_model_weights(
884
+ model=model,
885
+ path=workspace.model_tabular_weights_path,
886
+ device=device,
887
+ )
888
+ else:
889
+ _LOG.warning("model weights not found; generating data with an untrained model")
882
890
 
883
891
  model.to(device)
884
892
  model.eval()
@@ -1023,7 +1031,7 @@ def generate(
1023
1031
 
1024
1032
  # fix SLEN by propagating sampled SLEN from first step
1025
1033
  slen_vals = {}
1026
- if seq_step > 0:
1034
+ if has_slen and seq_step > 0:
1027
1035
  slen = out_df[SLEN_SUB_COLUMN_PREFIX]
1028
1036
  slen = encode_positional_column(slen, max_seq_len=seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX)
1029
1037
  slen_vals = {
@@ -1102,16 +1110,14 @@ def generate(
1102
1110
  out_df[SIDX_SUB_COLUMN_PREFIX] = decode_positional_column(
1103
1111
  out_df, seq_len_max, prefix=SIDX_SUB_COLUMN_PREFIX
1104
1112
  )
1105
- out_df[SLEN_SUB_COLUMN_PREFIX] = decode_positional_column(
1106
- out_df, seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX
1107
- )
1108
- out_df[SLEN_SUB_COLUMN_PREFIX] = out_df[SLEN_SUB_COLUMN_PREFIX].clip(lower=seq_len_min)
1113
+ if has_slen:
1114
+ out_df[SLEN_SUB_COLUMN_PREFIX] = decode_positional_column(
1115
+ out_df, seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX
1116
+ )
1117
+ out_df[SLEN_SUB_COLUMN_PREFIX] = out_df[SLEN_SUB_COLUMN_PREFIX].clip(lower=seq_len_min)
1109
1118
  if has_ridx:
1110
- # set RIDX to SLEN for first step; beyond that decode RIDX
1111
- out_df[RIDX_SUB_COLUMN_PREFIX] = (
1112
- decode_positional_column(out_df, seq_len_max, prefix=RIDX_SUB_COLUMN_PREFIX)
1113
- if seq_step > 0
1114
- else out_df[SLEN_SUB_COLUMN_PREFIX]
1119
+ out_df[RIDX_SUB_COLUMN_PREFIX] = decode_positional_column(
1120
+ out_df, seq_len_max, prefix=RIDX_SUB_COLUMN_PREFIX
1115
1121
  )
1116
1122
  out_df[RIDX_SUB_COLUMN_PREFIX] = out_df[RIDX_SUB_COLUMN_PREFIX].clip(
1117
1123
  lower=seq_len_min - seq_step, upper=seq_len_max
@@ -462,6 +462,7 @@ def train(
462
462
  model_size=model_size,
463
463
  column_order=trn_column_order,
464
464
  device=device,
465
+ with_dp=with_dp,
465
466
  )
466
467
  _LOG.info(f"model class: {argn.__class__.__name__}")
467
468
 
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "mostlyai-engine"
3
- version = "1.5.4"
3
+ version = "1.5.6"
4
4
  description = "Synthetic Data Engine"
5
5
  authors = [{ name = "MOSTLY AI", email = "dev@mostly.ai" }]
6
6
  requires-python = ">=3.10"
@@ -1,30 +0,0 @@
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 torch
20
-
21
- _LOG = logging.getLogger(__name__)
22
-
23
-
24
- def load_model_weights(model: torch.nn.Module, path: Path, device: torch.device):
25
- try:
26
- t00 = time.time()
27
- model.load_state_dict(torch.load(f=path, map_location=device, weights_only=True))
28
- _LOG.info(f"loaded model weights in {time.time() - t00:.2f}s")
29
- except Exception as e:
30
- _LOG.warning(f"failed to load model weights: {e}")
File without changes