mostlyai-engine 1.4.8__tar.gz → 1.5.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 (51) hide show
  1. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/PKG-INFO +3 -3
  2. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/__init__.py +1 -1
  3. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_common.py +48 -46
  4. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/common.py +5 -2
  5. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/generation.py +10 -11
  6. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/training.py +4 -1
  7. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/xgrammar_utils.py +2 -2
  8. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/argn.py +46 -11
  9. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/encoding.py +67 -35
  10. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/generation.py +241 -166
  11. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/training.py +28 -23
  12. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/pyproject.toml +3 -3
  13. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/.gitignore +0 -0
  14. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/LICENSE +0 -0
  15. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/README.md +0 -0
  16. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_dtypes.py +0 -0
  17. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  18. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  19. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  20. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  21. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  22. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  23. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  24. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  25. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  26. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  27. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  28. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  29. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  30. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/__init__.py +0 -0
  31. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/encoding.py +0 -0
  32. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
  33. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/base.py +0 -0
  34. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  35. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  36. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/lstm.py +0 -0
  37. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  38. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_memory.py +0 -0
  39. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/__init__.py +0 -0
  40. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/common.py +0 -0
  41. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/fairness.py +0 -0
  42. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_training_utils.py +0 -0
  43. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_workspace.py +0 -0
  44. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/analysis.py +0 -0
  45. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/domain.py +0 -0
  46. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/encoding.py +0 -0
  47. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/generation.py +0 -0
  48. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/logging.py +0 -0
  49. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/random_state.py +0 -0
  50. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/splitting.py +0 -0
  51. {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/training.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mostlyai-engine
3
- Version: 1.4.8
3
+ Version: 1.5.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
@@ -28,7 +28,7 @@ Requires-Dist: accelerate>=1.5.0
28
28
  Requires-Dist: datasets>=3.0.0
29
29
  Requires-Dist: huggingface-hub[hf-xet]>=0.30.2
30
30
  Requires-Dist: joblib>=1.4.2
31
- Requires-Dist: json-repair<0.47.0,>=0.30.0
31
+ Requires-Dist: json-repair>=0.47.0
32
32
  Requires-Dist: numpy>=2.0.0
33
33
  Requires-Dist: opacus>=1.5.4
34
34
  Requires-Dist: pandas~=2.2.0
@@ -40,7 +40,7 @@ Requires-Dist: tokenizers>=0.21.0
40
40
  Requires-Dist: torch<2.7.1,>=2.7.0
41
41
  Requires-Dist: torchaudio<2.7.1,>=2.7.0
42
42
  Requires-Dist: torchvision<0.22.1,>=0.22.0
43
- Requires-Dist: transformers>=4.51.0
43
+ Requires-Dist: transformers<4.54.0,>=4.51.0
44
44
  Requires-Dist: xgrammar>=0.1.19
45
45
  Provides-Extra: gpu
46
46
  Requires-Dist: bitsandbytes==0.42.0; (sys_platform == 'darwin') and extra == 'gpu'
@@ -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.4.8"
25
+ __version__ = "1.5.0"
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.*")
@@ -48,11 +48,12 @@ ARGN_COLUMN = "argn_column"
48
48
  PREFIX_TABLE = ":"
49
49
  PREFIX_COLUMN = "/"
50
50
  PREFIX_SUB_COLUMN = "__"
51
- SLEN_SIDX_SDEC_COLUMN = f"{TGT}{PREFIX_TABLE}{PREFIX_COLUMN}"
52
- SLEN_SIDX_DIGIT_ENCODING_THRESHOLD = 100
53
- SLEN_SUB_COLUMN_PREFIX = f"{SLEN_SIDX_SDEC_COLUMN}{PREFIX_SUB_COLUMN}slen_" # sequence length
54
- SIDX_SUB_COLUMN_PREFIX = f"{SLEN_SIDX_SDEC_COLUMN}{PREFIX_SUB_COLUMN}sidx_" # sequence index
55
- SDEC_SUB_COLUMN_PREFIX = f"{SLEN_SIDX_SDEC_COLUMN}{PREFIX_SUB_COLUMN}sdec_" # sequence index decile
51
+ SIDX_RIDX_DIGIT_ENCODING_THRESHOLD = 100
52
+ POSITIONAL_COLUMN = f"{TGT}{PREFIX_TABLE}{PREFIX_COLUMN}"
53
+ SIDX_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}sidx_" # sequence index
54
+ RIDX_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}ridx_" # reverse index
55
+ SLEN_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}slen_" # sequence length
56
+ SDEC_SUB_COLUMN_PREFIX = f"{POSITIONAL_COLUMN}{PREFIX_SUB_COLUMN}sdec_" # sequence index decile
56
57
  TABLE_COLUMN_INFIX = "::" # this should be consistent as in mostly-data and mostlyai-qa
