@tuned-tensor/local 0.2.7 → 0.2.8

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 (60) hide show
  1. package/CHANGELOG.md +51 -0
  2. package/README.md +139 -0
  3. package/dist/artifacts.d.ts +92 -0
  4. package/dist/artifacts.js +591 -3
  5. package/dist/artifacts.js.map +1 -1
  6. package/dist/contracts.d.ts +5 -2
  7. package/dist/contracts.js +15 -10
  8. package/dist/contracts.js.map +1 -1
  9. package/dist/dataset.d.ts +3 -0
  10. package/dist/dataset.js +87 -5
  11. package/dist/dataset.js.map +1 -1
  12. package/dist/doctor.d.ts +13 -2
  13. package/dist/doctor.js +372 -49
  14. package/dist/doctor.js.map +1 -1
  15. package/dist/evaluation.d.ts +15 -0
  16. package/dist/evaluation.js +120 -16
  17. package/dist/evaluation.js.map +1 -1
  18. package/dist/huggingface-cache.d.ts +26 -0
  19. package/dist/huggingface-cache.js +68 -0
  20. package/dist/huggingface-cache.js.map +1 -0
  21. package/dist/index.d.ts +3 -0
  22. package/dist/index.js +741 -75
  23. package/dist/index.js.map +1 -1
  24. package/dist/local-project.d.ts +11 -0
  25. package/dist/local-project.js +96 -3
  26. package/dist/local-project.js.map +1 -1
  27. package/dist/model-registry.d.ts +21 -0
  28. package/dist/model-registry.js +198 -0
  29. package/dist/model-registry.js.map +1 -1
  30. package/dist/model-server.d.ts +32 -0
  31. package/dist/model-server.js +158 -0
  32. package/dist/model-server.js.map +1 -0
  33. package/dist/orchestrator.d.ts +3 -0
  34. package/dist/orchestrator.js +1001 -142
  35. package/dist/orchestrator.js.map +1 -1
  36. package/dist/prefetch.d.ts +13 -0
  37. package/dist/prefetch.js +107 -19
  38. package/dist/prefetch.js.map +1 -1
  39. package/dist/process-runner.d.ts +7 -0
  40. package/dist/process-runner.js +171 -16
  41. package/dist/process-runner.js.map +1 -1
  42. package/dist/process-training.d.ts +5 -1
  43. package/dist/process-training.js +32 -16
  44. package/dist/process-training.js.map +1 -1
  45. package/dist/run-reporter.d.ts +2 -0
  46. package/dist/run-reporter.js +10 -0
  47. package/dist/run-reporter.js.map +1 -1
  48. package/dist/store.d.ts +16 -2
  49. package/dist/store.js +189 -24
  50. package/dist/store.js.map +1 -1
  51. package/docs/architecture.md +32 -0
  52. package/docs/local-workflow-remediation-2026-07-13.md +130 -0
  53. package/docs/local-workflow-ux-review-2026-07-13.md +458 -0
  54. package/package.json +3 -1
  55. package/training/local-runner/src/evaluate.py +80 -20
  56. package/training/local-runner/src/prefetch.py +107 -8
  57. package/training/local-runner/src/serve.py +329 -0
  58. package/training/local-runner/src/train.py +47 -22
  59. package/training/local-runner/src/train_dpo.py +41 -25
  60. package/training/local-runner/uv.lock +2401 -0
@@ -23,6 +23,8 @@ HYPERPARAMETERS_PATH = Path(
23
23
  )
24
24
  MODEL_DIR = Path(os.environ.get("SM_MODEL_DIR", "/opt/ml/model"))
25
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
26
28
 
27
29
 
28
30
  def load_hyperparameters() -> dict[str, str]:
@@ -57,6 +59,11 @@ def hp_bool(name: str, default: bool) -> bool:
57
59
  return (hp(name, str(default)) or "").lower() in {"1", "true", "yes", "y"}
58
60
 
