@tuned-tensor/local 0.2.6 → 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.
- package/CHANGELOG.md +60 -0
- package/README.md +149 -0
- package/dist/artifacts.d.ts +92 -0
- package/dist/artifacts.js +591 -3
- package/dist/artifacts.js.map +1 -1
- package/dist/contracts.d.ts +5 -2
- package/dist/contracts.js +15 -10
- package/dist/contracts.js.map +1 -1
- package/dist/dataset.d.ts +3 -0
- package/dist/dataset.js +87 -5
- package/dist/dataset.js.map +1 -1
- package/dist/doctor.d.ts +13 -2
- package/dist/doctor.js +372 -49
- package/dist/doctor.js.map +1 -1
- package/dist/evaluation.d.ts +15 -0
- package/dist/evaluation.js +120 -16
- package/dist/evaluation.js.map +1 -1
- package/dist/huggingface-cache.d.ts +26 -0
- package/dist/huggingface-cache.js +68 -0
- package/dist/huggingface-cache.js.map +1 -0
- package/dist/index.d.ts +4 -0
- package/dist/index.js +766 -73
- package/dist/index.js.map +1 -1
- package/dist/local-project.d.ts +11 -0
- package/dist/local-project.js +96 -3
- package/dist/local-project.js.map +1 -1
- package/dist/model-registry.d.ts +21 -0
- package/dist/model-registry.js +198 -0
- package/dist/model-registry.js.map +1 -1
- package/dist/model-server.d.ts +32 -0
- package/dist/model-server.js +158 -0
- package/dist/model-server.js.map +1 -0
- package/dist/orchestrator.d.ts +3 -0
- package/dist/orchestrator.js +1001 -142
- package/dist/orchestrator.js.map +1 -1
- package/dist/prefetch.d.ts +43 -0
- package/dist/prefetch.js +192 -0
- package/dist/prefetch.js.map +1 -0
- package/dist/process-runner.d.ts +7 -0
- package/dist/process-runner.js +171 -16
- package/dist/process-runner.js.map +1 -1
- package/dist/process-training.d.ts +5 -1
- package/dist/process-training.js +32 -16
- package/dist/process-training.js.map +1 -1
- package/dist/run-reporter.d.ts +2 -0
- package/dist/run-reporter.js +10 -0
- package/dist/run-reporter.js.map +1 -1
- package/dist/store.d.ts +16 -2
- package/dist/store.js +189 -24
- package/dist/store.js.map +1 -1
- package/docs/architecture.md +32 -0
- package/docs/local-workflow-remediation-2026-07-13.md +130 -0
- package/docs/local-workflow-ux-review-2026-07-13.md +458 -0
- package/docs/spark.md +10 -0
- package/package.json +3 -1
- package/training/local-runner/pyproject.toml +1 -0
- package/training/local-runner/src/evaluate.py +80 -20
- package/training/local-runner/src/prefetch.py +178 -0
- package/training/local-runner/src/serve.py +329 -0
- package/training/local-runner/src/train.py +47 -22
- package/training/local-runner/src/train_dpo.py +41 -25
- 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
|
|
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 =
|
|
88
|
-
|
|
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
|
-
|
|
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
|
-
|
|
250
|
-
|
|
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":
|
|
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"),
|