Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -0,0 +1,360 @@
|
||||
"""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"},
|
||||
)
|
||||
Reference in New Issue
Block a user