diff --git a/README.md b/README.md index c4c5862..0c58278 100644 --- a/README.md +++ b/README.md @@ -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 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 Core ML Tools no longer maintains the legacy ONNX converter, so the Apple artifact is diff --git a/prepare_history_experiment.py b/prepare_history_experiment.py new file mode 100644 index 0000000..5fba6d7 --- /dev/null +++ b/prepare_history_experiment.py @@ -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()) diff --git a/requirements-mlx.txt b/requirements-mlx.txt index 3a761cd..f85a82d 100644 --- a/requirements-mlx.txt +++ b/requirements-mlx.txt @@ -1,2 +1,3 @@ -r requirements.txt mlx==0.32.0 +mlx-cpu==0.32.0; sys_platform == "linux" diff --git a/tests/test_prepare_history_experiment.py b/tests/test_prepare_history_experiment.py new file mode 100644 index 0000000..07f445f --- /dev/null +++ b/tests/test_prepare_history_experiment.py @@ -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() diff --git a/tests/test_train_mlx.py b/tests/test_train_mlx.py index 61360a3..b19f5cc 100644 --- a/tests/test_train_mlx.py +++ b/tests/test_train_mlx.py @@ -15,6 +15,38 @@ import train 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 Tokenizer: pad_token_id = 0 diff --git a/train_mlx.py b/train_mlx.py index f8496b0..5aab6fb 100644 --- a/train_mlx.py +++ b/train_mlx.py @@ -1,5 +1,5 @@ #!/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 @@ -36,17 +36,28 @@ DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1" 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: import mlx.core as mx import mlx.nn as nn import mlx.optimizers as optim except ImportError as exc: raise DataError( - "MLX training requires Apple Silicon and requirements-mlx.txt" + "MLX training requires requirements-mlx.txt" ) from exc - if not mx.metal.is_available(): - raise DataError("MLX training requires the Apple Silicon Metal backend") + _configure_mlx_device(mx, device) return mx, nn, optim @@ -310,7 +321,7 @@ def _linear_schedule( def train(args: argparse.Namespace) -> dict[str, Any]: - mx, nn, optim = _load_mlx() + mx, nn, optim = _load_mlx(args.device) try: from transformers import AutoTokenizer @@ -718,7 +729,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]: "baseModel": str(model_dir), "baseModelRevision": "local-checkpoint", "trainingBackend": "mlx", - "device": "metal", + "device": args.device, "fixedInputShape": [1, MAX_LENGTH], "truncation": { "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("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) 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("--epochs", type=int, default=3) parser.add_argument("--batch-size", type=int, default=32) diff --git a/verify_mlx.py b/verify_mlx.py index f8a11b8..8673e31 100644 --- a/verify_mlx.py +++ b/verify_mlx.py @@ -12,14 +12,18 @@ import numpy as np from purpose_data import LABELS, DataError, load_jsonl 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 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: import mlx.core as mx import mlx.nn as nn @@ -38,10 +42,9 @@ def verify(model_dir: Path, dataset: Path, records: int) -> None: ) except ImportError as exc: raise DataError( - "verification requires requirements-mlx.txt on Apple Silicon" + "verification requires requirements-mlx.txt" ) from exc - if not mx.metal.is_available(): - raise DataError("verification requires the MLX Metal backend") + _configure_mlx_device(mx, device) # Check the fake-quantization contract independently of the full model. Tiny # 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("--dataset", type=Path, default=DEFAULT_DATASET) 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 @@ -241,7 +250,12 @@ def main(argv: Sequence[str] | None = None) -> int: if args.records <= 0: raise SystemExit("--records must be positive") 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: print(f"error: {exc}") return 2