Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -227,6 +227,53 @@ routing-tier-drift gates pass. This is the current accuracy-qualified shipping c
|
|||||||
latency and energy/residency still require measurement on the target Apple and Windows
|
latency and energy/residency still require measurement on the target Apple and Windows
|
||||||
accelerator runtimes.
|
accelerator runtimes.
|
||||||
|
|
||||||
|
### First-prompt history augmentation experiment
|
||||||
|
|
||||||
|
`prepare_history_experiment.py` appends the labeled Nucleic first-prompt corpus to
|
||||||
|
training only. It preserves validation and test byte-for-byte, excludes `vague-eval`
|
||||||
|
records from optimization, removes exact base/evaluation overlap, and applies the
|
||||||
|
canonical 0.92 near-duplicate guard against evaluation fixtures and earlier history
|
||||||
|
records. Every exclusion is represented only by hashes and source line in
|
||||||
|
`history-exclusions.jsonl`; `manifest.json` binds all input and output hashes.
|
||||||
|
|
||||||
|
Build the augmented split and its teacher cache:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ml/purpose-classifier/venv/bin/python \
|
||||||
|
ml/purpose-classifier/prepare_history_experiment.py
|
||||||
|
ml/purpose-classifier/venv/bin/python -u \
|
||||||
|
ml/purpose-classifier/cache_teacher.py \
|
||||||
|
--dataset-dir \
|
||||||
|
ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \
|
||||||
|
--model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \
|
||||||
|
--output \
|
||||||
|
ml/purpose-classifier/outputs/purpose-lite-v1-history-first-prompts-teacher.pt \
|
||||||
|
--device mps --batch-size 16 --progress-steps 25
|
||||||
|
```
|
||||||
|
|
||||||
|
Then run the same validation-selected MLX recipe as the accepted baseline:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
ml/purpose-classifier/venv/bin/python -u ml/purpose-classifier/train_mlx.py \
|
||||||
|
--device metal \
|
||||||
|
--dataset-dir \
|
||||||
|
ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \
|
||||||
|
--model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \
|
||||||
|
--distillation-cache \
|
||||||
|
ml/purpose-classifier/outputs/purpose-lite-v1-history-first-prompts-teacher.pt \
|
||||||
|
--distillation-weight 0.9 --distillation-temperature 2 \
|
||||||
|
--distillation-selection-weight 0.5 --quantization-aware \
|
||||||
|
--epochs 4 --early-stopping-patience 1 \
|
||||||
|
--learning-rate 1e-6 --warmup-ratio 0 --boundary-weight 1 \
|
||||||
|
--progress-steps 1 \
|
||||||
|
--output-dir \
|
||||||
|
ml/purpose-classifier/outputs/purpose-lite-v1-distilled-qat-mlx-history-v1
|
||||||
|
```
|
||||||
|
|
||||||
|
MLX remains Metal-first. `--device cpu` is an explicit diagnostic fallback for parity
|
||||||
|
checks and bounded smoke tests; it is not an acceptable full-training path when Metal is
|
||||||
|
available.
|
||||||
|
|
||||||
### Convert and validate Core ML
|
### Convert and validate Core ML
|
||||||
|
|
||||||
Core ML Tools no longer maintains the legacy ONNX converter, so the Apple artifact is
|
Core ML Tools no longer maintains the legacy ONNX converter, so the Apple artifact is
|
||||||
|
|||||||
@@ -0,0 +1,323 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Build a training-only dataset augmented with labeled first prompts."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import shutil
|
||||||
|
import sys
|
||||||
|
from collections import Counter
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Sequence
|
||||||
|
|
||||||
|
from purpose_data import (
|
||||||
|
DataError,
|
||||||
|
SourceRecord,
|
||||||
|
curate_records,
|
||||||
|
distribution,
|
||||||
|
file_sha256,
|
||||||
|
jsonl_bytes,
|
||||||
|
load_classifiable_fixtures,
|
||||||
|
load_jsonl,
|
||||||
|
normalized_key,
|
||||||
|
prompt_hash,
|
||||||
|
validate_source_record,
|
||||||
|
write_json,
|
||||||
|
write_jsonl,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||||
|
REPOSITORY_ROOT = SCRIPT_DIR.parent.parent
|
||||||
|
DEFAULT_BASE_DATASET = SCRIPT_DIR / ".artifacts" / "dataset-v1"
|
||||||
|
DEFAULT_HISTORY = (
|
||||||
|
SCRIPT_DIR / ".artifacts" / "nucleic-history-first-prompts.labeled.jsonl"
|
||||||
|
)
|
||||||
|
DEFAULT_FIXTURES = (
|
||||||
|
REPOSITORY_ROOT
|
||||||
|
/ "Tests"
|
||||||
|
/ "NucleicCoreTests"
|
||||||
|
/ "Fixtures"
|
||||||
|
/ "purpose-prompts.json"
|
||||||
|
)
|
||||||
|
DEFAULT_OUTPUT = SCRIPT_DIR / ".artifacts" / "dataset-v1-history-first-prompts"
|
||||||
|
EXPERIMENT_VERSION = "history-first-prompts-training-augmentation-v1"
|
||||||
|
|
||||||
|
|
||||||
|
def _relative(path: Path) -> str:
|
||||||
|
try:
|
||||||
|
return str(path.resolve().relative_to(REPOSITORY_ROOT))
|
||||||
|
except ValueError:
|
||||||
|
return str(path.resolve())
|
||||||
|
|
||||||
|
|
||||||
|
def _validate_records(records: Sequence[dict[str, Any]], path: Path) -> None:
|
||||||
|
if not records:
|
||||||
|
raise DataError(f"{path}: split is empty")
|
||||||
|
for line, record in enumerate(records, 1):
|
||||||
|
validate_source_record(record, f"{path}:{line}")
|
||||||
|
|
||||||
|
|
||||||
|
def _indexed_records(
|
||||||
|
records: Sequence[dict[str, Any]], path: Path
|
||||||
|
) -> list[SourceRecord]:
|
||||||
|
return [
|
||||||
|
SourceRecord(value=record, source=path, line=line)
|
||||||
|
for line, record in enumerate(records, 1)
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _unique_split_keys(
|
||||||
|
splits: Sequence[tuple[str, Sequence[dict[str, Any]]]]
|
||||||
|
) -> dict[str, tuple[str, str]]:
|
||||||
|
seen: dict[str, tuple[str, str]] = {}
|
||||||
|
for split_name, records in splits:
|
||||||
|
for line, record in enumerate(records, 1):
|
||||||
|
key = normalized_key(record["prompt"])
|
||||||
|
previous = seen.get(key)
|
||||||
|
if previous is not None:
|
||||||
|
previous_split, previous_label = previous
|
||||||
|
detail = (
|
||||||
|
"conflicting labels"
|
||||||
|
if previous_label != record["purpose"]
|
||||||
|
else "duplicate prompt"
|
||||||
|
)
|
||||||
|
raise DataError(
|
||||||
|
f"{split_name}:{line}: {detail} also present in {previous_split}"
|
||||||
|
)
|
||||||
|
seen[key] = (split_name, record["purpose"])
|
||||||
|
return seen
|
||||||
|
|
||||||
|
|
||||||
|
def _sha256_bytes(value: bytes) -> str:
|
||||||
|
return hashlib.sha256(value).hexdigest()
|
||||||
|
|
||||||
|
|
||||||
|
def _exclusion(
|
||||||
|
record: SourceRecord,
|
||||||
|
*,
|
||||||
|
reason: str,
|
||||||
|
matched_prompt_hash: str | None = None,
|
||||||
|
similarity: float | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
value: dict[str, Any] = {
|
||||||
|
"promptHash": prompt_hash(record.value["prompt"]),
|
||||||
|
"sourceLine": record.line,
|
||||||
|
"reason": reason,
|
||||||
|
}
|
||||||
|
if matched_prompt_hash is not None:
|
||||||
|
value["matchedPromptHash"] = matched_prompt_hash
|
||||||
|
if similarity is not None:
|
||||||
|
value["similarity"] = round(similarity, 6)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def prepare(
|
||||||
|
*,
|
||||||
|
base_dataset: Path,
|
||||||
|
history_path: Path,
|
||||||
|
fixtures_path: Path,
|
||||||
|
output_dir: Path,
|
||||||
|
near_duplicate_threshold: float,
|
||||||
|
overwrite_output: bool,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
base_paths = {
|
||||||
|
split: base_dataset / f"{split}.jsonl"
|
||||||
|
for split in ("train", "validation", "test")
|
||||||
|
}
|
||||||
|
base = {split: load_jsonl(path) for split, path in base_paths.items()}
|
||||||
|
for split, path in base_paths.items():
|
||||||
|
_validate_records(base[split], path)
|
||||||
|
|
||||||
|
base_keys = _unique_split_keys(
|
||||||
|
[(split, base[split]) for split in ("train", "validation", "test")]
|
||||||
|
)
|
||||||
|
history = load_jsonl(history_path)
|
||||||
|
_validate_records(history, history_path)
|
||||||
|
fixtures = load_classifiable_fixtures(fixtures_path)
|
||||||
|
|
||||||
|
eval_fixtures = [
|
||||||
|
{"prompt": record["prompt"], "purpose": record["purpose"]}
|
||||||
|
for split in ("validation", "test")
|
||||||
|
for record in base[split]
|
||||||
|
] + fixtures
|
||||||
|
eval_hashes = {prompt_hash(record["prompt"]) for record in eval_fixtures}
|
||||||
|
|
||||||
|
eligible: list[SourceRecord] = []
|
||||||
|
exclusions: list[dict[str, Any]] = []
|
||||||
|
for record in _indexed_records(history, history_path):
|
||||||
|
if record.value["slice"] == "vague-eval":
|
||||||
|
exclusions.append(_exclusion(record, reason="vague-eval"))
|
||||||
|
continue
|
||||||
|
|
||||||
|
key = normalized_key(record.value["prompt"])
|
||||||
|
previous = base_keys.get(key)
|
||||||
|
if previous is not None:
|
||||||
|
split_name, previous_label = previous
|
||||||
|
if previous_label != record.value["purpose"]:
|
||||||
|
raise DataError(
|
||||||
|
f"{history_path}:{record.line}: label conflicts with exact "
|
||||||
|
f"{split_name} prompt"
|
||||||
|
)
|
||||||
|
exclusions.append(
|
||||||
|
_exclusion(
|
||||||
|
record,
|
||||||
|
reason=(
|
||||||
|
"exact-base-train-overlap"
|
||||||
|
if split_name == "train"
|
||||||
|
else "exact-eval-overlap"
|
||||||
|
),
|
||||||
|
matched_prompt_hash=prompt_hash(record.value["prompt"]),
|
||||||
|
similarity=1.0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
eligible.append(record)
|
||||||
|
|
||||||
|
curated = curate_records(
|
||||||
|
eligible,
|
||||||
|
eval_fixtures,
|
||||||
|
near_duplicate_threshold=near_duplicate_threshold,
|
||||||
|
)
|
||||||
|
for duplicate in curated.duplicates:
|
||||||
|
exclusions.append(
|
||||||
|
_exclusion(
|
||||||
|
duplicate.dropped,
|
||||||
|
reason=(
|
||||||
|
f"{duplicate.kind}-eval-overlap"
|
||||||
|
if duplicate.matched_prompt_hash in eval_hashes
|
||||||
|
else f"{duplicate.kind}-history-duplicate"
|
||||||
|
),
|
||||||
|
matched_prompt_hash=duplicate.matched_prompt_hash,
|
||||||
|
similarity=duplicate.similarity,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
accepted_history = [record.value for record in curated.records]
|
||||||
|
output_train = base["train"] + accepted_history
|
||||||
|
train_bytes = jsonl_bytes(output_train)
|
||||||
|
validation_bytes = base_paths["validation"].read_bytes()
|
||||||
|
test_bytes = base_paths["test"].read_bytes()
|
||||||
|
|
||||||
|
if output_dir.exists() and any(output_dir.iterdir()):
|
||||||
|
if not overwrite_output:
|
||||||
|
raise DataError(
|
||||||
|
f"{output_dir}: output is not empty; pass --overwrite-output "
|
||||||
|
"intentionally"
|
||||||
|
)
|
||||||
|
shutil.rmtree(output_dir)
|
||||||
|
output_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
(output_dir / "train.jsonl").write_bytes(train_bytes)
|
||||||
|
(output_dir / "validation.jsonl").write_bytes(validation_bytes)
|
||||||
|
(output_dir / "test.jsonl").write_bytes(test_bytes)
|
||||||
|
write_jsonl(
|
||||||
|
output_dir / "history-exclusions.jsonl",
|
||||||
|
sorted(exclusions, key=lambda value: value["sourceLine"]),
|
||||||
|
)
|
||||||
|
|
||||||
|
exclusion_counts = dict(
|
||||||
|
sorted(Counter(value["reason"] for value in exclusions).items())
|
||||||
|
)
|
||||||
|
accepted_sources = _indexed_records(accepted_history, history_path)
|
||||||
|
manifest = {
|
||||||
|
"schemaVersion": 1,
|
||||||
|
"experimentVersion": EXPERIMENT_VERSION,
|
||||||
|
"policy": {
|
||||||
|
"historyUsage": "training-only",
|
||||||
|
"vagueEval": "excluded from optimization",
|
||||||
|
"exactBaseTrainOverlap": "excluded",
|
||||||
|
"exactEvaluationOverlap": "excluded",
|
||||||
|
"nearEvaluationOverlap": "excluded",
|
||||||
|
"nearHistoryDuplicates": "excluded",
|
||||||
|
"nearDuplicateThreshold": near_duplicate_threshold,
|
||||||
|
"validationAndTestBytes": "identical to base dataset",
|
||||||
|
},
|
||||||
|
"sources": {
|
||||||
|
"baseDataset": {
|
||||||
|
"path": _relative(base_dataset),
|
||||||
|
"splits": {
|
||||||
|
split: {
|
||||||
|
"path": _relative(path),
|
||||||
|
"records": len(base[split]),
|
||||||
|
"sha256": file_sha256(path),
|
||||||
|
}
|
||||||
|
for split, path in base_paths.items()
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"history": {
|
||||||
|
"path": _relative(history_path),
|
||||||
|
"records": len(history),
|
||||||
|
"sha256": file_sha256(history_path),
|
||||||
|
},
|
||||||
|
"fixtures": {
|
||||||
|
"path": _relative(fixtures_path),
|
||||||
|
"classifiableRecords": len(fixtures),
|
||||||
|
"sha256": file_sha256(fixtures_path),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"augmentation": {
|
||||||
|
"inputHistoryRecords": len(history),
|
||||||
|
"acceptedHistoryRecords": len(accepted_history),
|
||||||
|
"excludedHistoryRecords": len(exclusions),
|
||||||
|
"exclusions": exclusion_counts,
|
||||||
|
"acceptedDistribution": distribution(accepted_sources),
|
||||||
|
},
|
||||||
|
"outputs": {
|
||||||
|
"train": {
|
||||||
|
"records": len(output_train),
|
||||||
|
"sha256": _sha256_bytes(train_bytes),
|
||||||
|
},
|
||||||
|
"validation": {
|
||||||
|
"records": len(base["validation"]),
|
||||||
|
"sha256": _sha256_bytes(validation_bytes),
|
||||||
|
},
|
||||||
|
"test": {
|
||||||
|
"records": len(base["test"]),
|
||||||
|
"sha256": _sha256_bytes(test_bytes),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
write_json(output_dir / "manifest.json", manifest)
|
||||||
|
return manifest
|
||||||
|
|
||||||
|
|
||||||
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
|
parser = argparse.ArgumentParser(description=__doc__)
|
||||||
|
parser.add_argument("--base-dataset", type=Path, default=DEFAULT_BASE_DATASET)
|
||||||
|
parser.add_argument("--history", type=Path, default=DEFAULT_HISTORY)
|
||||||
|
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)
|
||||||
|
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
|
||||||
|
parser.add_argument("--near-duplicate-threshold", type=float, default=0.92)
|
||||||
|
parser.add_argument("--overwrite-output", action="store_true")
|
||||||
|
return parser
|
||||||
|
|
||||||
|
|
||||||
|
def main(argv: Sequence[str] | None = None) -> int:
|
||||||
|
args = build_parser().parse_args(argv)
|
||||||
|
try:
|
||||||
|
manifest = prepare(
|
||||||
|
base_dataset=args.base_dataset.resolve(),
|
||||||
|
history_path=args.history.resolve(),
|
||||||
|
fixtures_path=args.fixtures.resolve(),
|
||||||
|
output_dir=args.output_dir.resolve(),
|
||||||
|
near_duplicate_threshold=args.near_duplicate_threshold,
|
||||||
|
overwrite_output=args.overwrite_output,
|
||||||
|
)
|
||||||
|
except (DataError, OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||||
|
print(f"error: {exc}", file=sys.stderr)
|
||||||
|
return 1
|
||||||
|
|
||||||
|
augmentation = manifest["augmentation"]
|
||||||
|
outputs = manifest["outputs"]
|
||||||
|
print(
|
||||||
|
f"Prepared {outputs['train']['records']} training records "
|
||||||
|
f"({augmentation['acceptedHistoryRecords']} from history); validation and "
|
||||||
|
f"test remain byte-identical to the base dataset."
|
||||||
|
)
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
raise SystemExit(main())
|
||||||
@@ -1,2 +1,3 @@
|
|||||||
-r requirements.txt
|
-r requirements.txt
|
||||||
mlx==0.32.0
|
mlx==0.32.0
|
||||||
|
mlx-cpu==0.32.0; sys_platform == "linux"
|
||||||
|
|||||||
@@ -0,0 +1,133 @@
|
|||||||
|
import json
|
||||||
|
import sys
|
||||||
|
import tempfile
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
MODULE_DIR = Path(__file__).resolve().parents[1]
|
||||||
|
sys.path.insert(0, str(MODULE_DIR))
|
||||||
|
|
||||||
|
import prepare_history_experiment
|
||||||
|
import purpose_data
|
||||||
|
|
||||||
|
|
||||||
|
def example(index: int, **overrides):
|
||||||
|
value = {
|
||||||
|
"prompt": f"Implement sample endpoint number {index} with stable pagination",
|
||||||
|
"purpose": purpose_data.LABELS[index % len(purpose_data.LABELS)],
|
||||||
|
"secondary": None,
|
||||||
|
"mixed": False,
|
||||||
|
"difficulty": 0.4,
|
||||||
|
"slice": "core",
|
||||||
|
"lang": "en",
|
||||||
|
}
|
||||||
|
value.update(overrides)
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
class PrepareHistoryExperimentTests(unittest.TestCase):
|
||||||
|
def test_history_is_training_only_and_eval_splits_are_byte_identical(self):
|
||||||
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
root = Path(directory)
|
||||||
|
base = root / "base"
|
||||||
|
base.mkdir()
|
||||||
|
train = [example(0)]
|
||||||
|
validation = [example(1)]
|
||||||
|
test = [example(2)]
|
||||||
|
purpose_data.write_jsonl(base / "train.jsonl", train)
|
||||||
|
purpose_data.write_jsonl(base / "validation.jsonl", validation)
|
||||||
|
purpose_data.write_jsonl(base / "test.jsonl", test)
|
||||||
|
validation_bytes = (base / "validation.jsonl").read_bytes()
|
||||||
|
test_bytes = (base / "test.jsonl").read_bytes()
|
||||||
|
|
||||||
|
history_path = root / "history.jsonl"
|
||||||
|
history = [
|
||||||
|
example(10, prompt="Add a durable upload endpoint", purpose="backendImpl"),
|
||||||
|
example(11, prompt=validation[0]["prompt"], purpose=validation[0]["purpose"]),
|
||||||
|
example(12, prompt="Can you assess this?", slice="vague-eval"),
|
||||||
|
]
|
||||||
|
purpose_data.write_jsonl(history_path, history)
|
||||||
|
fixtures = root / "fixtures.json"
|
||||||
|
fixtures.write_text(
|
||||||
|
json.dumps(
|
||||||
|
[{"prompt": "Review the release diff", "purpose": "review"}]
|
||||||
|
),
|
||||||
|
encoding="utf-8",
|
||||||
|
)
|
||||||
|
output = root / "output"
|
||||||
|
|
||||||
|
manifest = prepare_history_experiment.prepare(
|
||||||
|
base_dataset=base,
|
||||||
|
history_path=history_path,
|
||||||
|
fixtures_path=fixtures,
|
||||||
|
output_dir=output,
|
||||||
|
near_duplicate_threshold=0.92,
|
||||||
|
overwrite_output=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(2, manifest["outputs"]["train"]["records"])
|
||||||
|
self.assertEqual(1, manifest["augmentation"]["acceptedHistoryRecords"])
|
||||||
|
self.assertEqual(
|
||||||
|
{"exact-eval-overlap": 1, "vague-eval": 1},
|
||||||
|
manifest["augmentation"]["exclusions"],
|
||||||
|
)
|
||||||
|
self.assertEqual(validation_bytes, (output / "validation.jsonl").read_bytes())
|
||||||
|
self.assertEqual(test_bytes, (output / "test.jsonl").read_bytes())
|
||||||
|
self.assertNotIn(
|
||||||
|
history[0]["prompt"],
|
||||||
|
(output / "validation.jsonl").read_text(encoding="utf-8"),
|
||||||
|
)
|
||||||
|
self.assertNotIn(
|
||||||
|
history[0]["prompt"],
|
||||||
|
(output / "test.jsonl").read_text(encoding="utf-8"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_near_eval_overlap_is_excluded_and_output_fails_closed(self):
|
||||||
|
with tempfile.TemporaryDirectory() as directory:
|
||||||
|
root = Path(directory)
|
||||||
|
base = root / "base"
|
||||||
|
base.mkdir()
|
||||||
|
purpose_data.write_jsonl(base / "train.jsonl", [example(0)])
|
||||||
|
words = [f"token{index}" for index in range(100)]
|
||||||
|
validation_prompt = " ".join(words)
|
||||||
|
purpose_data.write_jsonl(
|
||||||
|
base / "validation.jsonl",
|
||||||
|
[example(1, prompt=validation_prompt, purpose="review")],
|
||||||
|
)
|
||||||
|
purpose_data.write_jsonl(base / "test.jsonl", [example(2)])
|
||||||
|
words[50] = "replacement"
|
||||||
|
history_path = root / "history.jsonl"
|
||||||
|
purpose_data.write_jsonl(
|
||||||
|
history_path,
|
||||||
|
[example(10, prompt=" ".join(words), purpose="review")],
|
||||||
|
)
|
||||||
|
fixtures = root / "fixtures.json"
|
||||||
|
fixtures.write_text("[]", encoding="utf-8")
|
||||||
|
output = root / "output"
|
||||||
|
|
||||||
|
manifest = prepare_history_experiment.prepare(
|
||||||
|
base_dataset=base,
|
||||||
|
history_path=history_path,
|
||||||
|
fixtures_path=fixtures,
|
||||||
|
output_dir=output,
|
||||||
|
near_duplicate_threshold=0.92,
|
||||||
|
overwrite_output=False,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
{"near-eval-overlap": 1},
|
||||||
|
manifest["augmentation"]["exclusions"],
|
||||||
|
)
|
||||||
|
with self.assertRaisesRegex(purpose_data.DataError, "output is not empty"):
|
||||||
|
prepare_history_experiment.prepare(
|
||||||
|
base_dataset=base,
|
||||||
|
history_path=history_path,
|
||||||
|
fixtures_path=fixtures,
|
||||||
|
output_dir=output,
|
||||||
|
near_duplicate_threshold=0.92,
|
||||||
|
overwrite_output=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -15,6 +15,38 @@ import train
|
|||||||
import train_mlx
|
import train_mlx
|
||||||
|
|
||||||
|
|
||||||
|
class MLXDeviceTests(unittest.TestCase):
|
||||||
|
class Metal:
|
||||||
|
def __init__(self, available):
|
||||||
|
self.available = available
|
||||||
|
|
||||||
|
def is_available(self):
|
||||||
|
return self.available
|
||||||
|
|
||||||
|
class MLX:
|
||||||
|
cpu = "cpu"
|
||||||
|
gpu = "gpu"
|
||||||
|
|
||||||
|
def __init__(self, metal_available):
|
||||||
|
self.metal = MLXDeviceTests.Metal(metal_available)
|
||||||
|
self.selected = None
|
||||||
|
|
||||||
|
def set_default_device(self, device):
|
||||||
|
self.selected = device
|
||||||
|
|
||||||
|
def test_cpu_is_an_explicit_fallback(self):
|
||||||
|
mlx = self.MLX(metal_available=False)
|
||||||
|
train_mlx._configure_mlx_device(mlx, "cpu")
|
||||||
|
self.assertEqual("cpu", mlx.selected)
|
||||||
|
|
||||||
|
def test_metal_fails_closed_when_unavailable(self):
|
||||||
|
with self.assertRaisesRegex(train.DataError, "requires Apple Silicon"):
|
||||||
|
train_mlx._configure_mlx_device(
|
||||||
|
self.MLX(metal_available=False),
|
||||||
|
"metal",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class FixedShapeTokenizerTests(unittest.TestCase):
|
class FixedShapeTokenizerTests(unittest.TestCase):
|
||||||
class Tokenizer:
|
class Tokenizer:
|
||||||
pad_token_id = 0
|
pad_token_id = 0
|
||||||
|
|||||||
+24
-7
@@ -1,5 +1,5 @@
|
|||||||
#!/usr/bin/env python3
|
#!/usr/bin/env python3
|
||||||
"""Fine-tune purpose-lite natively on Apple Silicon with MLX."""
|
"""Fine-tune purpose-lite with MLX, using Metal by default."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -36,17 +36,28 @@ DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
|
|||||||
DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1-mlx"
|
DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1-mlx"
|
||||||
|
|
||||||
|
|
||||||
def _load_mlx() -> tuple[Any, Any, Any]:
|
def _configure_mlx_device(mx: Any, device: str) -> None:
|
||||||
|
if device == "metal":
|
||||||
|
if not mx.metal.is_available():
|
||||||
|
raise DataError("MLX Metal training requires Apple Silicon")
|
||||||
|
mx.set_default_device(mx.gpu)
|
||||||
|
return
|
||||||
|
if device == "cpu":
|
||||||
|
mx.set_default_device(mx.cpu)
|
||||||
|
return
|
||||||
|
raise DataError(f"unsupported MLX device {device!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def _load_mlx(device: str) -> tuple[Any, Any, Any]:
|
||||||
try:
|
try:
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
import mlx.optimizers as optim
|
import mlx.optimizers as optim
|
||||||
except ImportError as exc:
|
except ImportError as exc:
|
||||||
raise DataError(
|
raise DataError(
|
||||||
"MLX training requires Apple Silicon and requirements-mlx.txt"
|
"MLX training requires requirements-mlx.txt"
|
||||||
) from exc
|
) from exc
|
||||||
if not mx.metal.is_available():
|
_configure_mlx_device(mx, device)
|
||||||
raise DataError("MLX training requires the Apple Silicon Metal backend")
|
|
||||||
return mx, nn, optim
|
return mx, nn, optim
|
||||||
|
|
||||||
|
|
||||||
@@ -310,7 +321,7 @@ def _linear_schedule(
|
|||||||
|
|
||||||
|
|
||||||
def train(args: argparse.Namespace) -> dict[str, Any]:
|
def train(args: argparse.Namespace) -> dict[str, Any]:
|
||||||
mx, nn, optim = _load_mlx()
|
mx, nn, optim = _load_mlx(args.device)
|
||||||
try:
|
try:
|
||||||
from transformers import AutoTokenizer
|
from transformers import AutoTokenizer
|
||||||
|
|
||||||
@@ -718,7 +729,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"baseModel": str(model_dir),
|
"baseModel": str(model_dir),
|
||||||
"baseModelRevision": "local-checkpoint",
|
"baseModelRevision": "local-checkpoint",
|
||||||
"trainingBackend": "mlx",
|
"trainingBackend": "mlx",
|
||||||
"device": "metal",
|
"device": args.device,
|
||||||
"fixedInputShape": [1, MAX_LENGTH],
|
"fixedInputShape": [1, MAX_LENGTH],
|
||||||
"truncation": {
|
"truncation": {
|
||||||
"strategy": "head-tail-pair",
|
"strategy": "head-tail-pair",
|
||||||
@@ -774,6 +785,12 @@ def build_parser() -> argparse.ArgumentParser:
|
|||||||
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
|
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
|
||||||
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
|
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
|
||||||
parser.add_argument("--model", type=Path, required=True)
|
parser.add_argument("--model", type=Path, required=True)
|
||||||
|
parser.add_argument(
|
||||||
|
"--device",
|
||||||
|
choices=("metal", "cpu"),
|
||||||
|
default="metal",
|
||||||
|
help="MLX execution device (default: metal; cpu is a diagnostic fallback)",
|
||||||
|
)
|
||||||
parser.add_argument("--seed", type=int, default=20260730)
|
parser.add_argument("--seed", type=int, default=20260730)
|
||||||
parser.add_argument("--epochs", type=int, default=3)
|
parser.add_argument("--epochs", type=int, default=3)
|
||||||
parser.add_argument("--batch-size", type=int, default=32)
|
parser.add_argument("--batch-size", type=int, default=32)
|
||||||
|
|||||||
+20
-6
@@ -12,14 +12,18 @@ import numpy as np
|
|||||||
|
|
||||||
from purpose_data import LABELS, DataError, load_jsonl
|
from purpose_data import LABELS, DataError, load_jsonl
|
||||||
from train import enable_quantization_aware_training, encode_fixed_shape
|
from train import enable_quantization_aware_training, encode_fixed_shape
|
||||||
from train_mlx import _checkpoint_config, encode_fixed_shape_numpy
|
from train_mlx import (
|
||||||
|
_checkpoint_config,
|
||||||
|
_configure_mlx_device,
|
||||||
|
encode_fixed_shape_numpy,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
SCRIPT_DIR = Path(__file__).resolve().parent
|
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||||
DEFAULT_DATASET = SCRIPT_DIR / ".artifacts" / "dataset-v1" / "validation.jsonl"
|
DEFAULT_DATASET = SCRIPT_DIR / ".artifacts" / "dataset-v1" / "validation.jsonl"
|
||||||
|
|
||||||
|
|
||||||
def verify(model_dir: Path, dataset: Path, records: int) -> None:
|
def verify(model_dir: Path, dataset: Path, records: int, device: str) -> None:
|
||||||
try:
|
try:
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
@@ -38,10 +42,9 @@ def verify(model_dir: Path, dataset: Path, records: int) -> None:
|
|||||||
)
|
)
|
||||||
except ImportError as exc:
|
except ImportError as exc:
|
||||||
raise DataError(
|
raise DataError(
|
||||||
"verification requires requirements-mlx.txt on Apple Silicon"
|
"verification requires requirements-mlx.txt"
|
||||||
) from exc
|
) from exc
|
||||||
if not mx.metal.is_available():
|
_configure_mlx_device(mx, device)
|
||||||
raise DataError("verification requires the MLX Metal backend")
|
|
||||||
|
|
||||||
# Check the fake-quantization contract independently of the full model. Tiny
|
# Check the fake-quantization contract independently of the full model. Tiny
|
||||||
# backend-specific floating-point differences can cross later quantization
|
# backend-specific floating-point differences can cross later quantization
|
||||||
@@ -233,6 +236,12 @@ def build_parser() -> argparse.ArgumentParser:
|
|||||||
parser.add_argument("--model", type=Path, required=True)
|
parser.add_argument("--model", type=Path, required=True)
|
||||||
parser.add_argument("--dataset", type=Path, default=DEFAULT_DATASET)
|
parser.add_argument("--dataset", type=Path, default=DEFAULT_DATASET)
|
||||||
parser.add_argument("--records", type=int, default=8)
|
parser.add_argument("--records", type=int, default=8)
|
||||||
|
parser.add_argument(
|
||||||
|
"--device",
|
||||||
|
choices=("metal", "cpu"),
|
||||||
|
default="metal",
|
||||||
|
help="MLX execution device (default: metal; cpu is a diagnostic fallback)",
|
||||||
|
)
|
||||||
return parser
|
return parser
|
||||||
|
|
||||||
|
|
||||||
@@ -241,7 +250,12 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||||||
if args.records <= 0:
|
if args.records <= 0:
|
||||||
raise SystemExit("--records must be positive")
|
raise SystemExit("--records must be positive")
|
||||||
try:
|
try:
|
||||||
verify(args.model.expanduser(), args.dataset.expanduser(), args.records)
|
verify(
|
||||||
|
args.model.expanduser(),
|
||||||
|
args.dataset.expanduser(),
|
||||||
|
args.records,
|
||||||
|
args.device,
|
||||||
|
)
|
||||||
except (AssertionError, DataError) as exc:
|
except (AssertionError, DataError) as exc:
|
||||||
print(f"error: {exc}")
|
print(f"error: {exc}")
|
||||||
return 2
|
return 2
|
||||||
|
|||||||
Reference in New Issue
Block a user