@tuned-tensor/local 0.2.9 → 0.3.0

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 (82) hide show
  1. package/CHANGELOG.md +59 -0
  2. package/README.md +84 -210
  3. package/dist/artifacts.d.ts +3 -4
  4. package/dist/artifacts.js +55 -33
  5. package/dist/artifacts.js.map +1 -1
  6. package/dist/compare.d.ts +1 -8
  7. package/dist/compare.js +1 -23
  8. package/dist/compare.js.map +1 -1
  9. package/dist/contracts.d.ts +59 -366
  10. package/dist/contracts.js +73 -224
  11. package/dist/contracts.js.map +1 -1
  12. package/dist/dataset.d.ts +2 -4
  13. package/dist/dataset.js +31 -125
  14. package/dist/dataset.js.map +1 -1
  15. package/dist/doctor.d.ts +2 -3
  16. package/dist/doctor.js +23 -161
  17. package/dist/doctor.js.map +1 -1
  18. package/dist/evaluation.d.ts +11 -71
  19. package/dist/evaluation.js +212 -572
  20. package/dist/evaluation.js.map +1 -1
  21. package/dist/huggingface-cache.d.ts +8 -7
  22. package/dist/huggingface-cache.js +16 -8
  23. package/dist/huggingface-cache.js.map +1 -1
  24. package/dist/index.d.ts +0 -10
  25. package/dist/index.js +203 -626
  26. package/dist/index.js.map +1 -1
  27. package/dist/local-project.d.ts +1 -11
  28. package/dist/local-project.js +33 -61
  29. package/dist/local-project.js.map +1 -1
  30. package/dist/model-registry.d.ts +3 -16
  31. package/dist/model-registry.js +76 -292
  32. package/dist/model-registry.js.map +1 -1
  33. package/dist/model-server.d.ts +1 -2
  34. package/dist/model-server.js +9 -18
  35. package/dist/model-server.js.map +1 -1
  36. package/dist/orchestrator.d.ts +11 -25
  37. package/dist/orchestrator.js +246 -566
  38. package/dist/orchestrator.js.map +1 -1
  39. package/dist/prefetch.d.ts +1 -5
  40. package/dist/prefetch.js +14 -19
  41. package/dist/prefetch.js.map +1 -1
  42. package/dist/process-runner.d.ts +10 -29
  43. package/dist/process-runner.js +63 -169
  44. package/dist/process-runner.js.map +1 -1
  45. package/dist/process-training.d.ts +0 -1
  46. package/dist/process-training.js +10 -87
  47. package/dist/process-training.js.map +1 -1
  48. package/dist/store.d.ts +1 -3
  49. package/dist/store.js +92 -290
  50. package/dist/store.js.map +1 -1
  51. package/docs/architecture.md +87 -152
  52. package/docs/spark.md +66 -97
  53. package/examples/dry-runner.json +14 -0
  54. package/examples/local-runner.json +8 -6
  55. package/examples/smoke-spec.json +41 -0
  56. package/package.json +6 -6
  57. package/training/local-runner/pyproject.toml +0 -5
  58. package/training/local-runner/src/evaluate.py +89 -234
  59. package/training/local-runner/src/model_contract.py +47 -0
  60. package/training/local-runner/src/prefetch.py +18 -4
  61. package/training/local-runner/src/serve.py +54 -115
  62. package/training/local-runner/src/sft_data.py +97 -0
  63. package/training/local-runner/src/train.py +146 -346
  64. package/training/local-runner/uv.lock +0 -1197
  65. package/dist/labeling-sanitize.d.ts +0 -31
  66. package/dist/labeling-sanitize.js +0 -158
  67. package/dist/labeling-sanitize.js.map +0 -1
  68. package/dist/labeling.d.ts +0 -155
  69. package/dist/labeling.js +0 -496
  70. package/dist/labeling.js.map +0 -1
  71. package/dist/openrouter.d.ts +0 -29
  72. package/dist/openrouter.js +0 -66
  73. package/dist/openrouter.js.map +0 -1
  74. package/dist/server.d.ts +0 -9
  75. package/dist/server.js +0 -211
  76. package/dist/server.js.map +0 -1
  77. package/docs/local-workflow-remediation-2026-07-13.md +0 -130
  78. package/docs/local-workflow-ux-review-2026-07-13.md +0 -458
  79. package/examples/dpo-preferences.jsonl +0 -2
  80. package/examples/dpo-run-request.json +0 -34
  81. package/examples/smoke-run-request.json +0 -34
  82. package/training/local-runner/src/train_dpo.py +0 -286
