Files
nucleic/ml/purpose-classifier/verify_mlx.py
T

267 lines
8.9 KiB
Python

#!/usr/bin/env python3
"""Verify MLX/PyTorch parity and checkpoint round-tripping before MLX training."""
from __future__ import annotations
import argparse
import tempfile
from pathlib import Path
from typing import Sequence
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,
_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, device: str) -> None:
try:
import mlx.core as mx
import mlx.nn as nn
import torch
from mlx.utils import tree_flatten
from safetensors import safe_open
from transformers import AutoModelForSequenceClassification, AutoTokenizer
from mlx_model import (
BertClassifierConfig,
BertForSequenceClassification,
fake_quantize_activation,
fake_quantize_linear_weight,
load_hugging_face_weights,
save_hugging_face_weights,
)
except ImportError as exc:
raise DataError(
"verification requires requirements-mlx.txt"
) from exc
_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
# thresholds, especially on padded tokens, so full QAT logits are expected to
# have more drift than the ordinary float graph.
activation_values = np.asarray(
[[-2.75, -0.125, 0.0, 0.625], [1.25, 3.5, -1.0, 0.25]],
dtype=np.float32,
)
torch_activation = torch.from_numpy(activation_values)
activation_minimum = min(0.0, float(torch_activation.amin().item()))
activation_maximum = max(0.0, float(torch_activation.amax().item()))
activation_scale = max(
(activation_maximum - activation_minimum) / 255.0,
torch.finfo(torch.float32).eps,
)
activation_zero_point = max(
0,
min(255, round(-activation_minimum / activation_scale)),
)
torch_quantized_activation = torch.fake_quantize_per_tensor_affine(
torch_activation,
activation_scale,
activation_zero_point,
0,
255,
).numpy()
mlx_quantized_activation = np.asarray(
fake_quantize_activation(mx.array(activation_values))
)
np.testing.assert_array_equal(
mlx_quantized_activation,
torch_quantized_activation,
)
weight_values = np.asarray(
[[-1.5, -0.25, 0.75, 1.25], [0.125, -0.875, 2.0, -1.25]],
dtype=np.float32,
)
torch_weight = torch.from_numpy(weight_values)
weight_scales = torch_weight.abs().amax(dim=1).div(127.0).clamp_min(
torch.finfo(torch.float32).eps
)
torch_quantized_weight = torch.fake_quantize_per_channel_affine(
torch_weight,
weight_scales,
torch.zeros_like(weight_scales, dtype=torch.int32),
0,
-127,
127,
).numpy()
mlx_quantized_weight = np.asarray(
fake_quantize_linear_weight(mx.array(weight_values))
)
np.testing.assert_array_equal(mlx_quantized_weight, torch_quantized_weight)
print("fake-quant primitives: exact")
checkpoint = model_dir / "model.safetensors"
if not checkpoint.is_file():
raise DataError(f"{checkpoint}: checkpoint is missing")
validation = load_jsonl(dataset)[:records]
if not validation:
raise DataError(f"{dataset}: no validation records")
texts = [record["prompt"] for record in validation]
labels = np.asarray(
[LABELS.index(record["purpose"]) for record in validation],
dtype=np.int32,
)
tokenizer = AutoTokenizer.from_pretrained(model_dir, local_files_only=True)
numpy_tokens = encode_fixed_shape_numpy(tokenizer, texts)
torch_tokens = encode_fixed_shape(tokenizer, texts, torch)
config_json = _checkpoint_config(model_dir)
config = BertClassifierConfig.from_hugging_face(config_json)
mlx_float = BertForSequenceClassification(
config,
quantization_aware=False,
)
load_hugging_face_weights(mlx_float, checkpoint)
mlx_float.eval()
mlx_float_logits = mlx_float(
**{key: mx.array(value) for key, value in numpy_tokens.items()}
)
mx.eval(mlx_float_logits)
mlx_float_logits = np.asarray(mlx_float_logits)
torch_model = AutoModelForSequenceClassification.from_pretrained(
model_dir,
local_files_only=True,
)
torch_model.eval()
with torch.inference_mode():
torch_float_logits = torch_model(**torch_tokens).logits.numpy()
np.testing.assert_allclose(
mlx_float_logits,
torch_float_logits,
rtol=1e-4,
atol=2e-5,
)
if not np.array_equal(
mlx_float_logits.argmax(axis=-1),
torch_float_logits.argmax(axis=-1),
):
raise DataError("MLX and PyTorch float predictions differ")
print(
"float parity: "
f"max_abs_error={np.max(np.abs(mlx_float_logits - torch_float_logits)):.3g}"
)
with tempfile.TemporaryDirectory(prefix="purpose-mlx-roundtrip-") as temp:
roundtrip = Path(temp) / "model.safetensors"
save_hugging_face_weights(mlx_float, roundtrip)
with safe_open(checkpoint, framework="np") as original, safe_open(
roundtrip,
framework="np",
) as converted:
if set(original.keys()) != set(converted.keys()):
raise DataError("MLX checkpoint round-trip changed parameter keys")
for key in original.keys():
np.testing.assert_array_equal(
original.get_tensor(key),
converted.get_tensor(key),
)
print("checkpoint round-trip: exact")
del mlx_float
mx.clear_cache()
mlx_qat = BertForSequenceClassification(
config,
quantization_aware=True,
)
load_hugging_face_weights(mlx_qat, checkpoint)
mlx_qat.eval()
mlx_qat_logits = mlx_qat(
**{key: mx.array(value) for key, value in numpy_tokens.items()}
)
mx.eval(mlx_qat_logits)
mlx_qat_logits = np.asarray(mlx_qat_logits)
enable_quantization_aware_training(torch, torch_model)
torch_model.eval()
with torch.inference_mode():
torch_qat_logits = torch_model(**torch_tokens).logits.numpy()
qat_max_abs_error = float(
np.max(np.abs(mlx_qat_logits - torch_qat_logits))
)
if not np.isfinite(qat_max_abs_error) or qat_max_abs_error > 1.0:
raise DataError(
"MLX and PyTorch QAT logits have excessive backend drift: "
f"{qat_max_abs_error:.3g}"
)
if not np.array_equal(
mlx_qat_logits.argmax(axis=-1),
torch_qat_logits.argmax(axis=-1),
):
raise DataError("MLX and PyTorch QAT predictions differ")
print(
"QAT parity: "
f"predictions=exact max_abs_error={qat_max_abs_error:.3g}"
)
mlx_qat.train()
mlx_labels = mx.array(labels)
def loss_function() -> object:
logits = mlx_qat(
**{key: mx.array(value) for key, value in numpy_tokens.items()}
)
return nn.losses.cross_entropy(logits, mlx_labels, reduction="mean")
loss, gradients = nn.value_and_grad(mlx_qat, loss_function)()
flat_gradients = [gradient for _, gradient in tree_flatten(gradients)]
mx.eval(loss, gradients)
if not all(
bool(mx.all(mx.isfinite(gradient)).item())
for gradient in flat_gradients
):
raise DataError("MLX QAT produced non-finite gradients")
if not any(
float(mx.max(mx.abs(gradient)).item()) > 0
for gradient in flat_gradients
):
raise DataError("MLX QAT produced only zero gradients")
print(f"QAT gradient smoke: loss={float(loss.item()):.6f}")
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
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
def main(argv: Sequence[str] | None = None) -> int:
args = build_parser().parse_args(argv)
if args.records <= 0:
raise SystemExit("--records must be positive")
try:
verify(
args.model.expanduser(),
args.dataset.expanduser(),
args.records,
args.device,
)
except (AssertionError, DataError) as exc:
print(f"error: {exc}")
return 2
return 0
if __name__ == "__main__":
raise SystemExit(main())