Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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"},
|
||||||
|
)
|
||||||
@@ -0,0 +1,2 @@
|
|||||||
|
-r requirements.txt
|
||||||
|
mlx==0.32.0
|
||||||
@@ -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()
|
||||||
@@ -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
@@ -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
@@ -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())
|
||||||
Reference in New Issue
Block a user