"""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"}, )