57
58
 
58
59
  ANALYZE_MIN_MAX_TOP_N = 1000 # the number of min/max values to be kept from each partition
@@ -314,11 +315,14 @@ def get_argn_name(
314
315
  return "".join(name)
315
316
 
316
317
 
317
- def get_cardinalities(stats: dict) -> dict[str, int]:
318
+ def get_cardinalities(
319
+ stats: dict, has_slen: bool | None = None, has_ridx: bool | None = None, has_sdec: bool | None = None
320
+ ) -> dict[str, int]:
318
321
  cardinalities: dict[str, int] = {}
322
+
319
323
  if stats.get("is_sequential", False):
320
324
  max_seq_len = get_sequence_length_stats(stats)["max"]
321
- cardinalities |= get_slen_sidx_sdec_cardinalities(max_seq_len)
325
+ cardinalities |= get_positional_cardinalities(max_seq_len, has_slen, has_ridx, has_sdec)
322
326
 
323
327
  for i, column in enumerate(stats.get("columns", [])):
324
328
  column_stats = stats["columns"][column]
@@ -512,72 +516,70 @@ def skip_if_error(func: Callable) -> Callable:
512
516
  return skip_if_error_wrapper
513
517
 
514
518
 
515
- def encode_slen_sidx_sdec(vals: pd.Series, max_seq_len: int, prefix: str = "") -> pd.DataFrame:
519
+ def encode_positional_column(vals: pd.Series, max_seq_len: int, prefix: str = "") -> pd.DataFrame:
516
520
  assert is_integer_dtype(vals)
517
- if max_seq_len < SLEN_SIDX_DIGIT_ENCODING_THRESHOLD or prefix == SDEC_SUB_COLUMN_PREFIX:
518
- # encode slen and sidx as numeric_discrete
521
+ if max_seq_len < SIDX_RIDX_DIGIT_ENCODING_THRESHOLD:
522
+ # encode positional column as numeric_discrete
519
523
  df = pd.DataFrame({f"{prefix}cat": vals})
520
524
  else:
521
- # encode as numeric_digit
525
+ # encode positional column as numeric_digit
522
526
  n_digits = len(str(max_seq_len))
523
527
  df = pd.DataFrame(vals.astype(str).str.pad(width=n_digits, fillchar="0").apply(list).tolist()).astype(int)
524
528
  df.columns = [f"{prefix}E{i}" for i in range(n_digits - 1, -1, -1)]
525
529
  return df
526
530
 
527
531
 
528
- def decode_slen_sidx_sdec(df_encoded: pd.DataFrame, max_seq_len: int, prefix: str = "") -> pd.Series:
529
- if max_seq_len < SLEN_SIDX_DIGIT_ENCODING_THRESHOLD or prefix == SDEC_SUB_COLUMN_PREFIX:
530
- # decode slen and sidx as numeric_discrete
532
+ def decode_positional_column(df_encoded: pd.DataFrame, max_seq_len: int, prefix: str = "") -> pd.Series:
533
+ if max_seq_len < SIDX_RIDX_DIGIT_ENCODING_THRESHOLD:
534
+ # decode positional column as numeric_discrete
531
535
  vals = df_encoded[f"{prefix}cat"]
532
536
  else:
533
- # decode slen and sidx as numeric_digit
537
+ # decode positional column as numeric_digit
534
538
  n_digits = len(str(max_seq_len))
535
539
  vals = sum([df_encoded[f"{prefix}E{d}"] * 10 ** int(d) for d in list(range(n_digits))])
536
540
  return vals
537
541
 
538
542
 
539
- def get_slen_sidx_sdec_cardinalities(max_seq_len) -> dict[str, int]:
540
- if max_seq_len < SLEN_SIDX_DIGIT_ENCODING_THRESHOLD:
541
- # encode slen and sidx as numeric_discrete
542
- slen_cardinalities = {f"{SLEN_SUB_COLUMN_PREFIX}cat": max_seq_len + 1}
543
+ def get_positional_cardinalities(
544
+ max_seq_len: int, has_slen: bool | None, has_ridx: bool | None, has_sdec: bool | None
545
+ ) -> dict[str, int]:
546
+ # the latest version of the model uses SIDX/SLEN/RIDX positional column
547
+ has_slen = has_slen if has_slen is not None else True
548
+ has_ridx = has_ridx if has_ridx is not None else True
549
+ has_sdec = has_sdec if has_sdec is not None else False
550
+
551
+ if max_seq_len < SIDX_RIDX_DIGIT_ENCODING_THRESHOLD:
552
+ # encode positional columns as numeric_discrete
543
553
  sidx_cardinalities = {f"{SIDX_SUB_COLUMN_PREFIX}cat": max_seq_len + 1}
554
+ slen_cardinalities = {f"{SLEN_SUB_COLUMN_PREFIX}cat": max_seq_len + 1}
555
+ ridx_cardinalities = {f"{RIDX_SUB_COLUMN_PREFIX}cat": max_seq_len + 1}
544
556
  else:
545
- # encode slen and sidx as numeric_digit
557
+ # encode positional columns as numeric_digit
546
558
  digits = [int(digit) for digit in str(max_seq_len)]
547
- slen_cardinalities = {}
548
559
  sidx_cardinalities = {}
560
+ slen_cardinalities = {}
561
+ ridx_cardinalities = {}
549
562
  for idx, digit in enumerate(digits):
550
563
  # cap cardinality of the most significant position
551
564
  # less significant positions allow any digit
552
565
  card = digit + 1 if idx == 0 else 10
553
566
  e_idx = len(digits) - idx - 1
554
- slen_cardinalities[f"{SLEN_SUB_COLUMN_PREFIX}E{e_idx}"] = card
555
567
  sidx_cardinalities[f"{SIDX_SUB_COLUMN_PREFIX}E{e_idx}"] = card
556
- # order is important: slen first, then sidx, as the former has highest priority
568
+ ridx_cardinalities[f"{RIDX_SUB_COLUMN_PREFIX}E{e_idx}"] = card
569
+ slen_cardinalities[f"{SLEN_SUB_COLUMN_PREFIX}E{e_idx}"] = card
557
570
  sdec_cardinalities = {f"{SDEC_SUB_COLUMN_PREFIX}cat": 10}
558
- return slen_cardinalities | sidx_cardinalities | sdec_cardinalities
559
-
560
-
561
- def trim_sequences(syn: pd.DataFrame, tgt_context_key: str, seq_len_min: int, seq_len_max: int):
562
- if syn.empty:
563
- return syn
564
-
565
- # use SIDX and SLEN to determine sequence length
566
- syn[SIDX_SUB_COLUMN_PREFIX] = decode_slen_sidx_sdec(syn, seq_len_max, prefix=SIDX_SUB_COLUMN_PREFIX)
567
- syn[SLEN_SUB_COLUMN_PREFIX] = decode_slen_sidx_sdec(syn, seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX)
568
- # ensure that seq_len_min is respected
569
- syn[SLEN_SUB_COLUMN_PREFIX] = np.maximum(seq_len_min, syn[SLEN_SUB_COLUMN_PREFIX])
570
- syn = syn[syn[SIDX_SUB_COLUMN_PREFIX] < syn[SLEN_SUB_COLUMN_PREFIX]].reset_index(drop=True)
571
- # discarded padded context rows, ie where context key has been set to None
572
- syn = syn.dropna(subset=[tgt_context_key])
573
- # discard SLEN and SIDX columns
574
- syn.drop(
575
- [c for c in syn.columns if c.startswith(SLEN_SIDX_SDEC_COLUMN)],
576
- axis=1,
577
- inplace=True,
578
- )
579
- syn.reset_index(drop=True, inplace=True)
580
- return syn
571
+ match has_slen, has_ridx, has_sdec:
572
+ case True, True, False:
573
+ # SIDX/SLEN/RIDX model
574
+ return sidx_cardinalities | slen_cardinalities | ridx_cardinalities
575
+ case True, False, True:
576
+ # SLEN/SIDX/SDEC model
577
+ return slen_cardinalities | sidx_cardinalities | sdec_cardinalities
578
+ case True, False, False:
579
+ # SLEN/SIDX model
580
+ return slen_cardinalities | sidx_cardinalities
581
+ case _:
582
+ raise ValueError(f"Invalid positional encoding: {has_slen=}, {has_ridx=}, {has_sdec=}")
581
583
 
582
584
 
583
585
  def persist_data_part(df: pd.DataFrame, output_path: Path, infix: str):
@@ -27,6 +27,7 @@ from transformers import (
27
27
  PretrainedConfig,
28
28
  PreTrainedModel,
29
29
  )
30
+ from transformers.quantizers import AutoQuantizationConfig
30
31
 
31
32
  from mostlyai.engine._language.lstm import LSTMFromScratchConfig
32
33
 
@@ -84,7 +85,9 @@ def load_base_model_and_config(
84
85
  else:
85
86
  attn_implementation = None
86
87
  torch_dtype = torch.float32
87
- if is_gpu_training and is_bitsandbytes_available:
88
+ if hasattr(config, "quantization_config"):
89
+ quantization_config = AutoQuantizationConfig.from_dict(config.quantization_config)
90
+ elif is_gpu_training and is_bitsandbytes_available:
88
91
  quantization_config = BitsAndBytesConfig(
89
92
  load_in_4bit=True,
90
93
  bnb_4bit_quant_type="nf4",
@@ -117,7 +120,7 @@ def load_base_model_and_config(
117
120
  quantization_config=quantization_config,
118
121
  torch_dtype=torch_dtype,
119
122
  )
120
- if quantization_config:
123
+ if isinstance(quantization_config, BitsAndBytesConfig):
121
124
  # convert all non-kbit layers to float32
122
125
  model = prepare_model_for_kbit_training(model, use_gradient_checkpointing=False)
123
126
  if is_gpu_training and model.supports_gradient_checkpointing:
@@ -14,7 +14,6 @@
14
14
 
15
15
  import contextlib
16
16
  import importlib
17
- import json
18
17
  import logging
19
18
  import os
20
19
  import platform
@@ -72,10 +71,10 @@ def decode_buffered_samples(
72
71
 
73
72
  def parse_json(x, columns: list[str]):
74
73
  try:
75
- parsed_x = json_repair.loads(x)
74
+ parsed_x = json_repair.loads(x, stream_stable=True)
76
75
  if not isinstance(parsed_x, dict):
77
76
  raise ValueError("parsed_x has to be a dictionary")
78
- except (json.decoder.JSONDecodeError, ValueError):
77
+ except Exception:
79
78
  parsed_x = {}
80
79
  return [parsed_x.get(c, INVALID_VALUE) for c in columns]
81
80
 
@@ -183,9 +182,9 @@ def generate(
183
182
  enable_flexible_generation = model_configs.get("enable_flexible_generation", True)
184
183
  _LOG.info(f"{enable_flexible_generation=}")
185
184
 
186
- # resolve potential conflict between sample_seed and sample_size
185
+ # resolve potential conflict between seed_data and sample_size
187
186
  if seed_data is not None:
188
- assert sample_size is None, "either sample_seed or sample_size can be provided, not both"
187
+ assert sample_size is None, "either seed_data or sample_size can be provided, not both"
189
188
  sample_size = len(seed_data)
190
189
 
191
190
  if has_context:
@@ -222,7 +221,7 @@ def generate(
222
221
  ctx_primary_key = tgt_context_key = DUMMY_CONTEXT_KEY
223
222
  ctx_data = pd.DataFrame({ctx_primary_key: range(sample_size)})
224
223
 
225
- # ensure sample_seed exists; ensure valid columns
224
+ # ensure seed_data exists; ensure valid columns
226
225
  if seed_data is None:
227
226
  # build dummy seed
228
227
  seed_data = pd.DataFrame(index=list(range(sample_size)))
@@ -230,7 +229,7 @@ def generate(
230
229
  _LOG.info(f"{seed_data.shape=}")
231
230
 
232
231
  if not enable_flexible_generation:
233
- # validate sample_seed maintains the same column order as the one from training
232
+ # validate seed_data maintains the same column order as the one from training
234
233
  seed_columns = seed_data.columns.tolist()
235
234
  if seed_columns != tgt_text_columns[: len(seed_columns)]:
236
235
  raise ValueError(
@@ -238,7 +237,7 @@ def generate(
238
237
  "A change in column order is only permitted for models that were trained with `enable_flexible_generation=True`."
239
238
  )
240
239
 
241
- # sanity check: at this point sample seed and context data should have the same number of rows
240
+ # sanity check: at this point seed data and context data should have the same number of rows
242
241
  assert len(seed_data) == len(ctx_data)
243
242
 
244
243
  # early exit in case generation context is empty
@@ -306,14 +305,14 @@ def generate(
306
305
  samples_processed = 0
307
306
  while samples_processed < sample_size:
308
307
  encoded_ctx_batch = encoded_ctx_data.iloc[samples_processed : samples_processed + batch_size]
309
- sample_seed_batch = seed_data.iloc[samples_processed : samples_processed + batch_size]
308
+ seed_data_batch = seed_data.iloc[samples_processed : samples_processed + batch_size]
310
309
  ctx_batch = ctx_data.iloc[samples_processed : samples_processed + batch_size]
311
310
  ctx_keys = ctx_batch[ctx_primary_key]
312
311
 
313
312
  if enforce_json_output and not initialize_logits_processors_once:
314
313
  t0 = time.time()
315
314
  schemas = create_schemas(
316
- seed_df=sample_seed_batch,
315
+ seed_df=seed_data_batch,
317
316
  stats=tgt_stats,
318
317
  rare_category_replacement_method=rare_category_replacement_method,
319
318
  )
@@ -328,7 +327,7 @@ def generate(
328
327
  total_tokenize_fn_time += metrics.tokenize_time
329
328
  total_generate_fn_time += metrics.generate_time
330
329
 
331
- buffer.add((outputs, ctx_keys, sample_seed_batch))
330
+ buffer.add((outputs, ctx_keys, seed_data_batch))
332
331
  if buffer.is_full():
333
332
  decoded_data = decode_buffered_samples(
334
333
  buffer, engine.tokenizer, tgt_stats, tgt_context_key, max_new_tokens
@@ -319,7 +319,10 @@ def train(
319
319
  )
320
320
  ):
321
321
  device = torch.device("cuda:0")
322
- _LOG.info("device set to single gpu (cuda:0) because model is too small or differential privacy is enabled")
322
+ if torch.cuda.device_count() > 1:
323
+ _LOG.info(
324
+ "device set to single gpu (cuda:0) because model is too small or differential privacy is enabled"
325
+ )
323
326
 
324
327
  if not with_dp:
325
328
  if device.type == "cuda":
@@ -42,7 +42,7 @@ def prepend_grammar_root_with_space(grammar: str) -> str:
42
42
  return grammar.replace(start_of_grammar, start_of_grammar_with_space)
43
43
 
44
44
 
45
- def ensure_seed_can_be_tokenized(sample_seed: pd.DataFrame, tokenizer: PreTrainedTokenizerBase) -> pd.DataFrame:
45
+ def ensure_seed_can_be_tokenized(seed_data: pd.DataFrame, tokenizer: PreTrainedTokenizerBase) -> pd.DataFrame:
46
46
  def transform(x: str | pd._libs.missing.NAType) -> str:
47
47
  if pd.isna(x):
48
48
  null = tokenizer.decode(tokenizer.encode(JSON_NULL), skip_special_tokens=True)
@@ -53,7 +53,7 @@ def ensure_seed_can_be_tokenized(sample_seed: pd.DataFrame, tokenizer: PreTraine
53
53
  # skip tokens unseen during training
54
54
  return tokenizer.decode(tokenizer.encode(x), skip_special_tokens=True)
55
55
 
56
- return sample_seed.astype(STRING).map(transform)
56
+ return seed_data.astype(STRING).map(transform)
57
57
 
58
58
 
59
59
  def create_schemas(
@@ -36,6 +36,8 @@ from torch import nn
36
36
  from mostlyai.engine._common import (
37
37
  CTXFLT,
38
38
  CTXSEQ,
39
+ RIDX_SUB_COLUMN_PREFIX,
40
+ SLEN_SUB_COLUMN_PREFIX,
39
41
  get_columns_from_cardinalities,
40
42
  get_sub_columns_from_cardinalities,
41
43
  get_sub_columns_lookup,
@@ -214,9 +216,19 @@ class Embedders(nn.Module):
214
216
  self.embedders = nn.ModuleDict()
215
217
 
216
218
  # embedding layers for each sub column defined in cardinalities
219
+ last_slen_sub_col = next(
220
+ (sub_col for sub_col in reversed(self.cardinalities) if sub_col.startswith(SLEN_SUB_COLUMN_PREFIX)), None
221
+ )
222
+ last_ridx_sub_col = next(
223
+ (sub_col for sub_col in reversed(self.cardinalities) if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX)), None
224
+ )
217
225
  for sub_col, dim_input in self.cardinalities.items():
218
226
  dim_output = _embedding_heuristic(id=self.id(sub_col), model_size=model_size, dim_input=dim_input)
219
227
  embedder = nn.Embedding(num_embeddings=dim_input, embedding_dim=dim_output, device=device)
228
+ # the embeddings of the last slen and ridx sub columns are never used
229
+ # so we explicitly freeze them to make opacus not complain about "per sample gradient is not initialized"
230
+ if sub_col in [last_ridx_sub_col, last_slen_sub_col]:
231
+ embedder.weight.requires_grad = False
220
232
  self.add(sub_column=sub_col, embedder=embedder)
221
233
  self.dims.append(dim_output)
222
234
 
@@ -715,7 +727,7 @@ def _make_permutation_mask(
715
727
  # create mask in provided order
716
728
  order = torch.tensor([columns.index(c) for c in column_order], dtype=torch.int32)
717
729
  elif is_sequential and n_cols >= 1:
718
- # create mask in random order, but keep SLEN/SIDX/SDEC column at first position
730
+ # create mask in random order, but keep positional columns at first position
719
731
  order = torch.randperm(n_cols - 1) + 1
720
732
  order = torch.cat((torch.zeros(1, dtype=torch.int32), order), dim=0)
721
733
  else:
@@ -1241,13 +1253,17 @@ class SequentialModel(nn.Module):
1241
1253
  if mode == "trn":
1242
1254
  # forward pass through sub column embedders
1243
1255
  tgt_embeds = self.embedders(x)
1256
+ tgt_embeds_slen_ridx_masked = {
1257
+ k: torch.zeros_like(v) if k.startswith((RIDX_SUB_COLUMN_PREFIX, SLEN_SUB_COLUMN_PREFIX)) else v
1258
+ for k, v in tgt_embeds.items()
1259
+ }
1244
1260
 
1245
1261
  # forward pass through column embedders
1246
- tgt_col_embeds = self.column_embedders(tgt_embeds)
1262
+ tgt_col_embeds = self.column_embedders(tgt_embeds_slen_ridx_masked)
1247
1263
 
1248
1264
  # history
1249
1265
  # time shift: remove last time step; add zeros for first time step; add randoms for all others
1250
- embeddings = torch.cat(list(tgt_embeds.values()), dim=-1)
1266
+ embeddings = torch.cat(list(tgt_embeds_slen_ridx_masked.values()), dim=-1)
1251
1267
  history_in = embeddings[:, :-1, :]
1252
1268
  history_in = nn.ConstantPad2d((0, 0, 1, 0), 0)(history_in)
1253
1269
  history, _ = self.history_compressor(history_in)
@@ -1274,9 +1290,17 @@ class SequentialModel(nn.Module):
1274
1290
  masked_col_embeds = [torch.mul(col_mask[lookup.col_idx, :].int(), col_embeddings)]
1275
1291
 
1276
1292
  # collect previous sub column embeddings for current column
1277
- prev_sub_col_embeds = [
1278
- tgt_embeds[sub_col] for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1279
- ]
1293
+ prev_sub_col_embeds = {
1294
+ sub_col: tgt_embeds[sub_col]
1295
+ for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1296
+ }
1297
+ if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1298
+ # RIDX sub-columns should not see SLEN sub-columns
1299
+ prev_sub_col_embeds = {
1300
+ k: torch.zeros_like(v) if k.startswith(SLEN_SUB_COLUMN_PREFIX) else v
1301
+ for k, v in prev_sub_col_embeds.items()
1302
+ }
1303
+ prev_sub_col_embeds = list(prev_sub_col_embeds.values())
1280
1304
 
1281
1305
  # regressor
1282
1306
  regressor_in = context_history + masked_col_embeds + prev_sub_col_embeds
@@ -1335,10 +1359,17 @@ class SequentialModel(nn.Module):
1335
1359
 
1336
1360
  else: # sample from distribution
1337
1361
  # collect previous sub column embeddings for current column
1338
- prev_sub_col_embeds = [
1339
- tgt_embeds[sub_col]
1362
+ prev_sub_col_embeds = {
1363
+ sub_col: tgt_embeds[sub_col]
1340
1364
  for sub_col in self.tgt_sub_columns[lookup.sub_col_offset : lookup.sub_col_cum]
1341
- ]
1365
+ }
1366
+ if sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1367
+ # RIDX sub-columns should not see SLEN sub-columns
1368
+ prev_sub_col_embeds = {
1369
+ k: torch.zeros_like(v) if k.startswith(SLEN_SUB_COLUMN_PREFIX) else v
1370
+ for k, v in prev_sub_col_embeds.items()
1371
+ }
1372
+ prev_sub_col_embeds = list(prev_sub_col_embeds.values())
1342
1373
 
1343
1374
  # regressor
1344
1375
  regressor_in = context_history + [col_embeddings] + prev_sub_col_embeds
@@ -1362,16 +1393,20 @@ class SequentialModel(nn.Module):
1362
1393
 
1363
1394
  # update current sub column embedding
1364
1395
  tgt_embeds[sub_col] = self.embedders.get(sub_col)(out)
1396
+ tgt_embeds_slen_ridx_masked = {
1397
+ k: torch.zeros_like(v) if k.startswith((SLEN_SUB_COLUMN_PREFIX, RIDX_SUB_COLUMN_PREFIX)) else v
1398
+ for k, v in tgt_embeds.items()
1399
+ }
1365
1400
 
1366
1401
  # update current column embedding
1367
1402
  if sub_col in self.tgt_last_sub_cols:
1368
1403
  col_sub_cols = self.tgt_column_sub_columns[lookup.col_name]
1369
- col_embed_in = torch.cat([tgt_embeds[sc] for sc in col_sub_cols], dim=-1)
1404
+ col_embed_in = torch.cat([tgt_embeds_slen_ridx_masked[sc] for sc in col_sub_cols], dim=-1)
1370
1405
  tgt_col_embeds[lookup.col_name] = self.column_embedders.get(lookup.col_name)(col_embed_in)
1371
1406
  col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1372
1407
 
1373
1408
  # update history and hidden state
1374
- history_in = torch.cat([v for v in tgt_embeds.values()], dim=-1)
1409
+ history_in = torch.cat([v for v in tgt_embeds_slen_ridx_masked.values()], dim=-1)
1375
1410
  history, history_state = self.history_compressor(history_in, history_state=history_state)
1376
1411
 
1377
1412
  # order outputs according to tgt_sub_columns
@@ -24,13 +24,13 @@ from mostlyai.engine._common import (
24
24
  ARGN_COLUMN,
25
25
  ARGN_PROCESSOR,
26
26
  ARGN_TABLE,
27
- SDEC_SUB_COLUMN_PREFIX,
27
+ RIDX_SUB_COLUMN_PREFIX,
28
28
  SIDX_SUB_COLUMN_PREFIX,
29
29
  SLEN_SUB_COLUMN_PREFIX,
30
30
  TGT,
31
31
  ProgressCallback,
32
32
  ProgressCallbackWrapper,
33
- encode_slen_sidx_sdec,
33
+ encode_positional_column,
34
34
  get_argn_name,
35
35
  get_sequence_length_stats,
36
36
  is_a_list,
@@ -129,33 +129,33 @@ def _encode_partition(
129
129
  n_jobs=n_jobs,
130
130
  )
131
131
  # pad each list with one extra item
132
- df_ctx = pad_horizontally(df_ctx, padding_value=0, right=False)
132
+ df_ctx = pad_ctx_sequences(df_ctx)
133
133
 
134
134
  if is_sequential:
135
135
  assert isinstance(tgt_context_key, str)
136
136
  # trim sequences to (privacy-protected) max_len
137
137
  max_len = seq_len_stats["max"]
138
138
  df = df[df.groupby(tgt_context_key).cumcount() < max_len].reset_index(drop=True)
139
- # enrich with sequence lengths and sequence indexes
140
- df = _enrich_slen_sidx_sdec(df, tgt_context_key, max_len)
141
- # flatten to list columns
142
- df = flatten_frame(df, tgt_context_key)
139
+ # pad each list with one extra item
140
+ df = pad_tgt_sequences(df, context_key=tgt_context_key)
143
141
  # add empty records for IDs, that are present in context, but not in target; i.e., for zero-sequence records
144
142
  if has_context:
145
- zero_seq_ids = list(set(df_ctx[ctx_primary_key]) - set(df[tgt_context_key]))
146
- df_miss = pd.DataFrame({tgt_context_key: zero_seq_ids})
147
- df_pads = pd.DataFrame({c: [[]] for c in df.columns if c != tgt_context_key})
148
- df_miss = df_miss.merge(df_pads, how="cross")
149
- df = pd.concat([df, df_miss], axis=0).reset_index(drop=True)
150
- # pad each list with one extra item
151
- df = pad_horizontally(df, padding_value=0, right=True)
143
+ zero_seq_ids = set(df_ctx[ctx_primary_key]) - set(df[tgt_context_key])
144
+ df_miss = pd.DataFrame(
145
+ [{tgt_context_key: i, **{c: 0 for c in df.columns if c != tgt_context_key}} for i in zero_seq_ids]
146
+ )
147
+ df = pd.concat([df, df_miss], ignore_index=True)
148
+ # enrich with positional columns
149
+ df = _enrich_positional_columns(df, tgt_context_key, max_len)
150
+ # flatten to list columns
151
+ df = flatten_frame(df, tgt_context_key)
152
152
  elif has_context:
153
153
  # add 0-rows for IDs, that are present in context, but not in target; i.e., for zero-sequence records
154
- zero_seq_ids = list(set(df_ctx[ctx_primary_key]) - set(df[tgt_context_key]))
155
- df_miss = pd.DataFrame({tgt_context_key: zero_seq_ids})
156
- df_pads = pd.DataFrame({c: [0] for c in df.columns if c != tgt_context_key})
157
- df_miss = df_miss.merge(df_pads, how="cross")
158
- df = pd.concat([df, df_miss], axis=0).reset_index(drop=True)
154
+ zero_seq_ids = set(df_ctx[ctx_primary_key]) - set(df[tgt_context_key])
155
+ df_miss = pd.DataFrame(
156
+ [{tgt_context_key: i, **{c: 0 for c in df.columns if c != tgt_context_key}} for i in zero_seq_ids]
157
+ )
158
+ df = pd.concat([df, df_miss], ignore_index=True)
159
159
  # ensure that max 1 item is retained per context_id for flat mode
160
160
  df = df[df.groupby(tgt_context_key).cumcount() < 1]
161
161
 
@@ -380,29 +380,61 @@ def flatten_frame(df: pd.DataFrame, group_key: str) -> pd.DataFrame:
380
380
  return flattened_data
381
381
 
382
382
 
383
- def _enrich_slen_sidx_sdec(df: pd.DataFrame, context_key: str, max_seq_len: int) -> pd.DataFrame:
383
+ def _enrich_positional_columns(df: pd.DataFrame, context_key: str, max_seq_len: int) -> pd.DataFrame:
384
384
  df = df.reset_index(drop=True)
385
- slen = df.groupby(context_key)[context_key].transform("size") # sequence length
386
385
  sidx = df.groupby(context_key).cumcount(ascending=True) # sequence index
387
- sdec = (10 * sidx / slen.clip(lower=1)).astype(int) # sequence index decile
388
- slen = encode_slen_sidx_sdec(slen, max_seq_len=max_seq_len, prefix=SLEN_SUB_COLUMN_PREFIX)
389
- sidx = encode_slen_sidx_sdec(sidx, max_seq_len=max_seq_len, prefix=SIDX_SUB_COLUMN_PREFIX)
390
- sdec = encode_slen_sidx_sdec(sdec, max_seq_len=max_seq_len, prefix=SDEC_SUB_COLUMN_PREFIX)
391
- df = pd.concat([slen, sidx, sdec, df], axis=1)
386
+ slen = df.groupby(context_key)[context_key].transform("size") - 1 # sequence length; -1 to account for padding
387
+ ridx = df.groupby(context_key).cumcount(ascending=False) # sequence remainder
388
+ sidx = encode_positional_column(sidx, max_seq_len=max_seq_len, prefix=SIDX_SUB_COLUMN_PREFIX)
389
+ slen = encode_positional_column(slen, max_seq_len=max_seq_len, prefix=SLEN_SUB_COLUMN_PREFIX)
390
+ ridx = encode_positional_column(ridx, max_seq_len=max_seq_len, prefix=RIDX_SUB_COLUMN_PREFIX)
391
+ df = pd.concat([sidx, slen, ridx, df], axis=1)
392
392
  return df
393
393
 
394
394
 
395
- def pad_horizontally(df: pd.DataFrame, padding_value: int, right=True) -> pd.DataFrame:
396
- if df.shape[0] == 0:
397
- return df
398
- list_cols = [c for c in df.columns if is_a_list(df.loc[0, c])]
395
+ def pad_tgt_sequences(df: pd.DataFrame, context_key: str, padding_value: int = 0) -> pd.DataFrame:
396
+ """
397
+ Pad one extra row to for each subject in the target data frame.
399
398
 
400
- def pad_right(x):
401
- return x + [padding_value] if len(x) == 0 else x
399
+ Args:
400
+ df: Exploded (unflattened) target data frame. Each event is a row.
401
+ context_key: Context key.
402
+ padding_value: Value to pad with for columns other than context_key.
402
403
 
403
- def pad_left(x):
404
- return [padding_value] + x if len(x) == 0 else x
404
+ Returns:
405
+ Padded target data frame.
406
+ """
407
+
408
+ def pad_row(x):
409
+ return pd.concat(
410
+ [
411
+ x,
412
+ pd.DataFrame(
413
+ {col: [padding_value] for col in x.columns if col not in [context_key]}
414
+ | {context_key: [x.iloc[0][context_key]]}
415
+ ),
416
+ ],
417
+ axis=0,
418
+ )
419
+
420
+ return df.groupby(context_key)[df.columns].apply(pad_row).reset_index(drop=True)
421
+
422
+
423
+ def pad_ctx_sequences(df: pd.DataFrame, padding_value: int = 0) -> pd.DataFrame:
424
+ """
425
+ Pad one extra item to the context sequences with a given padding value.
426
+
427
+ Args:
428
+ df: Flattened context data frame.
429
+ padding_value: Value to pad with.
405
430
 
431
+ Returns:
432
+ Padded context data frame.
433
+ """
434
+ if df.shape[0] == 0:
435
+ return df
436
+ list_cols = [c for c in df.columns if is_a_list(df.loc[0, c])]
406
437
  for col in list_cols:
407
- df[col] = df[col].apply(pad_right if right else pad_left)
438
+ # Note: only pad empty sequences to keep the backward compatibility
439
+ df[col] = df[col].apply(lambda x: x + [padding_value] if len(x) == 0 else x)
408
440
  return df