needle-mlx / model.py
Blankyy's picture
Upload folder using huggingface_hub
3d5dbbb verified
Raw
History Blame Contribute Delete
24.9 kB
import json
import math
from dataclasses import dataclass
from pathlib import Path
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_flatten
def default_init():
return nn.init.normal(std=0.02)
def residual_init(num_layers):
return nn.init.normal(std=0.02 / math.sqrt(2 * num_layers))
DTYPE_MAP = {
"float32": mx.float32,
"bfloat16": mx.bfloat16,
"float16": mx.float16,
}
def _map_hf_weight(name):
if name == "model.embed_tokens.weight":
return "embedding.weight"
if name == "model.encoder.final_norm.weight":
return "encoder.final_norm.scale"
if name == "model.decoder.norm.weight":
return "decoder.final_norm.scale"
if name == "lm_head.weight":
return None
if name.startswith("model.encoder.layers."):
name = name.removeprefix("model.")
name = name.replace(".input_layernorm.weight", ".norm.scale")
elif name.startswith("model.decoder.layers."):
name = name.removeprefix("model.")
name = name.replace(".input_layernorm.weight", ".self_norm.scale")
name = name.replace(
".encoder_attn_layer_norm.weight", ".cross_norm.scale"
)
name = name.replace(".encoder_attn.", ".cross_attn.")
else:
raise ValueError(f"Unsupported checkpoint tensor: {name}")
return name.replace("_norm.weight", "_norm.scale")
def _linear(in_dims, out_dims, dtype, *, bias=False, init=None):
layer = nn.Linear(in_dims, out_dims, bias=bias)
layer.weight = (init or default_init())(layer.weight)
if bias:
layer.bias = mx.zeros_like(layer.bias)
layer.set_dtype(dtype)
return layer
def _dropout(x, rate, deterministic):
if deterministic or rate == 0:
return x
keep = 1 - rate
return mx.random.bernoulli(keep, x.shape) * x / keep
class ZCRMSNorm(nn.Module):
def __init__(self, dims, dtype=mx.bfloat16, epsilon=1e-6):
super().__init__()
self.scale = mx.zeros((dims,))
self.dtype = dtype
self.epsilon = epsilon
def __call__(self, x):
rms = mx.sqrt(
mx.mean(x.astype(mx.float32) ** 2, axis=-1, keepdims=True)
+ self.epsilon
)
return ((1 + self.scale) * x / rms).astype(self.dtype)
@dataclass
class TransformerConfig:
vocab_size: int = 8192
d_model: int = 128
num_heads: int = 4
num_kv_heads: int = 2
num_encoder_layers: int = 2
num_decoder_layers: int = 2
d_ff: int = 512
max_seq_len: int = 128
pad_token_id: int = 0
rope_theta: float = 10_000.0
dtype: str = "bfloat16"
activation: str = "drelu"
num_memory_slots: int = 64
dropout_rate: float = 0.1
contrastive_dim: int = 128
no_feedforward: bool = True
def __init__(self, **kwargs):
valid = self.__dataclass_fields__
for key, value in kwargs.items():
if key in valid:
setattr(self, key, value)
@property
def mlx_dtype(self):
return DTYPE_MAP[self.dtype]
@property
def total_layers(self):
return self.num_encoder_layers + self.num_decoder_layers
def precompute_rope_freqs(head_dim, seq_len, theta=10_000.0):
freqs = 1 / theta ** (mx.arange(0, head_dim, 2, dtype=mx.float32) / head_dim)
angles = mx.outer(mx.arange(seq_len, dtype=mx.float32), freqs)
return mx.cos(angles), mx.sin(angles)
def apply_rope(x, cos, sin):
length = x.shape[2]
half = x.shape[-1] // 2
cos = cos[:length][None, None, :, :]
sin = sin[:length][None, None, :, :]
x1, x2 = x[..., :half], x[..., half:]
return mx.concatenate((x1 * cos - x2 * sin, x2 * cos + x1 * sin), axis=-1)
class MultiHeadAttention(nn.Module):
def __init__(
self,
num_heads,
num_kv_heads,
d_model,
num_layers,
dtype=mx.bfloat16,
rope_keys_only=False,
):
super().__init__()
if d_model % num_heads or num_heads % num_kv_heads:
raise ValueError(
"d_model must be divisible by num_heads, and num_heads by num_kv_heads"
)
self.num_heads = num_heads
self.num_kv_heads = num_kv_heads
self.d_model = d_model
self.head_dim = d_model // num_heads
self.repeats = num_heads // num_kv_heads
self.dtype = dtype
self.rope_keys_only = rope_keys_only
kv_dim = num_kv_heads * self.head_dim
self.q_proj = _linear(d_model, d_model, dtype)
self.k_proj = _linear(d_model, kv_dim, dtype)
self.v_proj = _linear(d_model, kv_dim, dtype)
self.out_proj = _linear(
d_model, d_model, dtype, init=residual_init(num_layers)
)
self.q_norm = ZCRMSNorm(self.head_dim, dtype)
self.k_norm = ZCRMSNorm(self.head_dim, dtype)
def __call__(self, q_input, kv_input, mask=None, rope=None):
batch, q_len, _ = q_input.shape
kv_len = kv_input.shape[1]
q = self.q_proj(q_input.astype(self.dtype))
k = self.k_proj(kv_input.astype(self.dtype))
v = self.v_proj(kv_input.astype(self.dtype))
q = q.reshape(batch, q_len, self.num_heads, self.head_dim).transpose(
0, 2, 1, 3
)
k = k.reshape(batch, kv_len, self.num_kv_heads, self.head_dim).transpose(
0, 2, 1, 3
)
v = v.reshape(batch, kv_len, self.num_kv_heads, self.head_dim).transpose(
0, 2, 1, 3
)
q = self.q_norm(q)
k = self.k_norm(k)
if self.repeats > 1:
k = mx.repeat(k, self.repeats, axis=1)
v = mx.repeat(v, self.repeats, axis=1)
if rope is not None:
cos, sin = rope
if not self.rope_keys_only:
q = apply_rope(q, cos, sin)
k = apply_rope(k, cos, sin)
scale = mx.sqrt(mx.array(self.head_dim, dtype=mx.float32))
weights = mx.matmul(q, k.transpose(0, 1, 3, 2)) / scale
if mask is not None:
weights = mx.where(mask, weights, mx.finfo(weights.dtype).min)
weights = nn.softmax(weights, axis=-1)
out = mx.matmul(weights, v)
out = out.transpose(0, 2, 1, 3).reshape(batch, q_len, self.d_model)
return self.out_proj(out.astype(self.dtype))
class FeedForward(nn.Module):
def __init__(
self, d_model, d_ff, num_layers, dtype=mx.bfloat16, activation="drelu"
):
super().__init__()
self.dtype = dtype
self.activation = activation
self.gate_proj = _linear(d_model, d_ff, dtype)
self.up_proj = _linear(d_model, d_ff, dtype)
self.down_proj = _linear(
d_ff, d_model, dtype, init=residual_init(num_layers)
)
def __call__(self, x, ffn_mask=None):
gate = self.gate_proj(x.astype(self.dtype))
up = self.up_proj(x.astype(self.dtype))
if self.activation == "swiglu":
h = nn.silu(gate) * up
elif self.activation == "geglu":
h = nn.gelu_approx(gate) * up
else:
h = nn.relu(gate) * nn.relu(up)
if ffn_mask is not None:
h = h * ffn_mask[:, None, :]
return self.down_proj(h.astype(self.dtype))
class EncoderBlock(nn.Module):
def __init__(
self,
num_heads,
num_kv_heads,
d_model,
d_ff,
num_layers,
dtype=mx.bfloat16,
activation="drelu",
dropout_rate=0.0,
no_feedforward=True,
):
super().__init__()
self.dtype = dtype
self.dropout_rate = dropout_rate
self.no_feedforward = no_feedforward
self.attn_gate = mx.zeros(())
self.norm = ZCRMSNorm(d_model, dtype)
self.self_attn = MultiHeadAttention(
num_heads, num_kv_heads, d_model, num_layers, dtype
)
if not no_feedforward:
self.ffn_gate = mx.zeros(())
self.ffn_norm = ZCRMSNorm(d_model, dtype)
self.ffn = FeedForward(
d_model, d_ff, num_layers, dtype, activation
)
def __call__(
self, x, mask=None, rope=None, ffn_mask=None, deterministic=True
):
residual = x
normed = self.norm(x)
x = self.self_attn(normed, normed, mask=mask, rope=rope)
gate = nn.sigmoid(self.attn_gate).astype(self.dtype)
x = residual + gate * _dropout(x, self.dropout_rate, deterministic)
if not self.no_feedforward:
residual = x
x = self.ffn(self.ffn_norm(x), ffn_mask=ffn_mask)
gate = nn.sigmoid(self.ffn_gate).astype(self.dtype)
x = residual + gate * _dropout(x, self.dropout_rate, deterministic)
return x
class Encoder(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.layers = [
EncoderBlock(
config.num_heads,
config.num_kv_heads,
config.d_model,
config.d_ff,
config.total_layers,
config.mlx_dtype,
config.activation,
config.dropout_rate,
config.no_feedforward,
)
for _ in range(config.num_encoder_layers)
]
self.final_norm = ZCRMSNorm(config.d_model, config.mlx_dtype)
def __call__(
self, x, mask=None, rope=None, ffn_mask=None, deterministic=True
):
x = x.astype(self.config.mlx_dtype)
for layer in self.layers:
x = layer(x, mask, rope, ffn_mask, deterministic)
return self.final_norm(x), mask
class DecoderBlock(nn.Module):
def __init__(
self,
num_heads,
num_kv_heads,
d_model,
d_ff,
num_layers,
dtype=mx.bfloat16,
activation="drelu",
dropout_rate=0.0,
no_feedforward=True,
):
super().__init__()
self.dtype = dtype
self.dropout_rate = dropout_rate
self.no_feedforward = no_feedforward
self.self_attn_gate = mx.zeros(())
self.self_norm = ZCRMSNorm(d_model, dtype)
self.self_attn = MultiHeadAttention(
num_heads, num_kv_heads, d_model, num_layers, dtype
)
self.cross_attn_gate = mx.zeros(())
self.cross_norm = ZCRMSNorm(d_model, dtype)
self.cross_attn = MultiHeadAttention(
num_heads, num_kv_heads, d_model, num_layers, dtype
)
if not no_feedforward:
self.ffn_gate = mx.zeros(())
self.ffn_norm = ZCRMSNorm(d_model, dtype)
self.ffn = FeedForward(
d_model, d_ff, num_layers, dtype, activation
)
def __call__(
self,
x,
encoder_out,
self_mask=None,
cross_mask=None,
rope=None,
ffn_mask=None,
deterministic=True,
):
residual = x
normed = self.self_norm(x)
x = self.self_attn(normed, normed, mask=self_mask, rope=rope)
gate = nn.sigmoid(self.self_attn_gate).astype(self.dtype)
x = residual + gate * _dropout(x, self.dropout_rate, deterministic)
residual = x
x = self.cross_attn(
self.cross_norm(x), encoder_out, mask=cross_mask
)
gate = nn.sigmoid(self.cross_attn_gate).astype(self.dtype)
x = residual + gate * _dropout(x, self.dropout_rate, deterministic)
if not self.no_feedforward:
residual = x
x = self.ffn(self.ffn_norm(x), ffn_mask=ffn_mask)
gate = nn.sigmoid(self.ffn_gate).astype(self.dtype)
x = residual + gate * _dropout(x, self.dropout_rate, deterministic)
return x
class Decoder(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.layers = [
DecoderBlock(
config.num_heads,
config.num_kv_heads,
config.d_model,
config.d_ff,
config.total_layers,
config.mlx_dtype,
config.activation,
config.dropout_rate,
config.no_feedforward,
)
for _ in range(config.num_decoder_layers)
]
self.final_norm = ZCRMSNorm(config.d_model, config.mlx_dtype)
def __call__(
self,
x,
encoder_out,
self_mask=None,
cross_mask=None,
rope=None,
ffn_mask=None,
deterministic=True,
):
x = x.astype(self.config.mlx_dtype)
for layer in self.layers:
x = layer(
x,
encoder_out,
self_mask,
cross_mask,
rope,
ffn_mask,
deterministic,
)
return self.final_norm(x)
class SimpleAttentionNetwork(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
self.embedding = nn.Embedding(config.vocab_size, config.d_model)
self.embedding.weight = default_init()(self.embedding.weight)
self.embed_scale = math.sqrt(config.d_model)
self.encoder = Encoder(config)
self.decoder = Decoder(config)
self.contrastive_hidden = _linear(
config.d_model,
config.d_model // 4,
config.mlx_dtype,
bias=True,
)
self.contrastive_proj = _linear(
config.d_model // 4, config.contrastive_dim, config.mlx_dtype
)
self.log_temp = mx.zeros(())
@classmethod
def from_pretrained(cls, model_dir="weights"):
model_dir = Path(model_dir)
with (model_dir / "config.json").open() as file:
hf_config = json.load(file)
config = TransformerConfig(
vocab_size=hf_config["vocab_size"],
d_model=hf_config["d_model"],
num_heads=hf_config["num_heads"],
num_kv_heads=hf_config["num_kv_heads"],
num_encoder_layers=hf_config["num_encoder_layers"],
num_decoder_layers=hf_config["num_decoder_layers"],
pad_token_id=hf_config["pad_token_id"],
rope_theta=hf_config["rope_theta"],
dtype=hf_config.get("torch_dtype", "bfloat16"),
no_feedforward=True,
)
model = cls(config)
checkpoint = mx.load(str(model_dir / "model.safetensors"))
if "model.embed_tokens.weight" in checkpoint:
if not mx.array_equal(
checkpoint["model.embed_tokens.weight"],
checkpoint["lm_head.weight"],
).item():
raise ValueError("Checkpoint input and output embeddings are not tied")
mapped = []
for name, value in checkpoint.items():
target = _map_hf_weight(name)
if target is not None:
if target.endswith("_gate"):
value = value.reshape(())
mapped.append((target, value))
else:
mapped = list(checkpoint.items())
model.load_weights(mapped, strict=False)
mx.eval(model.parameters())
return model
def generate(self, src, max_new_tokens=512):
"""Greedy generation. Decoder starts with EOS, as in the JAX runtime."""
if src.shape[0] != 1:
raise ValueError("generate currently supports batch size 1")
encoder_out, enc_mask = self.encode_text(
src, src_mask=make_padding_mask(src, self.config.pad_token_id)
)
tokens = mx.full(
(1, max_new_tokens + 1),
self.config.pad_token_id,
dtype=mx.int32,
)
tokens[0, 0] = 1
causal = make_causal_mask(tokens.shape[1])
generated = []
# ponytail: full-buffer decode matches upstream; add a KV cache if latency matters.
for position in range(max_new_tokens):
logits = self.decode(
tokens,
encoder_out,
self_mask=causal,
cross_mask=enc_mask,
)
token = int(mx.argmax(logits[0, position]).item())
if token == 1:
break
generated.append(token)
tokens[0, position + 1] = token
return generated
def _rope(self, seq_len):
head_dim = self.config.d_model // self.config.num_heads
return precompute_rope_freqs(
head_dim, seq_len, self.config.rope_theta
)
def encode_text(
self, src, src_mask=None, ffn_mask=None, deterministic=True
):
x = self.embedding(src) * self.embed_scale
return self.encoder(
x,
mask=src_mask,
rope=self._rope(src.shape[1]),
ffn_mask=ffn_mask,
deterministic=deterministic,
)
def encode(self, src, src_mask=None):
return self.encode_text(src, src_mask=src_mask)
def decode(
self,
tgt,
encoder_out,
self_mask=None,
cross_mask=None,
deterministic=True,
):
x = self.embedding(tgt) * self.embed_scale
x = self.decoder(
x,
encoder_out,
self_mask=self_mask,
cross_mask=cross_mask,
rope=self._rope(tgt.shape[1]),
deterministic=deterministic,
)
return self.embedding.as_linear(x.astype(mx.float32))
def _mean_pool(self, encoder_out, enc_mask):
if enc_mask is not None:
mask_2d = enc_mask[:, 0, 0, :]
else:
mask_2d = mx.ones(encoder_out.shape[:2], dtype=encoder_out.dtype)
mask_3d = mask_2d[:, :, None].astype(encoder_out.dtype)
summed = mx.sum(encoder_out * mask_3d, axis=1)
counts = mx.maximum(mx.sum(mask_2d, axis=1, keepdims=True), 1.0)
return summed / counts
def encode_contrastive(self, tokens, deterministic=True):
src_mask = make_padding_mask(tokens, self.config.pad_token_id)
encoder_out, enc_mask = self.encode_text(
tokens, src_mask=src_mask, deterministic=deterministic
)
pooled = self._mean_pool(encoder_out, enc_mask)
projected = self.contrastive_proj(
nn.relu(self.contrastive_hidden(pooled))
)
denom = mx.sqrt(
mx.sum(
projected.astype(mx.float32) ** 2,
axis=-1,
keepdims=True,
)
+ 1e-12
)
return projected / denom.astype(projected.dtype)
def forward_contrastive(
self, query_tokens, tool_tokens, deterministic=True
):
q_emb = self.encode_contrastive(query_tokens, deterministic)
t_emb = self.encode_contrastive(tool_tokens, deterministic)
return q_emb, t_emb, self.log_temp
def __call__(
self, src, tgt, src_mask=None, tgt_mask=None, cross_mask=None
):
encoder_out, enc_mask = self.encode_text(src, src_mask=src_mask)
return self.decode(
tgt,
encoder_out,
self_mask=tgt_mask,
cross_mask=cross_mask if cross_mask is not None else enc_mask,
)
def _run_decoder(
self,
encoder_out,
tgt,
tgt_mask=None,
cross_mask=None,
ffn_mask=None,
deterministic=True,
):
x = self.embedding(tgt) * self.embed_scale
x = self.decoder(
x,
encoder_out,
self_mask=tgt_mask,
cross_mask=cross_mask,
rope=self._rope(tgt.shape[1]),
ffn_mask=ffn_mask,
deterministic=deterministic,
)
return x.astype(mx.float32)
def forward_masked(
self,
src,
tgt,
src_mask=None,
tgt_mask=None,
cross_mask=None,
ffn_mask=None,
deterministic=True,
):
encoder_out, enc_mask = self.encode_text(
src,
src_mask=src_mask,
ffn_mask=ffn_mask,
deterministic=deterministic,
)
x = self._run_decoder(
encoder_out,
tgt,
tgt_mask=tgt_mask,
cross_mask=cross_mask if cross_mask is not None else enc_mask,
ffn_mask=ffn_mask,
deterministic=deterministic,
)
return self.embedding.as_linear(x), 0.0
def _make_eval_ffn_mask(self, ff_width, batch, dtype):
mask = (mx.arange(self.config.d_ff) < ff_width).astype(dtype)
return mx.broadcast_to(mask[None, :], (batch, self.config.d_ff))
def forward_with_aux(
self,
src,
tgt,
src_mask=None,
tgt_mask=None,
cross_mask=None,
mat_ff_widths=None,
):
batch = src.shape[0]
encoder_out, enc_mask = self.encode_text(src, src_mask=src_mask)
cm = cross_mask if cross_mask is not None else enc_mask
x = self._run_decoder(
encoder_out, tgt, tgt_mask=tgt_mask, cross_mask=cm
)
logits = self.embedding.as_linear(x)
mat_logits = []
if mat_ff_widths is not None:
for width in mat_ff_widths:
mask = self._make_eval_ffn_mask(width, batch, x.dtype)
enc_m, _ = self.encode_text(
src, src_mask=src_mask, ffn_mask=mask
)
x_m = self._run_decoder(
enc_m,
tgt,
tgt_mask=tgt_mask,
cross_mask=cm,
ffn_mask=mask,
)
mat_logits.append(self.embedding.as_linear(x_m))
return logits, 0.0, mat_logits
def init_all(self, src, tgt):
src_mask = make_padding_mask(src, self.config.pad_token_id)
tgt_mask = make_causal_mask(tgt.shape[1]) & make_padding_mask(
tgt, self.config.pad_token_id
)
encoder_out, enc_mask = self.encode_text(src, src_mask=src_mask)
self._run_decoder(
encoder_out, tgt, tgt_mask=tgt_mask, cross_mask=enc_mask
)
self.encode_contrastive(src)
return mx.zeros(())
def make_causal_mask(seq_len):
return mx.tril(mx.ones((seq_len, seq_len), dtype=mx.bool_))[None, None]
def make_padding_mask(tokens, pad_token_id):
return (tokens != pad_token_id)[:, None, None, :]
def make_packing_mask(seg_ids):
mask = (seg_ids[:, :, None] == seg_ids[:, None, :]) & (
seg_ids[:, :, None] > 0
)
return mask[:, None, :, :]
def make_causal_packing_mask(seg_ids):
length = seg_ids.shape[1]
causal = mx.tril(mx.ones((length, length), dtype=mx.bool_))
block = (seg_ids[:, :, None] == seg_ids[:, None, :]) & (
seg_ids[:, :, None] > 0
)
return (block & causal[None, :, :])[:, None, :, :]
def make_cross_packing_mask(enc_seg_ids, dec_seg_ids):
mask = (dec_seg_ids[:, :, None] == enc_seg_ids[:, None, :]) & (
dec_seg_ids[:, :, None] > 0
)
return mask[:, None, :, :]
def _demo():
config = TransformerConfig(
vocab_size=32,
d_model=16,
num_heads=4,
num_kv_heads=2,
num_encoder_layers=1,
num_decoder_layers=1,
d_ff=32,
contrastive_dim=8,
dropout_rate=0.0,
no_feedforward=False,
dtype="float32",
)
model = SimpleAttentionNetwork(config)
src = mx.array([[1, 2, 3, 0], [4, 5, 0, 0]])
tgt = mx.array([[1, 6, 7, 0], [1, 8, 0, 0]])
src_mask = make_padding_mask(src, config.pad_token_id)
tgt_mask = make_causal_mask(tgt.shape[1]) & make_padding_mask(
tgt, config.pad_token_id
)
logits = model(src, tgt, src_mask, tgt_mask)
q_emb, t_emb, _ = model.forward_contrastive(src, tgt)
aux_logits, _, widths = model.forward_with_aux(
src, tgt, src_mask, tgt_mask, mat_ff_widths=[16]
)
mx.eval(logits, q_emb, t_emb, aux_logits, widths[0])
assert logits.shape == (2, 4, 32)
assert q_emb.shape == t_emb.shape == (2, 8)
assert aux_logits.shape == widths[0].shape == logits.shape
assert mx.all(mx.isfinite(logits)).item()
def _test_pretrained(model_dir):
model = SimpleAttentionNetwork.from_pretrained(model_dir)
src = mx.array([[2, 42, 314, 1]])
tgt = mx.array([[1, 27, 1]])
logits = model(
src,
tgt,
make_padding_mask(src, model.config.pad_token_id),
make_causal_mask(tgt.shape[1]),
)
mx.eval(logits)
assert logits.shape == (1, 3, model.config.vocab_size)
assert mx.all(mx.isfinite(logits)).item()
count = sum(value.size for _, value in tree_flatten(model.parameters()))
print(f"model parameters: {count:,}")
print(f"logits: shape={logits.shape}, dtype={logits.dtype}")
if __name__ == "__main__":
import sys
_test_pretrained(sys.argv[1]) if len(sys.argv) > 1 else _demo()