59
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
+
60
67
  def supported_kwargs(callable_obj: Any, values: dict[str, Any]) -> dict[str, Any]:
61
68
  accepted = set(inspect.signature(callable_obj).parameters)
62
69
  return {key: value for key, value in values.items() if key in accepted}
@@ -81,22 +88,36 @@ def load_preference_rows() -> list[dict[str, str]]:
81
88
  return rows
82
89
 
83
90
 
84
- def resolve_model_source() -> str:
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:
85
113
  archive = BASE_MODEL_DIR / "model.tar.gz"
86
114
  if archive.is_file():
87
- extracted = Path("/tmp/base-model")
88
- extracted.mkdir(parents=True, exist_ok=True)
89
- with tarfile.open(archive, "r:gz") as tar:
90
- extracted_root = extracted.resolve()
91
- for member in tar.getmembers():
92
- destination = (extracted / member.name).resolve()
93
- try:
94
- destination.relative_to(extracted_root)
95
- except ValueError:
96
- raise ValueError(f"Unsafe archive member: {member.name}")
97
- tar.extractall(extracted)
115
+ extracted = tmp / "base-model"
116
+ safe_extract_archive(archive, extracted)
98
117
  candidates = [path for path in extracted.iterdir() if path.is_dir()]
99
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)
100
121
  return hp("base_model") or "Qwen/Qwen3.5-2B"
101
122
 
102
123
 
@@ -115,16 +136,7 @@ def resolve_adapter_path(value: str | None, tmp: Path) -> str | None:
115
136
  path = Path(path_value)
116
137
  if path.is_file() and path.name.endswith(".tar.gz"):
117
138
  extracted = tmp / "parent-adapter"
118
- extracted.mkdir(parents=True, exist_ok=True)
119
- with tarfile.open(path, "r:gz") as tar:
120
- extracted_root = extracted.resolve()
121
- for member in tar.getmembers():
122
- destination = (extracted / member.name).resolve()
123
- try:
124
- destination.relative_to(extracted_root)
125
- except ValueError:
126
- raise ValueError(f"Unsafe archive member: {member.name}")
127
- tar.extractall(extracted)
139
+ safe_extract_archive(path, extracted)
128
140
  candidates = [candidate for candidate in extracted.rglob("adapter_config.json")]
129
141
  if candidates:
130
142
  return str(candidates[0].parent)
@@ -138,6 +150,7 @@ def create_model_and_tokenizer(model_source: str):
138
150
  token = os.getenv("HF_TOKEN")
139
151
  tokenizer = AutoTokenizer.from_pretrained(
140
152
  model_source,
153
+ **model_revision_kwargs(model_source),
141
154
  trust_remote_code=trust_remote_code,
142
155
  token=token,
143
156
  )
@@ -147,6 +160,7 @@ def create_model_and_tokenizer(model_source: str):
147
160
  dtype = torch.bfloat16 if torch.cuda.is_available() and torch.cuda.is_bf16_supported() else torch.float16
148
161
  model = AutoModelForCausalLM.from_pretrained(
149
162
  model_source,
163
+ **model_revision_kwargs(model_source),
150
164
  trust_remote_code=trust_remote_code,
151
165
  token=token,
152
166
  torch_dtype=dtype if torch.cuda.is_available() else None,
@@ -246,14 +260,16 @@ def main() -> None:
246
260
  OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
247
261
 
248
262
  rows = load_preference_rows()
249
- model_source = resolve_model_source()
250
- metrics = run_training(rows, model_source)
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)
251
266
 
252
267
  archive_path = create_model_archive()
253
268
  output_metrics = {
254
269
  "training_method": "dpo",
255
270
  "preference_rows": len(rows),
256
- "model_source": model_source,
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"),
257
273
  "parent_model_artifact": hp("parent_model_artifact"),
258
274
  "dpo_beta": hp_float("dpo_beta", 0.1),
259
275
  "dpo_loss_type": hp("dpo_loss_type", "sigmoid"),