Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 23:50:55 -07:00
parent 5687c90756
commit ac71c67e9a
7 changed files with 580 additions and 13 deletions
+20 -6
View File
@@ -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