Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
+20
-6
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user