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.
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/PKG-INFO +3 -3
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/__init__.py +1 -1
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_common.py +48 -46
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/common.py +5 -2
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/generation.py +10 -11
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/training.py +4 -1
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/xgrammar_utils.py +2 -2
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/argn.py +46 -11
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/encoding.py +67 -35
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/generation.py +241 -166
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/training.py +28 -23
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/pyproject.toml +3 -3
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/.gitignore +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/LICENSE +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/README.md +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_dtypes.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/__init__.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/__init__.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/categorical.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/datetime.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/numeric.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/language/text.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/__init__.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/categorical.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/character.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/datetime.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/itt.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/lat_long.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_encoding_types/tabular/numeric.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/__init__.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/encoding.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/__init__.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/base.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/hf_engine.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/engine/vllm_engine.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/lstm.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_language/tokenizer_utils.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_memory.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/__init__.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/common.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_tabular/fairness.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_training_utils.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/_workspace.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/analysis.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/domain.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/encoding.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/generation.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/logging.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/random_state.py +0 -0
- {mostlyai_engine-1.4.8 → mostlyai_engine-1.5.0}/mostlyai/engine/splitting.py +0 -0
- {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.
|
|
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
|
|
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
|
|
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.
|
|
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
|
-
|
|
52
|
-
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
|
|
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(
|
|
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 |=
|
|
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
|
|
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 <
|
|
518
|
-
# encode
|
|
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
|
|
529
|
-
if max_seq_len <
|
|
530
|
-
# decode
|
|
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
|
|
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
|
|
540
|
-
|
|
541
|
-
|
|
542
|
-
|
|
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
|
|
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
|
-
|
|
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
|
-
|
|
559
|
-
|
|
560
|
-
|
|
561
|
-
|
|
562
|
-
|
|
563
|
-
|
|
564
|
-
|
|
565
|
-
|
|
566
|
-
|
|
567
|
-
|
|
568
|
-
|
|
569
|
-
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
-
|
|
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=
|
|
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,
|
|
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
|
-
|
|
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(
|
|
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
|
|
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
|
|
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(
|
|
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(
|
|
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]
|
|
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([
|
|
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
|
|
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
|
-
|
|
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
|
-
|
|
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 =
|
|
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
|
-
#
|
|
140
|
-
df =
|
|
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 =
|
|
146
|
-
df_miss = pd.DataFrame(
|
|
147
|
-
|
|
148
|
-
|
|
149
|
-
df = pd.concat([df, df_miss],
|
|
150
|
-
#
|
|
151
|
-
df =
|
|
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 =
|
|
155
|
-
df_miss = pd.DataFrame(
|
|
156
|
-
|
|
157
|
-
|
|
158
|
-
df = pd.concat([df, df_miss],
|
|
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
|
|
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
|
-
|
|
388
|
-
|
|
389
|
-
sidx =
|
|
390
|
-
|
|
391
|
-
|
|
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
|
|
396
|
-
|
|
397
|
-
|
|
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
|
-
|
|
401
|
-
|
|
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
|
-
|
|
404
|
-
|
|
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
|
-
|
|
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
|