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()