Files

402 lines
15 KiB
Python

"""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 load_checkpoint_weights(
model: ModernBertForPurposeClassification,
checkpoint: Path,
) -> dict[str, int]:
"""Strictly restore a trained purpose-deep checkpoint, including task heads."""
if not checkpoint.is_file():
raise DataError(f"{checkpoint}: purpose-deep checkpoint is missing")
weights = mx.load(str(checkpoint))
parameters = dict(tree_flatten(model.parameters()))
missing = sorted(set(parameters) - set(weights))
unexpected = sorted(set(weights) - set(parameters))
if missing or unexpected:
details = []
if missing:
details.append(
f"missing {len(missing)} tensors ({', '.join(missing[:3])})"
)
if unexpected:
details.append(
f"has {len(unexpected)} unexpected tensors "
f"({', '.join(unexpected[:3])})"
)
raise DataError(
f"{checkpoint}: trained checkpoint " + " and ".join(details)
)
for key, parameter in parameters.items():
if tuple(weights[key].shape) != tuple(parameter.shape):
raise DataError(
f"purpose-deep tensor {key} has shape {weights[key].shape}; "
f"expected {parameter.shape}"
)
model.load_weights(list(weights.items()), strict=True)
mx.eval(model.parameters())
return {
"loaded": len(weights),
"ignored": 0,
"freshTaskHeads": 0,
}
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"},
)