@@ -1,286 +0,0 @@
1
- from __future__ import annotations
2
-
3
- import inspect
4
- import json
5
- import os
6
- import tarfile
7
- import time
8
- from pathlib import Path
9
- from tempfile import TemporaryDirectory
10
- from typing import Any
11
-
12
- import torch
13
- from datasets import Dataset as HFDataset
14
- from peft import LoraConfig, PeftModel
15
- from transformers import AutoModelForCausalLM, AutoTokenizer
16
- from trl import DPOConfig, DPOTrainer
17
-
18
-
19
- TRAINING_DIR = Path(os.environ.get("SM_CHANNEL_TRAINING", "/opt/ml/input/data/training"))
20
- BASE_MODEL_DIR = Path(os.environ.get("SM_CHANNEL_BASE_MODEL", "/opt/ml/input/data/base_model"))
21
- HYPERPARAMETERS_PATH = Path(
22
- os.environ.get("TT_HYPERPARAMETERS_PATH", "/opt/ml/input/config/hyperparameters.json")
23
- )
24
- MODEL_DIR = Path(os.environ.get("SM_MODEL_DIR", "/opt/ml/model"))
25
- OUTPUT_DIR = Path(os.environ.get("SM_OUTPUT_DIR", "/opt/ml/output"))
26
- MAX_ARCHIVE_MEMBERS = 20_000
27
- MAX_ARCHIVE_EXPANDED_BYTES = 20 * 1024 * 1024 * 1024
28
-
29
-
30
- def load_hyperparameters() -> dict[str, str]:
31
- if not HYPERPARAMETERS_PATH.is_file():
32
- return {}
33
- raw = json.loads(HYPERPARAMETERS_PATH.read_text())
34
- return {str(key): str(value) for key, value in raw.items()}
35
-
36
-
37
- HP = load_hyperparameters()
38
-
39
-
40
- def hp(name: str, default: str | None = None) -> str | None:
41
- value = os.getenv(f"SM_HP_{name.upper()}", HP.get(name, default))
42
- if value is None:
43
- return None
44
- value = str(value).strip()
45
- if len(value) >= 2 and value[0] == value[-1] == '"':
46
- return value[1:-1]
47
- return value
48
-
49
-
50
- def hp_int(name: str, default: int) -> int:
51
- return int(hp(name, str(default)) or default)
52
-
53
-
54
- def hp_float(name: str, default: float) -> float:
55
- return float(hp(name, str(default)) or default)
56
-
57
-
58
- def hp_bool(name: str, default: bool) -> bool:
59
- return (hp(name, str(default)) or "").lower() in {"1", "true", "yes", "y"}
60
-
61
-
62
- def model_revision_kwargs(model_source: str) -> dict[str, str]:
63
- revision = hp("base_model_revision")
64
- return {"revision": revision} if revision and not Path(model_source).exists() else {}
65
-
66
-
67
- def supported_kwargs(callable_obj: Any, values: dict[str, Any]) -> dict[str, Any]:
68
- accepted = set(inspect.signature(callable_obj).parameters)
69
- return {key: value for key, value in values.items() if key in accepted}
70
-
71
-
72
- def load_preference_rows() -> list[dict[str, str]]:
73
- rows: list[dict[str, str]] = []
74
- for path in sorted(TRAINING_DIR.rglob("*.jsonl")):
75
- for line_number, line in enumerate(path.read_text().splitlines(), start=1):
76
- if not line.strip():
77
- continue
78
- row = json.loads(line)
79
- cleaned: dict[str, str] = {}
80
- for key in ("prompt", "chosen", "rejected"):
81
- value = row.get(key)
82
- if not isinstance(value, str) or not value.strip():
83
- raise ValueError(f"{path}:{line_number} missing non-empty string field {key}")
84
- cleaned[key] = value
85
- rows.append(cleaned)
86
- if not rows:
87
- raise ValueError(f"No preference JSONL rows found under {TRAINING_DIR}")
88
- return rows
89
-
90
-
91
- def safe_extract_archive(path: Path, destination: Path) -> None:
92
- destination.mkdir(parents=True, exist_ok=True)
93
- destination_root = destination.resolve()
94
- with tarfile.open(path, "r:gz") as tar:
95
- members = tar.getmembers()
96
- if len(members) > MAX_ARCHIVE_MEMBERS:
97
- raise ValueError(f"Model archive exceeds {MAX_ARCHIVE_MEMBERS} members")
98
- expanded_bytes = sum(max(0, member.size) for member in members if member.isfile())
99
- if expanded_bytes > MAX_ARCHIVE_EXPANDED_BYTES:
100
- raise ValueError("Model archive exceeds the 20 GiB expanded-size limit")
101
- for member in members:
102
- destination_path = (destination / member.name).resolve()
103
- try:
104
- destination_path.relative_to(destination_root)
105
- except ValueError:
106
- raise ValueError(f"Unsafe archive member: {member.name}")
107
- if member.issym() or member.islnk() or member.isdev():
108
- raise ValueError(f"Unsafe archive member type: {member.name}")
109
- tar.extractall(destination)
110
-
111
-
112
- def resolve_model_source(tmp: Path) -> str:
113
- archive = BASE_MODEL_DIR / "model.tar.gz"
114
- if archive.is_file():
115
- extracted = tmp / "base-model"
116
- safe_extract_archive(archive, extracted)
117
- candidates = [path for path in extracted.iterdir() if path.is_dir()]
118
- return str(candidates[0] if len(candidates) == 1 else extracted)
119
- if BASE_MODEL_DIR.is_dir() and any(BASE_MODEL_DIR.iterdir()):
120
- return str(BASE_MODEL_DIR)
121
- return hp("base_model") or "Qwen/Qwen3.5-2B"
122
-
123
-
124
- def strip_file_uri(value: str | None) -> str | None:
125
- if not value:
126
- return value
127
- if value.startswith("file://"):
128
- return value[7:]
129
- return value
130
-
131
-
132
- def resolve_adapter_path(value: str | None, tmp: Path) -> str | None:
133
- path_value = strip_file_uri(value)
134
- if not path_value:
135
- return None
136
- path = Path(path_value)
137
- if path.is_file() and path.name.endswith(".tar.gz"):
138
- extracted = tmp / "parent-adapter"
139
- safe_extract_archive(path, extracted)
140
- candidates = [candidate for candidate in extracted.rglob("adapter_config.json")]
141
- if candidates:
142
- return str(candidates[0].parent)
143
- directories = [candidate for candidate in extracted.iterdir() if candidate.is_dir()]
144
- return str(directories[0] if len(directories) == 1 else extracted)
145
- return str(path)
146
-
147
-
148
- def create_model_and_tokenizer(model_source: str):
149
- trust_remote_code = hp_bool("trust_remote_code", True)
150
- token = os.getenv("HF_TOKEN")
151
- tokenizer = AutoTokenizer.from_pretrained(
152
- model_source,
153
- **model_revision_kwargs(model_source),
154
- trust_remote_code=trust_remote_code,
155
- token=token,
156
- )
157
- if tokenizer.pad_token is None:
158
- tokenizer.pad_token = tokenizer.eos_token
159
-
160
- dtype = torch.bfloat16 if torch.cuda.is_available() and torch.cuda.is_bf16_supported() else torch.float16
161
- model = AutoModelForCausalLM.from_pretrained(
162
- model_source,
163
- **model_revision_kwargs(model_source),
164
- trust_remote_code=trust_remote_code,
165
- token=token,
166
- torch_dtype=dtype if torch.cuda.is_available() else None,
167
- device_map="auto" if torch.cuda.is_available() else None,
168
- )
169
- model.config.use_cache = False
170
- return model, tokenizer
171
-
172
-
173
- def lora_target_modules(default: str) -> list[str] | str:
174
- raw = hp("lora_target_modules", default) or default
175
- if raw == "all-linear":
176
- return raw
177
- return [item.strip() for item in raw.split(",") if item.strip()]
178
-
179
-
180
- def lora_config() -> LoraConfig:
181
- return LoraConfig(
182
- r=hp_int("lora_rank", 16),
183
- lora_alpha=hp_int("lora_alpha", 32),
184
- lora_dropout=hp_float("lora_dropout", 0.05),
185
- bias="none",
186
- task_type="CAUSAL_LM",
187
- target_modules=lora_target_modules("all-linear"),
188
- )
189
-
190
-
191
- def create_model_archive() -> Path:
192
- OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
193
- archive_path = OUTPUT_DIR / "model.tar.gz"
194
- with tarfile.open(archive_path, "w:gz") as tar:
195
- tar.add(MODEL_DIR, arcname="model")
196
- return archive_path
197
-
198
-
199
- def dpo_config(output_dir: str) -> DPOConfig:
200
- max_prompt_length = hp("max_prompt_length")
201
- max_completion_length = hp("max_completion_length")
202
- config_values: dict[str, Any] = {
203
- "output_dir": output_dir,
204
- "num_train_epochs": hp_int("n_epochs", 3),
205
- "learning_rate": hp_float("learning_rate", 0.00001),
206
- "per_device_train_batch_size": hp_int("per_device_train_batch_size", 1),
207
- "gradient_accumulation_steps": hp_int("gradient_accumulation_steps", 8),
208
- "logging_steps": 1,
209
- "save_strategy": "no",
210
- "report_to": "none",
211
- "remove_unused_columns": False,
212
- "bf16": torch.cuda.is_available() and torch.cuda.is_bf16_supported(),
213
- "fp16": torch.cuda.is_available() and not torch.cuda.is_bf16_supported(),
214
- "gradient_checkpointing": True,
215
- "gradient_checkpointing_kwargs": {"use_reentrant": False},
216
- "beta": hp_float("dpo_beta", 0.1),
217
- "loss_type": hp("dpo_loss_type", "sigmoid"),
218
- "label_smoothing": hp_float("dpo_label_smoothing", 0.0),
219
- "reference_free": hp_bool("dpo_reference_free", False),
220
- "max_length": hp_int("max_seq_length", 2048),
221
- "max_prompt_length": int(max_prompt_length) if max_prompt_length else None,
222
- "max_completion_length": int(max_completion_length) if max_completion_length else None,
223
- }
224
- return DPOConfig(**supported_kwargs(
225
- DPOConfig,
226
- {key: value for key, value in config_values.items() if value is not None},
227
- ))
228
-
229
-
230
- def run_training(rows: list[dict[str, str]], model_source: str) -> dict[str, Any]:
231
- with TemporaryDirectory() as tmp:
232
- parent_adapter = resolve_adapter_path(hp("parent_model_artifact"), Path(tmp))
233
- model, tokenizer = create_model_and_tokenizer(model_source)
234
- if parent_adapter:
235
- model = PeftModel.from_pretrained(model, parent_adapter, is_trainable=True)
236
- dataset = HFDataset.from_list(rows)
237
-
238
- args = dpo_config(tmp)
239
- trainer_values: dict[str, Any] = {
240
- "model": model,
241
- "ref_model": None,
242
- "args": args,
243
- "train_dataset": dataset,
244
- "processing_class": tokenizer,
245
- "tokenizer": tokenizer,
246
- }
247
- if not parent_adapter:
248
- trainer_values["peft_config"] = lora_config()
249
- trainer = DPOTrainer(**supported_kwargs(DPOTrainer, trainer_values))
250
- result = trainer.train()
251
- trainer.save_model(MODEL_DIR)
252
-
253
- tokenizer.save_pretrained(MODEL_DIR)
254
- return result.metrics
255
-
256
-
257
- def main() -> None:
258
- started = time.time()
259
- MODEL_DIR.mkdir(parents=True, exist_ok=True)
260
- OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
261
-
262
- rows = load_preference_rows()
263
- with TemporaryDirectory(prefix="tt-local-base-model-") as base_tmp:
264
- model_source = resolve_model_source(Path(base_tmp))
265
- metrics = run_training(rows, model_source)
266
-
267
- archive_path = create_model_archive()
268
- output_metrics = {
269
- "training_method": "dpo",
270
- "preference_rows": len(rows),
271
- "model_source": str(BASE_MODEL_DIR) if BASE_MODEL_DIR.is_dir() and any(BASE_MODEL_DIR.iterdir()) else hp("base_model"),
272
- "base_model_revision": hp("base_model_revision"),
273
- "parent_model_artifact": hp("parent_model_artifact"),
274
- "dpo_beta": hp_float("dpo_beta", 0.1),
275
- "dpo_loss_type": hp("dpo_loss_type", "sigmoid"),
276
- "train_runtime": round(time.time() - started, 3),
277
- **{key: float(value) for key, value in metrics.items() if isinstance(value, (int, float))},
278
- "model_archive": str(archive_path),
279
- }
280
- (MODEL_DIR / "training-metrics.json").write_text(json.dumps(output_metrics, indent=2))
281
- (OUTPUT_DIR / "training-metrics.json").write_text(json.dumps(output_metrics, indent=2))
282
- print(json.dumps(output_metrics, indent=2))
283
-
284
-
285
- if __name__ == "__main__":
286
- main()