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

This commit is contained in:
2026-07-30 20:14:37 -07:00
parent 09f98cdd00
commit 90092332db
8 changed files with 1668 additions and 0 deletions
+43
View File
@@ -145,6 +145,49 @@ backpropagation, selection, and ordinary checkpoint reload. The current shared C
then showed severe post-batch throttling, so no full candidate result is claimed from that then showed severe post-batch throttling, so no full candidate result is claimed from that
canary. canary.
### Native Apple Silicon training with MLX
Use the MLX backend when training on Apple Silicon. It implements the same six-layer BERT
classifier, fixed head-tail tokenization, export-matched QAT graph, cached-teacher
distillation, validation selection, and early stopping with native MLX arrays. Fake
quantization is decomposed into Metal-supported round, clip, and straight-through-gradient
operations, avoiding PyTorch's unsupported MPS fake-quant operator. Selected weights are
written back with the original Hugging Face parameter names, so the existing PyTorch
`export.py` and `eval.py` paths remain unchanged.
Install the additional pinned dependency into the macOS virtual environment:
```bash
ml/purpose-classifier/venv/bin/python -m pip install \
-r ml/purpose-classifier/requirements-mlx.txt
```
Before the first full run on a new MLX or Transformers version, run the fail-closed parity
check. It requires exact fake-quant primitives, float-logit parity, matching QAT
predictions with bounded backend drift, healthy QAT gradients, and an exact Hugging Face
→ MLX → Hugging Face weight round trip:
```bash
ml/purpose-classifier/venv/bin/python ml/purpose-classifier/verify_mlx.py \
--model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model
```
Then run the distilled QAT candidate natively on Metal:
```bash
ml/purpose-classifier/venv/bin/python -u ml/purpose-classifier/train_mlx.py \
--model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \
--distillation-cache \
ml/purpose-classifier/outputs/purpose-lite-v1-boundary-teacher.pt \
--distillation-weight 0.9 --distillation-temperature 2 \
--distillation-selection-weight 0.5 --quantization-aware \
--epochs 2 --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 \
--overwrite-output
```
For a wiring smoke test, use a small deterministic prefix: For a wiring smoke test, use a small deterministic prefix:
```bash ```bash
+39
View File
@@ -0,0 +1,39 @@
"""Hugging Face <-> MLX parameter-name conversion for purpose-lite BERT."""
from __future__ import annotations
_HF_TO_MLX_REPLACEMENTS = (
(".layer.", ".layers."),
(".self.key.", ".key_proj."),
(".self.query.", ".query_proj."),
(".self.value.", ".value_proj."),
(".attention.output.dense.", ".attention.out_proj."),
(".attention.output.LayerNorm.", ".ln1."),
(".output.LayerNorm.", ".ln2."),
(".intermediate.dense.", ".linear1."),
(".output.dense.", ".linear2."),
(".embeddings.LayerNorm.", ".embeddings.norm."),
(".pooler.dense.", ".pooler."),
)
_MLX_TO_HF_REPLACEMENTS = tuple(
(mlx, hugging_face) for hugging_face, mlx in reversed(_HF_TO_MLX_REPLACEMENTS)
)
def hugging_face_to_mlx_key(key: str) -> str:
"""Return the MLX BERT parameter name corresponding to a Transformers key."""
for hugging_face, mlx in _HF_TO_MLX_REPLACEMENTS:
key = key.replace(hugging_face, mlx)
return key
def mlx_to_hugging_face_key(key: str) -> str:
"""Return the Transformers parameter name corresponding to an MLX BERT key."""
for mlx, hugging_face in _MLX_TO_HF_REPLACEMENTS:
key = key.replace(mlx, hugging_face)
return key
+354
View File
@@ -0,0 +1,354 @@
"""Native-MLX BERT sequence classifier used by purpose-lite training.
The module layout follows Apple's reference MLX BERT implementation while adding
the classifier, training dropout, and the project's export-matched fake QAT graph.
Weights retain a reversible mapping to Hugging Face ``BertForSequenceClassification``.
"""
from __future__ import annotations
import math
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
from mlx_checkpoint import hugging_face_to_mlx_key, mlx_to_hugging_face_key
@dataclass(frozen=True)
class BertClassifierConfig:
vocab_size: int
hidden_size: int
num_hidden_layers: int
num_attention_heads: int
intermediate_size: int
max_position_embeddings: int
type_vocab_size: int
layer_norm_eps: float
hidden_dropout_prob: float
attention_probs_dropout_prob: float
classifier_dropout: float
num_labels: int
@classmethod
def from_hugging_face(cls, config: dict[str, Any]) -> "BertClassifierConfig":
classifier_dropout = config.get("classifier_dropout")
if classifier_dropout is None:
classifier_dropout = config["hidden_dropout_prob"]
return cls(
vocab_size=int(config["vocab_size"]),
hidden_size=int(config["hidden_size"]),
num_hidden_layers=int(config["num_hidden_layers"]),
num_attention_heads=int(config["num_attention_heads"]),
intermediate_size=int(config["intermediate_size"]),
max_position_embeddings=int(config["max_position_embeddings"]),
type_vocab_size=int(config["type_vocab_size"]),
layer_norm_eps=float(config["layer_norm_eps"]),
hidden_dropout_prob=float(config["hidden_dropout_prob"]),
attention_probs_dropout_prob=float(
config["attention_probs_dropout_prob"]
),
classifier_dropout=float(classifier_dropout),
num_labels=int(config.get("num_labels", len(config["id2label"]))),
)
def _affine_parameters(value: Any) -> tuple[Any, Any]:
detached = mx.stop_gradient(value.astype(mx.float32))
zero = mx.array(0.0, dtype=mx.float32)
minimum = mx.minimum(zero, mx.min(detached))
maximum = mx.maximum(zero, mx.max(detached))
scale = mx.maximum(
(maximum - minimum) / 255.0,
mx.array(mx.finfo(mx.float32).eps),
)
zero_point = mx.clip(mx.round(-minimum / scale), 0, 255)
return scale, zero_point
def _fake_quantize(
value: Any,
scale: Any,
zero_point: Any,
quant_min: int,
quant_max: int,
) -> Any:
"""Fake-quantize with an identity straight-through gradient."""
quantized = mx.clip(
mx.round(value / scale) + zero_point,
quant_min,
quant_max,
)
dequantized = (quantized - zero_point) * scale
return value + mx.stop_gradient(dequantized - value)
def fake_quantize_activation(value: Any) -> Any:
scale, zero_point = _affine_parameters(value)
return _fake_quantize(value, scale, zero_point, 0, 255)
def fake_quantize_linear_weight(weight: Any) -> Any:
detached = mx.stop_gradient(weight.astype(mx.float32))
scales = mx.maximum(
mx.max(mx.abs(detached), axis=1, keepdims=True) / 127.0,
mx.array(mx.finfo(mx.float32).eps),
)
return _fake_quantize(weight, scales, 0.0, -127, 127)
class QATLinear(nn.Linear):
def __call__(self, value: Any) -> Any:
value = fake_quantize_activation(value)
weight = fake_quantize_linear_weight(self.weight)
result = value @ weight.T
if "bias" in self:
result = result + self.bias
return fake_quantize_activation(result)
class QATEmbedding(nn.Embedding):
def __call__(self, indexes: Any) -> Any:
embedded = self.weight[indexes]
weight_scale, weight_zero_point = _affine_parameters(self.weight)
embedded = _fake_quantize(
embedded,
weight_scale,
weight_zero_point,
0,
255,
)
return fake_quantize_activation(embedded)
class BertSelfAttention(nn.Module):
def __init__(
self,
dims: int,
num_heads: int,
dropout: float,
linear: type[nn.Linear],
) -> None:
super().__init__()
if dims % num_heads:
raise ValueError("BERT hidden size must be divisible by attention heads")
self.num_heads = num_heads
self.head_dims = dims // num_heads
self.query_proj = linear(dims, dims, bias=True)
self.key_proj = linear(dims, dims, bias=True)
self.value_proj = linear(dims, dims, bias=True)
self.out_proj = linear(dims, dims, bias=True)
self.probability_dropout = nn.Dropout(dropout)
def __call__(self, value: Any, mask: Any | None) -> Any:
batch, length, _ = value.shape
def split_heads(projected: Any) -> Any:
return projected.reshape(
batch,
length,
self.num_heads,
self.head_dims,
).transpose(0, 2, 1, 3)
queries = split_heads(self.query_proj(value))
keys = split_heads(self.key_proj(value))
values = split_heads(self.value_proj(value))
scores = (queries @ keys.transpose(0, 1, 3, 2)) / math.sqrt(
self.head_dims
)
if mask is not None:
scores = scores + mask
probabilities = self.probability_dropout(mx.softmax(scores, axis=-1))
context = probabilities @ values
context = context.transpose(0, 2, 1, 3).reshape(batch, length, -1)
return self.out_proj(context)
class BertEncoderLayer(nn.Module):
def __init__(
self,
config: BertClassifierConfig,
linear: type[nn.Linear],
) -> None:
super().__init__()
self.attention = BertSelfAttention(
config.hidden_size,
config.num_attention_heads,
config.attention_probs_dropout_prob,
linear,
)
self.ln1 = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
)
self.ln2 = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
)
self.linear1 = linear(
config.hidden_size,
config.intermediate_size,
bias=True,
)
self.linear2 = linear(
config.intermediate_size,
config.hidden_size,
bias=True,
)
self.gelu = nn.GELU(approx="none")
self.attention_output_dropout = nn.Dropout(config.hidden_dropout_prob)
self.output_dropout = nn.Dropout(config.hidden_dropout_prob)
def __call__(self, value: Any, mask: Any | None) -> Any:
attention = self.attention_output_dropout(self.attention(value, mask))
value = self.ln1(value + attention)
feed_forward = self.linear2(self.gelu(self.linear1(value)))
return self.ln2(value + self.output_dropout(feed_forward))
class BertEncoder(nn.Module):
def __init__(
self,
config: BertClassifierConfig,
linear: type[nn.Linear],
) -> None:
super().__init__()
self.layers = [
BertEncoderLayer(config, linear)
for _ in range(config.num_hidden_layers)
]
def __call__(self, value: Any, mask: Any | None) -> Any:
for layer in self.layers:
value = layer(value, mask)
return value
class BertEmbeddings(nn.Module):
def __init__(
self,
config: BertClassifierConfig,
embedding: type[nn.Embedding],
) -> None:
super().__init__()
self.word_embeddings = embedding(config.vocab_size, config.hidden_size)
self.token_type_embeddings = embedding(
config.type_vocab_size,
config.hidden_size,
)
self.position_embeddings = embedding(
config.max_position_embeddings,
config.hidden_size,
)
self.norm = nn.LayerNorm(
config.hidden_size,
eps=config.layer_norm_eps,
)
self.dropout = nn.Dropout(config.hidden_dropout_prob)
def __call__(self, input_ids: Any, token_type_ids: Any | None) -> Any:
if token_type_ids is None:
token_type_ids = mx.zeros_like(input_ids)
position_ids = mx.broadcast_to(
mx.arange(input_ids.shape[1]),
input_ids.shape,
)
embeddings = (
self.word_embeddings(input_ids)
+ self.position_embeddings(position_ids)
+ self.token_type_embeddings(token_type_ids)
)
return self.dropout(self.norm(embeddings))
class BertModel(nn.Module):
def __init__(
self,
config: BertClassifierConfig,
linear: type[nn.Linear],
embedding: type[nn.Embedding],
) -> None:
super().__init__()
self.embeddings = BertEmbeddings(config, embedding)
self.encoder = BertEncoder(config, linear)
self.pooler = linear(config.hidden_size, config.hidden_size, bias=True)
def __call__(
self,
input_ids: Any,
attention_mask: Any | None,
token_type_ids: Any | None,
) -> tuple[Any, Any]:
value = self.embeddings(input_ids, token_type_ids)
additive_mask = None
if attention_mask is not None:
visible = attention_mask.astype(mx.bool_)[:, None, None, :]
additive_mask = mx.where(
visible,
mx.array(0.0, dtype=value.dtype),
mx.array(-1e4, dtype=value.dtype),
)
sequence = self.encoder(value, additive_mask)
pooled = mx.tanh(self.pooler(sequence[:, 0]))
return sequence, pooled
class BertForSequenceClassification(nn.Module):
def __init__(
self,
config: BertClassifierConfig,
*,
quantization_aware: bool,
) -> None:
super().__init__()
linear = QATLinear if quantization_aware else nn.Linear
embedding = QATEmbedding if quantization_aware else nn.Embedding
self.bert = BertModel(config, linear, embedding)
self.dropout = nn.Dropout(config.classifier_dropout)
self.classifier = linear(
config.hidden_size,
config.num_labels,
bias=True,
)
self.quantization_aware = quantization_aware
def __call__(
self,
input_ids: Any,
attention_mask: Any | None = None,
token_type_ids: Any | None = None,
) -> Any:
_, pooled = self.bert(
input_ids,
attention_mask,
token_type_ids,
)
return self.classifier(self.dropout(pooled))
def load_hugging_face_weights(model: nn.Module, checkpoint: Path) -> None:
weights = mx.load(str(checkpoint))
converted = [
(hugging_face_to_mlx_key(key), value) for key, value in weights.items()
]
model.load_weights(converted, strict=True)
mx.eval(model.parameters())
def save_hugging_face_weights(model: nn.Module, checkpoint: Path) -> None:
mx.eval(model.parameters())
weights = {
mlx_to_hugging_face_key(key): value
for key, value in tree_flatten(model.parameters())
}
mx.save_safetensors(
str(checkpoint),
weights,
metadata={"format": "pt"},
)
+2
View File
@@ -0,0 +1,2 @@
-r requirements.txt
mlx==0.32.0
+57
View File
@@ -0,0 +1,57 @@
import sys
import unittest
from pathlib import Path
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
from mlx_checkpoint import hugging_face_to_mlx_key, mlx_to_hugging_face_key
class MLXCheckpointTests(unittest.TestCase):
def test_representative_bert_keys_round_trip(self):
keys = (
"bert.embeddings.LayerNorm.weight",
"bert.embeddings.word_embeddings.weight",
"bert.encoder.layer.0.attention.self.query.weight",
"bert.encoder.layer.2.attention.self.key.bias",
"bert.encoder.layer.4.attention.self.value.weight",
"bert.encoder.layer.5.attention.output.dense.bias",
"bert.encoder.layer.1.attention.output.LayerNorm.weight",
"bert.encoder.layer.3.intermediate.dense.weight",
"bert.encoder.layer.3.output.dense.bias",
"bert.encoder.layer.3.output.LayerNorm.weight",
"bert.pooler.dense.weight",
"classifier.weight",
)
for key in keys:
with self.subTest(key=key):
self.assertEqual(
key,
mlx_to_hugging_face_key(hugging_face_to_mlx_key(key)),
)
def test_expected_mlx_names(self):
self.assertEqual(
"bert.encoder.layers.0.attention.query_proj.weight",
hugging_face_to_mlx_key(
"bert.encoder.layer.0.attention.self.query.weight"
),
)
self.assertEqual(
"bert.encoder.layers.0.ln1.bias",
hugging_face_to_mlx_key(
"bert.encoder.layer.0.attention.output.LayerNorm.bias"
),
)
self.assertEqual(
"bert.encoder.layers.0.ln2.weight",
hugging_face_to_mlx_key(
"bert.encoder.layer.0.output.LayerNorm.weight"
),
)
if __name__ == "__main__":
unittest.main()
+72
View File
@@ -0,0 +1,72 @@
import json
import sys
import tempfile
import unittest
from pathlib import Path
import numpy as np
import torch
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
import train
import train_mlx
class FixedShapeTokenizerTests(unittest.TestCase):
class Tokenizer:
pad_token_id = 0
cls_token_id = 1
sep_token_id = 2
padding_side = "right"
model_input_names = ["input_ids", "attention_mask", "token_type_ids"]
def __call__(self, texts, **_):
return {
"input_ids": [
list(range(10, 10 + int(text.split()[-1]))) for text in texts
]
}
@staticmethod
def num_special_tokens_to_add(pair=False):
return 3 if pair else 2
def test_mlx_encoding_matches_pytorch_encoding(self):
texts = ["tokens 3", "tokens 200"]
pytorch = train.encode_fixed_shape(self.Tokenizer(), texts, torch)
mlx = train_mlx.encode_fixed_shape_numpy(self.Tokenizer(), texts)
self.assertEqual(set(pytorch), set(mlx))
for key in pytorch:
with self.subTest(key=key):
np.testing.assert_array_equal(pytorch[key].numpy(), mlx[key])
class CheckpointConfigTests(unittest.TestCase):
def test_rejects_changed_label_order(self):
config = {
"model_type": "bert",
"hidden_size": 384,
"num_hidden_layers": 6,
"id2label": {
str(index): label
for index, label in enumerate(reversed(train.LABELS))
},
}
with tempfile.TemporaryDirectory() as temp:
model_dir = Path(temp)
(model_dir / "config.json").write_text(
json.dumps(config),
encoding="utf-8",
)
with self.assertRaisesRegex(
train.DataError,
"label order",
):
train_mlx._checkpoint_config(model_dir)
if __name__ == "__main__":
unittest.main()
+849
View File
@@ -0,0 +1,849 @@
#!/usr/bin/env python3
"""Fine-tune purpose-lite natively on Apple Silicon with MLX."""
from __future__ import annotations
import argparse
import json
import math
import random
import shutil
import sys
import time
from pathlib import Path
from typing import Any, Iterator, Sequence
import numpy as np
from purpose_data import LABELS, DataError, load_jsonl, write_json
from train import (
HEAD_TOKENS,
MAX_LENGTH,
TAIL_TOKENS,
_fit_temperature,
_validate_split,
choose_confidence_thresholds,
classification_metrics,
distillation_record_keys,
expected_calibration_error,
prepare_text,
training_weight,
)
SCRIPT_DIR = Path(__file__).resolve().parent
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]:
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"
) from exc
if not mx.metal.is_available():
raise DataError("MLX training requires the Apple Silicon Metal backend")
return mx, nn, optim
def encode_fixed_shape_numpy(
tokenizer: Any,
texts: Sequence[str],
) -> dict[str, np.ndarray]:
"""Apply the same fixed 128-token head-tail contract as train.py."""
normalized = [prepare_text(text) for text in texts]
raw = tokenizer(
normalized,
add_special_tokens=False,
padding=False,
truncation=False,
return_attention_mask=False,
return_token_type_ids=False,
verbose=False,
)
if not isinstance(raw.get("input_ids"), list):
raise DataError("tokenizer did not return input_ids")
if tokenizer.pad_token_id is None:
raise DataError("purpose-lite tokenizer must define a padding token")
if tokenizer.cls_token_id is None or tokenizer.sep_token_id is None:
raise DataError("purpose-lite tokenizer must define BERT CLS and SEP tokens")
if tokenizer.padding_side != "right":
raise DataError("purpose-lite tokenizer must use right padding")
input_rows: list[list[int]] = []
mask_rows: list[list[int]] = []
type_rows: list[list[int]] = []
include_token_types = "token_type_ids" in tokenizer.model_input_names
single_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=False)
pair_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=True)
if pair_budget != HEAD_TOKENS + TAIL_TOKENS:
raise DataError(
"purpose-lite tokenizer special-token layout changed; expected three "
"tokens for head-tail inputs"
)
for content in raw["input_ids"]:
if len(content) <= single_budget:
first = content
second = None
else:
first = content[:HEAD_TOKENS]
second = content[-TAIL_TOKENS:]
if second is None:
input_ids = [tokenizer.cls_token_id] + first + [tokenizer.sep_token_id]
token_types = [0] * len(input_ids)
else:
input_ids = (
[tokenizer.cls_token_id]
+ first
+ [tokenizer.sep_token_id]
+ second
+ [tokenizer.sep_token_id]
)
token_types = [0] * (len(first) + 2) + [1] * (len(second) + 1)
if len(input_ids) > MAX_LENGTH:
raise DataError("fixed-shape tokenizer exceeded its 128-token contract")
padding = MAX_LENGTH - len(input_ids)
input_rows.append(input_ids + [tokenizer.pad_token_id] * padding)
mask_rows.append([1] * len(input_ids) + [0] * padding)
if include_token_types:
type_rows.append(token_types + [0] * padding)
encoded = {
"input_ids": np.asarray(input_rows, dtype=np.int32),
"attention_mask": np.asarray(mask_rows, dtype=np.int32),
}
if include_token_types:
encoded["token_type_ids"] = np.asarray(type_rows, dtype=np.int32)
return encoded
def _encode_records(
tokenizer: Any,
records: Sequence[dict[str, Any]],
*,
chunk_size: int = 256,
) -> dict[str, np.ndarray]:
chunks: dict[str, list[np.ndarray]] = {}
for start in range(0, len(records), chunk_size):
encoded = encode_fixed_shape_numpy(
tokenizer,
[record["prompt"] for record in records[start : start + chunk_size]],
)
for key, value in encoded.items():
chunks.setdefault(key, []).append(value)
return {key: np.concatenate(values) for key, values in chunks.items()}
def _batch_indexes(
size: int,
batch_size: int,
*,
permutation: np.ndarray | None = None,
) -> Iterator[np.ndarray]:
indexes = permutation if permutation is not None else np.arange(size)
for start in range(0, size, batch_size):
yield indexes[start : start + batch_size]
def _mlx_batch(
mx: Any,
encoded: dict[str, np.ndarray],
indexes: np.ndarray,
) -> dict[str, Any]:
return {key: mx.array(value[indexes]) for key, value in encoded.items()}
def _evaluate(
mx: Any,
model: Any,
encoded: dict[str, np.ndarray],
labels: np.ndarray,
batch_size: int,
) -> tuple[np.ndarray, np.ndarray]:
model.eval()
logits: list[np.ndarray] = []
for indexes in _batch_indexes(len(labels), batch_size):
batch = _mlx_batch(mx, encoded, indexes)
output = model(**batch)
mx.eval(output)
logits.append(np.asarray(output))
return np.concatenate(logits), labels.copy()
def _teacher_cache(
path: Path,
train_records: Sequence[dict[str, Any]],
validation_records: Sequence[dict[str, Any]],
) -> tuple[np.ndarray, np.ndarray]:
try:
import torch
except ImportError as exc:
raise DataError(
"loading the existing teacher cache requires PyTorch"
) from exc
if not path.is_file():
raise DataError(f"{path}: distillation cache is missing")
try:
cache = torch.load(path, map_location="cpu", weights_only=True)
if cache["schemaVersion"] != 1 or cache["labels"] != list(LABELS):
raise DataError("distillation cache contract does not match purpose-lite")
if cache["trainRecordKeys"][: len(train_records)] != distillation_record_keys(
train_records
):
raise DataError("distillation cache does not match the training split")
if cache["validationRecordKeys"][
: len(validation_records)
] != distillation_record_keys(validation_records):
raise DataError("distillation cache does not match the validation split")
train_logits = (
cache["trainLogits"][: len(train_records)].float().numpy().copy()
)
validation_logits = (
cache["validationLogits"][: len(validation_records)]
.float()
.numpy()
.copy()
)
except DataError:
raise
except (KeyError, TypeError, ValueError, RuntimeError) as exc:
raise DataError(f"{path}: cannot load distillation cache: {exc}") from exc
if train_logits.shape != (len(train_records), len(LABELS)):
raise DataError("distillation training logits have the wrong shape")
if validation_logits.shape != (len(validation_records), len(LABELS)):
raise DataError("distillation validation logits have the wrong shape")
return train_logits, validation_logits
def _checkpoint_config(model_dir: Path) -> dict[str, Any]:
config_path = model_dir / "config.json"
try:
config = json.loads(config_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise DataError(f"{config_path}: cannot load model config: {exc}") from exc
if (
config.get("model_type") != "bert"
or config.get("hidden_size") != 384
or config.get("num_hidden_layers") != 6
or len(config.get("id2label", {})) != len(LABELS)
):
raise DataError("MLX purpose-lite requires the 6-layer 384-wide BERT classifier")
configured_labels = [
config["id2label"].get(str(index), config["id2label"].get(index))
for index in range(len(LABELS))
]
if configured_labels != list(LABELS):
raise DataError("MLX checkpoint label order does not match purpose-lite")
return config
def _save_checkpoint(
mx: Any,
model: Any,
source_dir: Path,
destination: Path,
config: dict[str, Any],
*,
quantization_aware: bool,
) -> None:
from mlx_model import save_hugging_face_weights
if destination.exists():
shutil.rmtree(destination)
destination.mkdir(parents=True)
for source in source_dir.iterdir():
if source.name.startswith("model") and source.suffix == ".safetensors":
continue
target = destination / source.name
if source.is_dir():
shutil.copytree(source, target)
else:
shutil.copy2(source, target)
output_config = dict(config)
output_config["purpose_classifier_training_backend"] = "mlx"
output_config["purpose_classifier_quantization_aware_training"] = bool(
quantization_aware
)
(destination / "config.json").write_text(
json.dumps(output_config, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
save_hugging_face_weights(model, destination / "model.safetensors")
mx.eval(model.parameters())
def _softmax(values: np.ndarray) -> np.ndarray:
shifted = values - values.max(axis=-1, keepdims=True)
exponentials = np.exp(shifted)
return exponentials / exponentials.sum(axis=-1, keepdims=True)
def _linear_schedule(
mx: Any,
learning_rate: float,
total_steps: int,
warmup_steps: int,
) -> Any:
def schedule(step: Any) -> Any:
step = step.astype(mx.float32)
if warmup_steps:
warmup = learning_rate * step / warmup_steps
else:
warmup = mx.array(learning_rate)
remaining = max(total_steps - warmup_steps, 1)
decay = learning_rate * mx.maximum(
0.0,
(total_steps - step) / remaining,
)
if warmup_steps:
return mx.where(step < warmup_steps, warmup, decay)
return decay
return schedule
def train(args: argparse.Namespace) -> dict[str, Any]:
mx, nn, optim = _load_mlx()
try:
from transformers import AutoTokenizer
from mlx_model import (
BertClassifierConfig,
BertForSequenceClassification,
QATEmbedding,
QATLinear,
load_hugging_face_weights,
)
except ImportError as exc:
raise DataError(
"MLX training dependencies are missing; install requirements-mlx.txt"
) from exc
model_dir = args.model.expanduser()
checkpoint = model_dir / "model.safetensors"
if not checkpoint.is_file():
raise DataError("--model must be a local Hugging Face safetensors checkpoint")
output_dir: Path = args.output_dir
try:
model_dir.resolve().relative_to(output_dir.resolve())
except ValueError:
pass
else:
raise DataError("--model must not be inside --output-dir")
if output_dir.exists() and any(output_dir.iterdir()):
if not args.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)
train_path = args.dataset_dir / "train.jsonl"
validation_path = args.dataset_dir / "validation.jsonl"
train_records = load_jsonl(train_path)
validation_records = load_jsonl(validation_path)
_validate_split(train_records, train_path)
_validate_split(validation_records, validation_path)
if args.max_train_records:
train_records = train_records[: args.max_train_records]
if args.max_validation_records:
validation_records = validation_records[: args.max_validation_records]
config_json = _checkpoint_config(model_dir)
config = BertClassifierConfig.from_hugging_face(config_json)
tokenizer = AutoTokenizer.from_pretrained(model_dir, local_files_only=True)
print("tokenizing train and validation splits", flush=True)
encoded_train = _encode_records(tokenizer, train_records)
encoded_validation = _encode_records(tokenizer, validation_records)
label_to_id = {label: index for index, label in enumerate(LABELS)}
train_labels = np.asarray(
[label_to_id[record["purpose"]] for record in train_records],
dtype=np.int32,
)
validation_labels = np.asarray(
[label_to_id[record["purpose"]] for record in validation_records],
dtype=np.int32,
)
sample_weights = np.asarray(
[
training_weight(record, args.boundary_weight)
for record in train_records
],
dtype=np.float32,
)
teacher_train_logits = None
teacher_validation_logits = None
if args.distillation_cache is not None:
teacher_train_logits, teacher_validation_logits = _teacher_cache(
args.distillation_cache.expanduser(),
train_records,
validation_records,
)
random.seed(args.seed)
np.random.seed(args.seed)
mx.random.seed(args.seed)
model = BertForSequenceClassification(
config,
quantization_aware=args.quantization_aware,
)
load_hugging_face_weights(model, checkpoint)
qat_modules = {
"linear": sum(isinstance(module, QATLinear) for module in model.modules()),
"embedding": sum(
isinstance(module, QATEmbedding) for module in model.modules()
),
}
validation_scorable = np.asarray(
[record.get("slice") != "vague-eval" for record in validation_records],
dtype=np.bool_,
)
initial_logits, _ = _evaluate(
mx,
model,
encoded_validation,
validation_labels,
args.eval_batch_size,
)
initial_predictions = initial_logits.argmax(axis=-1)
initial_metrics = classification_metrics(
validation_labels[validation_scorable].tolist(),
initial_predictions[validation_scorable].tolist(),
)
teacher_validation_predictions = (
teacher_validation_logits.argmax(axis=-1)
if teacher_validation_logits is not None
else None
)
def teacher_agreement(predictions: np.ndarray) -> float | None:
if teacher_validation_predictions is None:
return None
return float(
np.mean(
predictions[validation_scorable]
== teacher_validation_predictions[validation_scorable]
)
)
def selection_score(accuracy: float, agreement: float | None) -> float:
if agreement is None:
return accuracy
weight = args.distillation_selection_weight
return (accuracy + weight * agreement) / (1.0 + weight)
initial_agreement = teacher_agreement(initial_predictions)
initial_selection_score = selection_score(
initial_metrics["accuracy"],
initial_agreement,
)
if initial_agreement is not None:
initial_metrics["teacherAgreement"] = initial_agreement
initial_metrics["selectionScore"] = initial_selection_score
best_accuracy = initial_metrics["accuracy"]
best_selection_score = initial_selection_score
best_dir = output_dir / "model"
_save_checkpoint(
mx,
model,
model_dir,
best_dir,
config_json,
quantization_aware=args.quantization_aware,
)
print(
f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} "
f"macro_recall={initial_metrics['macroRecall']:.4%}"
+ (
f" teacher_agreement={initial_agreement:.4%} "
f"selection_score={initial_selection_score:.4%}"
if initial_agreement is not None
else ""
),
flush=True,
)
steps_per_epoch = math.ceil(len(train_records) / args.batch_size)
total_steps = steps_per_epoch * args.epochs
schedule = _linear_schedule(
mx,
args.learning_rate,
total_steps,
round(total_steps * args.warmup_ratio),
)
optimizer = optim.AdamW(
learning_rate=schedule,
weight_decay=args.weight_decay,
bias_correction=True,
)
def loss_function(
input_ids: Any,
attention_mask: Any,
token_type_ids: Any,
labels: Any,
weights: Any,
teacher_logits: Any | None,
) -> tuple[Any, Any, Any]:
logits = model(
input_ids=input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
)
label_loss = nn.losses.cross_entropy(
logits,
labels,
reduction="none",
)
distillation_loss = mx.zeros_like(label_loss)
if teacher_logits is not None:
temperature = args.distillation_temperature
student_log_probabilities = (
logits / temperature
- mx.logsumexp(logits / temperature, axis=-1, keepdims=True)
)
teacher_probabilities = mx.softmax(
teacher_logits / temperature,
axis=-1,
)
teacher_log_probabilities = mx.log(
mx.maximum(teacher_probabilities, 1e-12)
)
distillation_loss = (
mx.sum(
teacher_probabilities
* (teacher_log_probabilities - student_log_probabilities),
axis=-1,
)
* temperature
* temperature
)
per_record_loss = (
(1.0 - args.distillation_weight) * label_loss
+ args.distillation_weight * distillation_loss
)
denominator = mx.sum(weights)
loss = mx.sum(per_record_loss * weights) / denominator
mean_label = mx.sum(label_loss * weights) / denominator
mean_distillation = mx.sum(distillation_loss * weights) / denominator
return loss, mean_label, mean_distillation
loss_and_grad = nn.value_and_grad(model, loss_function)
rng = np.random.default_rng(args.seed)
epochs_without_improvement = 0
stopped_early = False
history: list[dict[str, Any]] = []
started = time.perf_counter()
for epoch in range(1, args.epochs + 1):
epoch_started = time.perf_counter()
model.train()
running_loss = 0.0
running_label_loss = 0.0
running_distillation_loss = 0.0
permutation = rng.permutation(len(train_records))
for step, indexes in enumerate(
_batch_indexes(
len(train_records),
args.batch_size,
permutation=permutation,
),
1,
):
batch = _mlx_batch(mx, encoded_train, indexes)
labels = mx.array(train_labels[indexes])
weights = mx.array(sample_weights[indexes])
teacher_logits = (
mx.array(teacher_train_logits[indexes])
if teacher_train_logits is not None
else None
)
(loss, label_loss, distillation_loss), gradients = loss_and_grad(
batch["input_ids"],
batch["attention_mask"],
batch.get("token_type_ids"),
labels,
weights,
teacher_logits,
)
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
optimizer.update(model, gradients)
mx.eval(
model.parameters(),
optimizer.state,
loss,
label_loss,
distillation_loss,
)
running_loss += float(loss.item())
running_label_loss += float(label_loss.item())
running_distillation_loss += float(distillation_loss.item())
if args.progress_steps and (
step % args.progress_steps == 0 or step == steps_per_epoch
):
print(
f"epoch {epoch} step {step}/{steps_per_epoch} "
f"mean_loss={running_loss / step:.4f} "
f"label_loss={running_label_loss / step:.4f} "
f"distill_loss={running_distillation_loss / step:.4f} "
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True,
)
logits, _ = _evaluate(
mx,
model,
encoded_validation,
validation_labels,
args.eval_batch_size,
)
predictions = logits.argmax(axis=-1)
metrics = classification_metrics(
validation_labels[validation_scorable].tolist(),
predictions[validation_scorable].tolist(),
)
agreement = teacher_agreement(predictions)
candidate_selection_score = selection_score(metrics["accuracy"], agreement)
if agreement is not None:
metrics["teacherAgreement"] = agreement
metrics["selectionScore"] = candidate_selection_score
metrics["epoch"] = epoch
metrics["meanTrainingLoss"] = running_loss / steps_per_epoch
metrics["meanLabelLoss"] = running_label_loss / steps_per_epoch
metrics["meanDistillationLoss"] = (
running_distillation_loss / steps_per_epoch
)
history.append(metrics)
print(
f"epoch {epoch}: loss={metrics['meanTrainingLoss']:.4f} "
f"validation_accuracy={metrics['accuracy']:.4%} "
f"macro_recall={metrics['macroRecall']:.4%}"
+ (
f" teacher_agreement={agreement:.4%} "
f"selection_score={candidate_selection_score:.4%}"
if agreement is not None
else ""
),
flush=True,
)
improvement = candidate_selection_score - best_selection_score
if improvement > args.minimum_improvement:
best_accuracy = metrics["accuracy"]
best_selection_score = candidate_selection_score
epochs_without_improvement = 0
_save_checkpoint(
mx,
model,
model_dir,
best_dir,
config_json,
quantization_aware=args.quantization_aware,
)
else:
epochs_without_improvement += 1
if epochs_without_improvement >= args.early_stopping_patience:
stopped_early = True
print(
f"early stopping after epoch {epoch}: no selection-score "
f"improvement greater than {args.minimum_improvement:.4%} "
f"for {args.early_stopping_patience} epoch(s)",
flush=True,
)
break
final_model = BertForSequenceClassification(
config,
quantization_aware=False,
)
load_hugging_face_weights(final_model, best_dir / "model.safetensors")
logits, labels = _evaluate(
mx,
final_model,
encoded_validation,
validation_labels,
args.eval_batch_size,
)
try:
import torch
except ImportError as exc:
raise DataError("final calibration requires PyTorch") from exc
temperature = _fit_temperature(
torch,
torch.from_numpy(logits),
torch.from_numpy(labels.astype(np.int64)),
)
calibrated = _softmax(logits / temperature)
sorted_indexes = np.argsort(calibrated, axis=-1)
top_indexes = sorted_indexes[:, -1]
second_indexes = sorted_indexes[:, -2]
row_indexes = np.arange(len(labels))
top_probabilities = calibrated[row_indexes, top_indexes]
margins = (
top_probabilities - calibrated[row_indexes, second_indexes]
)
correct = (
(top_indexes == labels) & validation_scorable
).tolist()
thresholds = choose_confidence_thresholds(
top_probabilities.tolist(),
margins.tolist(),
correct,
high_precision=args.high_precision,
accepted_precision=args.accepted_precision,
)
calibration = {
"schemaVersion": 1,
"modelVersion": "purpose-lite-v1",
"labels": list(LABELS),
"temperature": temperature,
"confidence": thresholds,
"validationECE": expected_calibration_error(
top_probabilities.tolist(),
correct,
),
}
metrics = {
"modelVersion": "purpose-lite-v1",
"baseModel": str(model_dir),
"baseModelRevision": "local-checkpoint",
"trainingBackend": "mlx",
"device": "metal",
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
},
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"validationRecords": len(validation_records),
"scoredValidationRecords": int(validation_scorable.sum()),
"vagueAbstentionValidationRecords": int(
(~validation_scorable).sum()
),
"boundaryTrainingWeight": args.boundary_weight,
"quantizationAwareTraining": args.quantization_aware,
"quantizationAwareModules": qat_modules,
"distillation": {
"cache": (
str(args.distillation_cache)
if args.distillation_cache is not None
else None
),
"weight": args.distillation_weight,
"temperature": args.distillation_temperature,
"selectionAgreementWeight": args.distillation_selection_weight,
},
"bestValidationAccuracy": best_accuracy,
"bestValidationSelectionScore": best_selection_score,
"initialValidation": initial_metrics,
"epochsCompleted": len(history),
"stoppedEarly": stopped_early,
"bestValidation": classification_metrics(
labels[validation_scorable].tolist(),
top_indexes[validation_scorable].tolist(),
),
"history": history,
"calibration": calibration,
}
write_json(output_dir / "calibration.json", calibration)
write_json(output_dir / "metrics.json", metrics)
write_json(
output_dir / "training-config.json",
{
key: str(value) if isinstance(value, Path) else value
for key, value in vars(args).items()
},
)
return metrics
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
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("--seed", type=int, default=20260730)
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--eval-batch-size", type=int, default=64)
parser.add_argument("--learning-rate", type=float, default=2e-5)
parser.add_argument("--weight-decay", type=float, default=0.01)
parser.add_argument("--warmup-ratio", type=float, default=0.1)
parser.add_argument("--max-grad-norm", type=float, default=1.0)
parser.add_argument("--progress-steps", type=int, default=50)
parser.add_argument("--early-stopping-patience", type=int, default=2)
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
parser.add_argument("--boundary-weight", type=float, default=1.0)
parser.add_argument("--quantization-aware", action="store_true")
parser.add_argument("--distillation-cache", type=Path)
parser.add_argument("--distillation-weight", type=float, default=0.0)
parser.add_argument("--distillation-temperature", type=float, default=2.0)
parser.add_argument("--distillation-selection-weight", type=float, default=0.0)
parser.add_argument("--high-precision", type=float, default=0.98)
parser.add_argument("--accepted-precision", type=float, default=0.95)
parser.add_argument("--max-train-records", type=int)
parser.add_argument("--max-validation-records", type=int)
parser.add_argument("--overwrite-output", action="store_true")
return parser
def main(argv: Sequence[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
for name in (
"epochs",
"batch_size",
"eval_batch_size",
"early_stopping_patience",
):
if getattr(args, name) <= 0:
parser.error(f"--{name.replace('_', '-')} must be positive")
if args.progress_steps < 0:
parser.error("--progress-steps must be non-negative")
if args.learning_rate <= 0:
parser.error("--learning-rate must be positive")
if not 0 <= args.warmup_ratio < 1:
parser.error("--warmup-ratio must be in [0, 1)")
if args.boundary_weight <= 0:
parser.error("--boundary-weight must be positive")
if not 0 <= args.distillation_weight <= 1:
parser.error("--distillation-weight must be in [0, 1]")
if args.distillation_temperature <= 0:
parser.error("--distillation-temperature must be positive")
if not 0 <= args.distillation_selection_weight <= 1:
parser.error("--distillation-selection-weight must be in [0, 1]")
if (args.distillation_cache is None) != (args.distillation_weight == 0):
parser.error(
"--distillation-cache and a positive --distillation-weight "
"must be supplied together"
)
if args.distillation_selection_weight and args.distillation_cache is None:
parser.error(
"--distillation-selection-weight requires --distillation-cache"
)
try:
metrics = train(args)
except DataError as exc:
print(f"error: {exc}", file=sys.stderr)
return 2
print(
f"selected validation accuracy: {metrics['bestValidationAccuracy']:.4%}",
flush=True,
)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+252
View File
@@ -0,0 +1,252 @@
#!/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, 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:
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 on Apple Silicon"
) from exc
if not mx.metal.is_available():
raise DataError("verification requires the MLX Metal backend")
# 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)
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)
except (AssertionError, DataError) as exc:
print(f"error: {exc}")
return 2
return 0
if __name__ == "__main__":
raise SystemExit(main())