mostlyai-engine 1.5.2__tar.gz → 1.5.5__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.5.2 → mostlyai_engine-1.5.5}/PKG-INFO +1 -1
  2. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/__init__.py +1 -1
  3. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_common.py +19 -16
  4. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/common.py +3 -3
  5. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_tabular/argn.py +34 -18
  6. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_tabular/generation.py +13 -10
  7. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/analysis.py +26 -7
  8. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/pyproject.toml +1 -1
  9. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/.gitignore +0 -0
  10. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/LICENSE +0 -0
  11. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/README.md +0 -0
  12. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_dtypes.py +0 -0
  13. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  14. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  15. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  16. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  17. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  18. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  19. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  20. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  21. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  22. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  23. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  24. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  25. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  26. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/__init__.py +0 -0
  27. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/encoding.py +0 -0
  28. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/engine/__init__.py +0 -0
  29. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/engine/base.py +0 -0
  30. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  31. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  32. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/generation.py +0 -0
  33. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/lstm.py +0 -0
  34. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  35. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/training.py +0 -0
  36. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
  37. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_memory.py +0 -0
  38. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_tabular/__init__.py +0 -0
  39. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_tabular/common.py +0 -0
  40. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_tabular/encoding.py +0 -0
  41. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_tabular/fairness.py +0 -0
  42. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_tabular/training.py +0 -0
  43. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_training_utils.py +0 -0
  44. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/_workspace.py +0 -0
  45. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/domain.py +0 -0
  46. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/encoding.py +0 -0
  47. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/generation.py +0 -0
  48. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/logging.py +0 -0
  49. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/random_state.py +0 -0
  50. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/mostlyai/engine/splitting.py +0 -0
  51. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.5}/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.2
3
+ Version: 1.5.5
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.2"
25
+ __version__ = "1.5.5"
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):
@@ -401,12 +407,11 @@ def get_sub_columns_lookup(
401
407
  return sub_cols_lookup
402
408
 
403
409
 
404
- class CtxSequenceLengthError(Exception):
405
- """Error raised when the cols of the same table do not have the same stats value"""
406
-
407
-
408
410
  def get_ctx_sequence_length(ctx_stats: dict, key: str) -> dict[str, int]:
409
- seq_stats: dict[str, int] = {}
411
+ """
412
+ Get the stats of sequence lengths from the first column_stats of each context table
413
+ """
414
+ ctxseq_stats: dict[str, int] = {}
410
415
 
411
416
  for column_stats in ctx_stats.get("columns", {}).values():
412
417
  if "seq_len" in column_stats:
@@ -414,12 +419,10 @@ def get_ctx_sequence_length(ctx_stats: dict, key: str) -> dict[str, int]:
414
419
  argn_processor=column_stats[ARGN_PROCESSOR],
415
420
  argn_table=column_stats[ARGN_TABLE],
416
421
  )
417
- cur_value = seq_stats.get(table)
418
- if cur_value and cur_value != column_stats["seq_len"][key]:
419
- raise CtxSequenceLengthError()
420
- seq_stats[table] = column_stats["seq_len"][key]
422
+ if table not in ctxseq_stats:
423
+ ctxseq_stats[table] = column_stats["seq_len"][key]
421
424
 
422
- return seq_stats
425
+ return ctxseq_stats
423
426
 
424
427
 
425
428
  def get_max_data_points_per_sample(stats: dict) -> int:
