"""MLX implementation of the ModernBERT ``purpose-deep`` multi-task model. Parameter names and tensor layouts match Hugging Face ModernBERT. The masked-LM backbone and prediction head therefore load directly from the pinned safetensors checkpoint; only the four small task heads start fresh. """ 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 purpose_data import LABELS, DataError @dataclass(frozen=True) class ModernBertPurposeConfig: vocab_size: int hidden_size: int intermediate_size: int num_hidden_layers: int num_attention_heads: int max_position_embeddings: int pad_token_id: int norm_eps: float norm_bias: bool attention_bias: bool attention_dropout: float layer_types: tuple[str, ...] local_attention: int embedding_dropout: float mlp_bias: bool mlp_dropout: float classifier_bias: bool classifier_dropout: float full_rope_theta: float local_rope_theta: float gradient_checkpointing: bool = True @classmethod def from_hugging_face( cls, config: dict[str, Any], *, gradient_checkpointing: bool = True, ) -> "ModernBertPurposeConfig": layer_types = config.get("layer_types") if layer_types is None: every = int(config.get("global_attn_every_n_layers", 3)) layer_types = [ "sliding_attention" if index % every else "full_attention" for index in range(int(config["num_hidden_layers"])) ] rope = config.get("rope_parameters") or {} return cls( vocab_size=int(config["vocab_size"]), hidden_size=int(config["hidden_size"]), intermediate_size=int(config["intermediate_size"]), num_hidden_layers=int(config["num_hidden_layers"]), num_attention_heads=int(config["num_attention_heads"]), max_position_embeddings=int(config["max_position_embeddings"]), pad_token_id=int(config["pad_token_id"]), norm_eps=float(config.get("norm_eps", 1e-5)), norm_bias=bool(config.get("norm_bias", False)), attention_bias=bool(config.get("attention_bias", False)), attention_dropout=float(config.get("attention_dropout", 0.0)), layer_types=tuple(layer_types), local_attention=int(config.get("local_attention", 128)), embedding_dropout=float(config.get("embedding_dropout", 0.0)), mlp_bias=bool(config.get("mlp_bias", False)), mlp_dropout=float(config.get("mlp_dropout", 0.0)), classifier_bias=bool(config.get("classifier_bias", False)), classifier_dropout=float(config.get("classifier_dropout", 0.0)), full_rope_theta=float( (rope.get("full_attention") or {}).get( "rope_theta", config.get("global_rope_theta", 160_000.0) ) ), local_rope_theta=float( (rope.get("sliding_attention") or {}).get( "rope_theta", config.get("local_rope_theta", 10_000.0) ) ), gradient_checkpointing=gradient_checkpointing, ) def _rotate_half(value: Any) -> Any: half = value.shape[-1] // 2 return mx.concatenate((-value[..., half:], value[..., :half]), axis=-1) class ModernBertEmbeddings(nn.Module): def __init__(self, config: ModernBertPurposeConfig) -> None: super().__init__() self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size) self.norm = nn.LayerNorm( config.hidden_size, eps=config.norm_eps, bias=config.norm_bias ) self.drop = nn.Dropout(config.embedding_dropout) def __call__(self, input_ids: Any) -> Any: return self.drop(self.norm(self.tok_embeddings(input_ids))) class ModernBertMLP(nn.Module): def __init__(self, config: ModernBertPurposeConfig) -> None: super().__init__() self.Wi = nn.Linear( config.hidden_size, config.intermediate_size * 2, bias=config.mlp_bias ) self.Wo = nn.Linear( config.intermediate_size, config.hidden_size, bias=config.mlp_bias ) self.drop = nn.Dropout(config.mlp_dropout) self.act = nn.GELU(approx="none") def __call__(self, hidden_states: Any) -> Any: projected = self.Wi(hidden_states) content, gate = mx.split(projected, 2, axis=-1) return self.Wo(self.drop(self.act(content) * gate)) class ModernBertAttention(nn.Module): def __init__( self, config: ModernBertPurposeConfig, attention_type: str ) -> None: super().__init__() if config.hidden_size % config.num_attention_heads: raise DataError("ModernBERT hidden size must divide evenly into heads") self.num_heads = config.num_attention_heads self.head_dim = config.hidden_size // config.num_attention_heads self.attention_type = attention_type self.Wqkv = nn.Linear( config.hidden_size, config.hidden_size * 3, bias=config.attention_bias ) self.Wo = nn.Linear( config.hidden_size, config.hidden_size, bias=config.attention_bias ) self.out_drop = nn.Dropout(config.attention_dropout) def __call__( self, hidden_states: Any, mask: Any | None, cos: Any, sin: Any, ) -> Any: batch, length, hidden = hidden_states.shape qkv = self.Wqkv(hidden_states).reshape( batch, length, 3, self.num_heads, self.head_dim ) query, key, value = ( qkv[:, :, index].transpose(0, 2, 1, 3) for index in range(3) ) cos = cos[None, None, :, :].astype(query.dtype) sin = sin[None, None, :, :].astype(query.dtype) query = query * cos + _rotate_half(query) * sin key = key * cos + _rotate_half(key) * sin context = mx.fast.scaled_dot_product_attention( query, key, value, scale=self.head_dim**-0.5, mask=mask, ) context = context.transpose(0, 2, 1, 3).reshape(batch, length, hidden) return self.out_drop(self.Wo(context)) class ModernBertEncoderLayer(nn.Module): def __init__( self, config: ModernBertPurposeConfig, layer_index: int, ) -> None: super().__init__() self.attention_type = config.layer_types[layer_index] self.attn_norm = ( nn.Identity() if layer_index == 0 else nn.LayerNorm( config.hidden_size, eps=config.norm_eps, bias=config.norm_bias ) ) self.attn = ModernBertAttention(config, self.attention_type) self.mlp_norm = nn.LayerNorm( config.hidden_size, eps=config.norm_eps, bias=config.norm_bias ) self.mlp = ModernBertMLP(config) def __call__(self, hidden_states: Any, mask: Any, cos: Any, sin: Any) -> Any: hidden_states = hidden_states + self.attn( self.attn_norm(hidden_states), mask, cos, sin ) return hidden_states + self.mlp(self.mlp_norm(hidden_states)) class ModernBertModel(nn.Module): def __init__(self, config: ModernBertPurposeConfig) -> None: super().__init__() self.config = config self.embeddings = ModernBertEmbeddings(config) self.layers = [ ModernBertEncoderLayer(config, index) for index in range(config.num_hidden_layers) ] self.final_norm = nn.LayerNorm( config.hidden_size, eps=config.norm_eps, bias=config.norm_bias ) def _rotary(self, length: int, theta: float) -> tuple[Any, Any]: head_dim = self.config.hidden_size // self.config.num_attention_heads inverse_frequency = 1.0 / ( theta ** (mx.arange(0, head_dim, 2).astype(mx.float32) / head_dim) ) frequency = mx.arange(length).astype(mx.float32)[:, None] * inverse_frequency[None] embedding = mx.concatenate((frequency, frequency), axis=-1) return mx.cos(embedding), mx.sin(embedding) def _masks(self, attention_mask: Any, dtype: Any) -> dict[str, Any]: length = attention_mask.shape[1] visible = attention_mask.astype(mx.bool_)[:, None, None, :] zero = mx.array(0.0, dtype=dtype) blocked = mx.array(-1e4, dtype=dtype) full = mx.where(visible, zero, blocked) positions = mx.arange(length) half_window = self.config.local_attention // 2 local_visible = mx.abs(positions[:, None] - positions[None, :]) <= half_window local = mx.where( visible & local_visible[None, None, :, :], zero, blocked ) return {"full_attention": full, "sliding_attention": local} def __call__(self, input_ids: Any, attention_mask: Any | None = None) -> Any: if input_ids.shape[1] > self.config.max_position_embeddings: raise DataError("ModernBERT input exceeds its position contract") if attention_mask is None: attention_mask = mx.ones_like(input_ids) hidden_states = self.embeddings(input_ids) masks = self._masks(attention_mask, hidden_states.dtype) rotary = { "full_attention": self._rotary( input_ids.shape[1], self.config.full_rope_theta ), "sliding_attention": self._rotary( input_ids.shape[1], self.config.local_rope_theta ), } for layer in self.layers: cos, sin = rotary[layer.attention_type] if self.config.gradient_checkpointing and self.training: hidden_states = mx.checkpoint(layer)( hidden_states, masks[layer.attention_type], cos, sin ) else: hidden_states = layer( hidden_states, masks[layer.attention_type], cos, sin ) return self.final_norm(hidden_states) class ModernBertPredictionHead(nn.Module): def __init__(self, config: ModernBertPurposeConfig) -> None: super().__init__() self.dense = nn.Linear( config.hidden_size, config.hidden_size, bias=config.classifier_bias ) self.act = nn.GELU(approx="none") self.norm = nn.LayerNorm( config.hidden_size, eps=config.norm_eps, bias=config.norm_bias ) def __call__(self, hidden_states: Any) -> Any: return self.norm(self.act(self.dense(hidden_states))) class ModernBertForPurposeClassification(nn.Module): def __init__(self, config: ModernBertPurposeConfig) -> None: super().__init__() self.config = config self.model = ModernBertModel(config) # Reuse ModernBERT's pretrained masked-LM prediction head as the shared # representation adapter before the four task-specific readouts. self.head = ModernBertPredictionHead(config) self.drop = nn.Dropout(config.classifier_dropout) self.purpose_classifier = nn.Linear( config.hidden_size, len(LABELS), bias=True ) self.secondary_classifier = nn.Linear( config.hidden_size, len(LABELS), bias=True ) self.mixed_classifier = nn.Linear(config.hidden_size, 1, bias=True) self.difficulty_regressor = nn.Linear(config.hidden_size, 1, bias=True) def __call__(self, input_ids: Any, attention_mask: Any | None = None) -> dict[str, Any]: sequence = self.model(input_ids, attention_mask) pooled = self.drop(self.head(sequence[:, 0])) return { "purpose_logits": self.purpose_classifier(pooled), "secondary_logits": self.secondary_classifier(pooled), "mixed_logits": self.mixed_classifier(pooled).squeeze(-1), "difficulty": mx.sigmoid( self.difficulty_regressor(pooled).squeeze(-1) ), } def load_pretrained_weights( model: ModernBertForPurposeClassification, checkpoint: Path, ) -> dict[str, int]: """Load every pretrained backbone/head tensor and reject partial checkpoints.""" if not checkpoint.is_file(): raise DataError(f"{checkpoint}: ModernBERT safetensors checkpoint is missing") weights = mx.load(str(checkpoint)) available = { key: value for key, value in weights.items() if key.startswith("model.") or key.startswith("head.") } parameters = dict(tree_flatten(model.parameters())) expected = { key for key in parameters if key.startswith("model.") or key.startswith("head.") } missing = sorted(expected - set(available)) if missing: preview = ", ".join(missing[:5]) raise DataError( f"ModernBERT checkpoint is missing {len(missing)} required tensors: {preview}" ) for key in expected: if tuple(available[key].shape) != tuple(parameters[key].shape): raise DataError( f"ModernBERT tensor {key} has shape {available[key].shape}; " f"expected {parameters[key].shape}" ) model.load_weights(list(available.items()), strict=False) mx.eval(model.parameters()) return { "loaded": len(available), "ignored": len(weights) - len(available), "freshTaskHeads": len(parameters) - len(expected), } def save_weights(model: ModernBertForPurposeClassification, checkpoint: Path) -> None: mx.eval(model.parameters()) mx.save_safetensors( str(checkpoint), dict(tree_flatten(model.parameters())), metadata={"format": "pt"}, )