From 90092332dbbf545c9ec74c253fd6f1281c6d5cdd Mon Sep 17 00:00:00 2001 From: Nucleic Date: Thu, 30 Jul 2026 20:14:37 -0700 Subject: [PATCH] Merge nucleic/sleek-ember-seal-uady into dev --- README.md | 43 ++ mlx_checkpoint.py | 39 ++ mlx_model.py | 354 +++++++++++++++ requirements-mlx.txt | 2 + tests/test_mlx_checkpoint.py | 57 +++ tests/test_train_mlx.py | 72 +++ train_mlx.py | 849 +++++++++++++++++++++++++++++++++++ verify_mlx.py | 252 +++++++++++ 8 files changed, 1668 insertions(+) create mode 100644 mlx_checkpoint.py create mode 100644 mlx_model.py create mode 100644 requirements-mlx.txt create mode 100644 tests/test_mlx_checkpoint.py create mode 100644 tests/test_train_mlx.py create mode 100644 train_mlx.py create mode 100644 verify_mlx.py diff --git a/README.md b/README.md index 5614468..f2b6c3f 100644 --- a/README.md +++ b/README.md @@ -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 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: ```bash diff --git a/mlx_checkpoint.py b/mlx_checkpoint.py new file mode 100644 index 0000000..35b78bb --- /dev/null +++ b/mlx_checkpoint.py @@ -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 + diff --git a/mlx_model.py b/mlx_model.py new file mode 100644 index 0000000..798d385 --- /dev/null +++ b/mlx_model.py @@ -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"}, + ) diff --git a/requirements-mlx.txt b/requirements-mlx.txt new file mode 100644 index 0000000..3a761cd --- /dev/null +++ b/requirements-mlx.txt @@ -0,0 +1,2 @@ +-r requirements.txt +mlx==0.32.0 diff --git a/tests/test_mlx_checkpoint.py b/tests/test_mlx_checkpoint.py new file mode 100644 index 0000000..e6ee2b3 --- /dev/null +++ b/tests/test_mlx_checkpoint.py @@ -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() diff --git a/tests/test_train_mlx.py b/tests/test_train_mlx.py new file mode 100644 index 0000000..61360a3 --- /dev/null +++ b/tests/test_train_mlx.py @@ -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() diff --git a/train_mlx.py b/train_mlx.py new file mode 100644 index 0000000..f8496b0 --- /dev/null +++ b/train_mlx.py @@ -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()) diff --git a/verify_mlx.py b/verify_mlx.py new file mode 100644 index 0000000..f8a11b8 --- /dev/null +++ b/verify_mlx.py @@ -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())