@@ -541,12 +544,11 @@ def decode_positional_column(df_encoded: pd.DataFrame, max_seq_len: int, prefix:
541
544
 
542
545
 
543
546
  def get_positional_cardinalities(
544
- 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
545
548
  ) -> 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
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
550
552
 
551
553
  if max_seq_len < SIDX_RIDX_DIGIT_ENCODING_THRESHOLD:
552
554
  # encode positional columns as numeric_discrete
@@ -867,6 +869,7 @@ def dp_non_rare(value_counts: dict[str, int], epsilon: float, threshold: int = 5
867
869
  noisy_counts = np.clip(np.array(list(value_counts.values())) + noise, 0, None).astype(int)
868
870
  for i, cat in enumerate(value_counts):
869
871
  value_counts[cat] = noisy_counts[i]
872
+ # NOTE: total_counts can be 0 in the edge case when the column only has null values
870
873
  total_counts = sum(value_counts.values())
871
874
 
872
875
  # 2. Collect all categories whose noisy count >= threshold
@@ -874,7 +877,7 @@ def dp_non_rare(value_counts: dict[str, int], epsilon: float, threshold: int = 5
874
877
 
875
878
  # 3. Compute the non-rare ratio
876
879
  noisy_total_counts = sum(selected.values())
877
- non_rare_ratio = noisy_total_counts / total_counts
880
+ non_rare_ratio = noisy_total_counts / total_counts if total_counts > 0 else 0
878
881
 
879
882
  return list(selected.keys()), non_rare_ratio
880
883
 
@@ -97,10 +97,10 @@ def load_base_model_and_config(
97
97
  else:
98
98
  quantization_config = None
99
99
 
100
- if device.type == "cuda" and device.index is not None:
101
- device_map = str(device)
102
- else:
100
+ if device.type == "cuda" and device.index is None:
103
101
  device_map = "auto"
102
+ else: # device is `cpu` or `cuda:0` (when using single GPU on a multi-GPU instance)
103
+ device_map = str(device)
104
104
 
105
105
  if hasattr(config, "text_config") and hasattr(config, "vision_config"):
106
106
  config.text_config.use_cache = use_cache
@@ -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)
@@ -1256,26 +1253,28 @@ class SequentialModel(nn.Module):
1256
1253
  if context is None:
1257
1254
  context = self.context_compressor(x)
1258
1255
 
1259
- # SLEN is only masked for models with RIDX
1260
1256
  has_ridx = any(sub_col.startswith(RIDX_SUB_COLUMN_PREFIX) for sub_col in self.tgt_cardinalities)
1261
- masked_positional_columns = (
1257
+
1258
+ # SLEN and RIDX are masked for history
1259
+ # NOTE: SLEN is not masked for models without RIDX (backwards compatibility)
1260
+ history_masked_sub_cols = (
1262
1261
  (SLEN_SUB_COLUMN_PREFIX, RIDX_SUB_COLUMN_PREFIX) if has_ridx else (RIDX_SUB_COLUMN_PREFIX,)
1263
1262
  )
1263
+ # SLEN is masked for column embeddings
1264
+ # NOTE: SLEN is not masked for models without RIDX (backwards compatibility)
1265
+ col_embeddings_masked_sub_cols = (SLEN_SUB_COLUMN_PREFIX,) if has_ridx else ()
1264
1266
 
1265
1267
  outputs = {}
1266
1268
  if mode == "trn":
1267
1269
  # forward pass through sub column embedders
1268
1270
  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
1271
 
1276
1272
  # history
1277
1273
  # 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)
1274
+ tgt_embeds_for_history = {
1275
+ k: torch.zeros_like(v) if k.startswith(history_masked_sub_cols) else v for k, v in tgt_embeds.items()
1276
+ }
1277
+ embeddings = torch.cat(list(tgt_embeds_for_history.values()), dim=-1)
1279
1278
  history_in = embeddings[:, :-1, :]
1280
1279
  history_in = nn.ConstantPad2d((0, 0, 1, 0), 0)(history_in)
1281
1280
  history, _ = self.history_compressor(history_in)
@@ -1287,6 +1286,13 @@ class SequentialModel(nn.Module):
1287
1286
  flat_ctx = self._repeat_flat_context(flat_ctx, history.size(1))
1288
1287
  context_history = [torch.cat(flat_ctx + seq_ctx + [history], -1)]
1289
1288
 
1289
+ # forward pass through column embedders
1290
+ tgt_embeds_for_col_embeddings = {
1291
+ k: torch.zeros_like(v) if k.startswith(col_embeddings_masked_sub_cols) else v
1292
+ for k, v in tgt_embeds.items()
1293
+ }
1294
+ tgt_col_embeds = self.column_embedders(tgt_embeds_for_col_embeddings)
1295
+
1290
1296
  # create batch-wise permutation mask
1291
1297
  col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1292
1298
  col_mask = _make_permutation_mask(
@@ -1327,7 +1333,9 @@ class SequentialModel(nn.Module):
1327
1333
  return outputs, {}
1328
1334
 
1329
1335
  else: # mode == "gen"
1336
+ is_0th_step = False
1330
1337
  if history is None or history_state is None:
1338
+ is_0th_step = True
1331
1339
  # initialize history
1332
1340
  history_in = torch.cat(
1333
1341
  [
@@ -1401,24 +1409,32 @@ class SequentialModel(nn.Module):
1401
1409
  )
1402
1410
 
1403
1411
  # update timestep output
1412
+ if is_0th_step and sub_col.startswith(RIDX_SUB_COLUMN_PREFIX):
1413
+ # overwrite output for RIDX sub-columns on 0th step with SLEN sub-columns
1414
+ slen_sub_col = sub_col.replace(RIDX_SUB_COLUMN_PREFIX, SLEN_SUB_COLUMN_PREFIX)
1415
+ out = outputs[slen_sub_col] if slen_sub_col in outputs else out
1404
1416
  outputs[sub_col] = out
1405
1417
 
1406
1418
  # update current sub column embedding
1407
1419
  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
1420
+ tgt_embeds_for_history = {
1421
+ k: torch.zeros_like(v) if k.startswith(history_masked_sub_cols) else v
1422
+ for k, v in tgt_embeds.items()
1423
+ }
1424
+ tgt_embeds_for_col_embeddings = {
1425
+ k: torch.zeros_like(v) if k.startswith(col_embeddings_masked_sub_cols) else v
1410
1426
  for k, v in tgt_embeds.items()
1411
1427
  }
1412
1428
 
1413
1429
  # update current column embedding
1414
1430
  if sub_col in self.tgt_last_sub_cols:
1415
1431
  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)
1432
+ col_embed_in = torch.cat([tgt_embeds_for_col_embeddings[sc] for sc in col_sub_cols], dim=-1)
1417
1433
  tgt_col_embeds[lookup.col_name] = self.column_embedders.get(lookup.col_name)(col_embed_in)
1418
1434
  col_embeddings = torch.cat(list(tgt_col_embeds.values()), dim=-1)
1419
1435
 
1420
1436
  # update history and hidden state
1421
- history_in = torch.cat([v for v in tgt_embeds_positional_masked.values()], dim=-1)
1437
+ history_in = torch.cat(list(tgt_embeds_for_history.values()), dim=-1)
1422
1438
  history, history_state = self.history_compressor(history_in, history_state=history_state)
1423
1439
 
1424
1440
  # order outputs according to tgt_sub_columns
@@ -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)
@@ -1023,7 +1028,7 @@ def generate(
1023
1028
 
1024
1029
  # fix SLEN by propagating sampled SLEN from first step
1025
1030
  slen_vals = {}
1026
- if seq_step > 0:
1031
+ if has_slen and seq_step > 0:
1027
1032
  slen = out_df[SLEN_SUB_COLUMN_PREFIX]
1028
1033
  slen = encode_positional_column(slen, max_seq_len=seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX)
1029
1034
  slen_vals = {
@@ -1102,16 +1107,14 @@ def generate(
1102
1107
  out_df[SIDX_SUB_COLUMN_PREFIX] = decode_positional_column(
1103
1108
  out_df, seq_len_max, prefix=SIDX_SUB_COLUMN_PREFIX
1104
1109
  )
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)
1110
+ if has_slen:
1111
+ out_df[SLEN_SUB_COLUMN_PREFIX] = decode_positional_column(
1112
+ out_df, seq_len_max, prefix=SLEN_SUB_COLUMN_PREFIX
1113
+ )
1114
+ out_df[SLEN_SUB_COLUMN_PREFIX] = out_df[SLEN_SUB_COLUMN_PREFIX].clip(lower=seq_len_min)
1109
1115
  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]
1116
+ out_df[RIDX_SUB_COLUMN_PREFIX] = decode_positional_column(
1117
+ out_df, seq_len_max, prefix=RIDX_SUB_COLUMN_PREFIX
1115
1118
  )
1116
1119
  out_df[RIDX_SUB_COLUMN_PREFIX] = out_df[RIDX_SUB_COLUMN_PREFIX].clip(
1117
1120
  lower=seq_len_min - seq_step, upper=seq_len_max
@@ -339,13 +339,25 @@ def _analyze_reduce(
339
339
  stats_list = [read_json(file) for file in stats_files]
340
340
  stats: dict[str, Any] = {"columns": {}}
341
341
 
342
+ # check how many context tables have sequential context
343
+ if mode == "ctx":
344
+ ctxseq_stats = {}
345
+ ctxseq_tables = []
346
+ for column, column_stats in stats_list[0]["columns"].items():
347
+ if "seq_len" in column_stats:
348
+ table_name = column.split(TABLE_COLUMN_INFIX)[0]
349
+ if table_name not in ctxseq_tables:
350
+ ctxseq_tables.append(table_name)
351
+ n_ctxseq_tables = len(ctxseq_tables)
352
+ _LOG.info(f"{n_ctxseq_tables = }")
353
+
342
354
  encoding_types = {
343
355
  column: column_stats.get("encoding_type") for column, column_stats in stats_list[0]["columns"].items()
344
356
  }
345
357
 
346
- # ctx: distribute the privacy budget across all columns
358
+ # ctx: distribute the privacy budget across all columns + sequence lengths of n_ctxseq_tables
347
359
  # tgt: distribute the privacy budget across all columns + sequence length
348
- n_dp_splits = len(encoding_types) if mode == "ctx" else len(encoding_types) + 1
360
+ n_dp_splits = len(encoding_types) + n_ctxseq_tables if mode == "ctx" else len(encoding_types) + 1
349
361
  _LOG.info(f"{value_protection = }")
350
362
  if value_protection_epsilon is not None and n_dp_splits > 0:
351
363
  _LOG.info(f"epsilon for analyzing each column and sequence length: {value_protection_epsilon / n_dp_splits}")
@@ -364,13 +376,13 @@ def _analyze_reduce(
364
376
  stats["columns"][column] = {"encoding_type": encoding_type}
365
377
  continue
366
378
 
367
- analyze_reduce_column_args = {
368
- "stats_list": column_stats_list,
379
+ value_protection_args = {
369
380
  "value_protection": value_protection,
370
381
  "value_protection_epsilon": value_protection_epsilon / n_dp_splits
371
382
  if value_protection_epsilon is not None
372
383
  else None,
373
384
  }
385
+ analyze_reduce_column_args = {"stats_list": column_stats_list} | value_protection_args
374
386
 
375
387
  match encoding_type:
376
388
  case ModelEncodingType.tabular_categorical:
@@ -413,9 +425,16 @@ def _analyze_reduce(
413
425
  if encoding_type in _VALUE_PROTECTION_ENCODING_TYPES:
414
426
  stats_col = {"value_protection": value_protection} | stats_col
415
427
 
416
- is_flat_column = "seq_len" not in column_stats_list[0]
417
- if not is_flat_column:
418
- stats_col["seq_len"] = _analyze_reduce_seq_len([column_stats_list[0]["seq_len"]])
428
+ is_ctxseq_column = "seq_len" in column_stats_list[0]
429
+ if is_ctxseq_column:
430
+ table_name = column.split(TABLE_COLUMN_INFIX)[0]
431
+ # only get the lengths from the first column of a ctxseq table and reuse the stats later
432
+ if table_name not in ctxseq_stats:
433
+ ctxseq_stats[table_name] = _analyze_reduce_seq_len(
434
+ stats_list=[column_stats_list[0]["seq_len"]], **value_protection_args
435
+ )
436
+ _LOG.info(f"analyzed sequence length for context table `{table_name}`")
437
+ stats_col["seq_len"] = ctxseq_stats[table_name]
419
438
 
420
439
  is_language_column = encoding_type in (
421
440
  ModelEncodingType.language_text,
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "mostlyai-engine"
3
- version = "1.5.2"
3
+ version = "1.5.5"
4
4
  description = "Synthetic Data Engine"
5
5
  authors = [{ name = "MOSTLY AI", email = "dev@mostly.ai" }]
6
6
  requires-python = ">=3.10"
File without changes