mostlyai-engine 1.5.2__tar.gz → 1.5.4__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.4}/PKG-INFO +1 -1
  2. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/__init__.py +1 -1
  3. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_common.py +9 -11
  4. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/common.py +3 -3
  5. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/analysis.py +26 -7
  6. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/pyproject.toml +1 -1
  7. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/.gitignore +0 -0
  8. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/LICENSE +0 -0
  9. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/README.md +0 -0
  10. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_dtypes.py +0 -0
  11. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/__init__.py +0 -0
  12. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
  13. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
  14. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
  15. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
  16. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/language/text.py +0 -0
  17. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
  18. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
  19. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
  20. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
  21. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
  22. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
  23. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
  24. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/__init__.py +0 -0
  25. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/encoding.py +0 -0
  26. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/engine/__init__.py +0 -0
  27. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/engine/base.py +0 -0
  28. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
  29. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
  30. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/generation.py +0 -0
  31. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/lstm.py +0 -0
  32. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
  33. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/training.py +0 -0
  34. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_language/xgrammar_utils.py +0 -0
  35. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_memory.py +0 -0
  36. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_tabular/__init__.py +0 -0
  37. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_tabular/argn.py +0 -0
  38. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_tabular/common.py +0 -0
  39. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_tabular/encoding.py +0 -0
  40. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_tabular/fairness.py +0 -0
  41. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_tabular/generation.py +0 -0
  42. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_tabular/training.py +0 -0
  43. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_training_utils.py +0 -0
  44. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/_workspace.py +0 -0
  45. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/domain.py +0 -0
  46. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/encoding.py +0 -0
  47. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/generation.py +0 -0
  48. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/logging.py +0 -0
  49. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/random_state.py +0 -0
  50. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/mostlyai/engine/splitting.py +0 -0
  51. {mostlyai_engine-1.5.2 → mostlyai_engine-1.5.4}/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.4
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.4"
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.*")
@@ -401,12 +401,11 @@ def get_sub_columns_lookup(
401
401
  return sub_cols_lookup
402
402
 
403
403
 
404
- class CtxSequenceLengthError(Exception):
405
- """Error raised when the cols of the same table do not have the same stats value"""
406
-
407
-
408
404
  def get_ctx_sequence_length(ctx_stats: dict, key: str) -> dict[str, int]:
409
- seq_stats: dict[str, int] = {}
405
+ """
406
+ Get the stats of sequence lengths from the first column_stats of each context table
407
+ """
408
+ ctxseq_stats: dict[str, int] = {}
410
409
 
411
410
  for column_stats in ctx_stats.get("columns", {}).values():
412
411
  if "seq_len" in column_stats:
@@ -414,12 +413,10 @@ def get_ctx_sequence_length(ctx_stats: dict, key: str) -> dict[str, int]:
414
413
  argn_processor=column_stats[ARGN_PROCESSOR],
415
414
  argn_table=column_stats[ARGN_TABLE],
416
415
  )
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]
416
+ if table not in ctxseq_stats:
417
+ ctxseq_stats[table] = column_stats["seq_len"][key]
421
418
 
422
- return seq_stats
419
+ return ctxseq_stats
423
420
 
424
421
 
425
422
  def get_max_data_points_per_sample(stats: dict) -> int:
@@ -867,6 +864,7 @@ def dp_non_rare(value_counts: dict[str, int], epsilon: float, threshold: int = 5
867
864
  noisy_counts = np.clip(np.array(list(value_counts.values())) + noise, 0, None).astype(int)
868
865
  for i, cat in enumerate(value_counts):
869
866
  value_counts[cat] = noisy_counts[i]
867
+ # NOTE: total_counts can be 0 in the edge case when the column only has null values
870
868
  total_counts = sum(value_counts.values())
871
869
 
872
870
  # 2. Collect all categories whose noisy count >= threshold
@@ -874,7 +872,7 @@ def dp_non_rare(value_counts: dict[str, int], epsilon: float, threshold: int = 5
874
872
 
875
873
  # 3. Compute the non-rare ratio
876
874
  noisy_total_counts = sum(selected.values())
877
- non_rare_ratio = noisy_total_counts / total_counts
875
+ non_rare_ratio = noisy_total_counts / total_counts if total_counts > 0 else 0
878
876
 
879
877
  return list(selected.keys()), non_rare_ratio
880
878
 
@@ -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
@@ -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.4"
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