#!/usr/bin/env python3 """MLX-only transformer model for small causal language modeling.""" from __future__ import annotations import math from dataclasses import dataclass from typing import List, Optional, Tuple import mlx.core as mx import mlx.nn as nn import mlx.nn.losses as losses @dataclass class TransformerConfig: vocab_size: int max_seq_len: int = 1024 d_model: int = 768 n_heads: int = 12 n_kv_heads: Optional[int] = None n_layers: int = 12 mlp_ratio: float = 4.0 mlp_multiple_of: int = 256 rope_base: float = 10000.0 bias: bool = False qk_norm: bool = True attention_impl: str = "fast" loss_dtype: str = "float32" ce_impl: str = "reference" ffn_impl: str = "reference" def __post_init__(self) -> None: if self.vocab_size <= 0: raise ValueError("vocab_size must be > 0") if self.max_seq_len <= 0: raise ValueError("max_seq_len must be > 0") if self.d_model <= 0: raise ValueError("d_model must be > 0") if self.n_heads <= 0 or self.d_model % self.n_heads != 0: raise ValueError("n_heads must divide d_model") if self.n_kv_heads is None: self.n_kv_heads = self.n_heads if self.n_kv_heads <= 0 or self.n_heads % self.n_kv_heads != 0: raise ValueError("n_kv_heads must be positive and divide n_heads") if self.n_layers <= 0: raise ValueError("n_layers must be > 0") if self.mlp_ratio <= 1.0: raise ValueError("mlp_ratio must be > 1.0") if self.mlp_multiple_of <= 0: raise ValueError("mlp_multiple_of must be > 0") if self.attention_impl not in {"fast", "vanilla"}: raise ValueError("attention_impl must be one of: fast, vanilla") if self.loss_dtype not in {"float32", "model"}: raise ValueError("loss_dtype must be one of: float32, model") if self.ce_impl not in {"reference", "metal"}: raise ValueError("ce_impl must be reference or metal") if self.ce_impl == "metal" and self.loss_dtype != "float32": raise ValueError("Metal CE uses FP32 reductions") if self.ffn_impl not in {"reference", "packed-metal"}: raise ValueError("ffn_impl must be reference or packed-metal") def _tree_leaves(tree): if isinstance(tree, dict): for v in tree.values(): yield from _tree_leaves(v) elif isinstance(tree, (list, tuple)): for v in tree: yield from _tree_leaves(v) else: yield tree def count_parameters(model: nn.Module) -> int: total = 0 for leaf in _tree_leaves(model.parameters()): if isinstance(leaf, mx.array): n = 1 for d in leaf.shape: n *= int(d) total += n return total class CausalSelfAttention(nn.Module): def __init__(self, cfg: TransformerConfig): super().__init__() self.n_heads = cfg.n_heads self.n_kv_heads = cfg.n_kv_heads self.head_dim = cfg.d_model // cfg.n_heads self.q_dim = cfg.n_heads * self.head_dim self.kv_dim = cfg.n_kv_heads * self.head_dim self.max_seq_len = cfg.max_seq_len self.qkv_proj = nn.Linear(cfg.d_model, self.q_dim + 2 * self.kv_dim, bias=cfg.bias) self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=cfg.bias) self.q_norm = nn.RMSNorm(self.head_dim) if cfg.qk_norm else None self.k_norm = nn.RMSNorm(self.head_dim) if cfg.qk_norm else None self.rope = nn.RoPE(self.head_dim, base=cfg.rope_base) self.attention_impl = cfg.attention_impl def __call__( self, x: mx.array, cache: Optional[Tuple[mx.array, mx.array]] = None, ) -> Tuple[mx.array, Tuple[mx.array, mx.array]]: bsz, seqlen, d_model = x.shape qkv = self.qkv_proj(x) q, k, v = mx.split(qkv, [self.q_dim, self.q_dim + self.kv_dim], axis=-1) def split_heads(t: mx.array, n_heads: int) -> mx.array: t = t.reshape(bsz, seqlen, n_heads, self.head_dim) return t.transpose(0, 2, 1, 3) q = split_heads(q, self.n_heads) k = split_heads(k, self.n_kv_heads) v = split_heads(v, self.n_kv_heads) if self.q_norm is not None: q = self.q_norm(q) k = self.k_norm(k) cache_offset = 0 if cache is not None and cache[0] is not None: cache_offset = int(cache[0].shape[2]) q = self.rope(q, offset=cache_offset) k = self.rope(k, offset=cache_offset) if cache is not None: k_cache, v_cache = cache if k_cache is not None and v_cache is not None: k = mx.concatenate([k_cache, k], axis=2) v = mx.concatenate([v_cache, v], axis=2) q_len = q.shape[2] k_len = k.shape[2] if q_len > self.max_seq_len or k_len > self.max_seq_len: raise ValueError( f"Sequence length {q_len}/{k_len} exceeds max_seq_len={self.max_seq_len}" ) scale = 1.0 / math.sqrt(self.head_dim) # Cache compact KV heads, not the expanded heads used by vanilla GQA. next_cache = (k, v) if self.attention_impl == "fast": attn = mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, mask="causal", ) else: repeats = self.n_heads // self.n_kv_heads if repeats > 1: k = mx.repeat(k, repeats, axis=1) v = mx.repeat(v, repeats, axis=1) full_mask = nn.MultiHeadAttention.create_additive_causal_mask(k_len) mask = full_mask[k_len - q_len : k_len, :k_len] mask = mx.expand_dims(mx.expand_dims(mask, axis=0), axis=0).astype(q.dtype) scores = mx.matmul(q, k.transpose(0, 1, 3, 2)) * scale scores = scores + mask probs = mx.softmax(scores.astype(mx.float32), axis=-1).astype(v.dtype) attn = mx.matmul(probs, v) out = attn.transpose(0, 2, 1, 3).reshape(bsz, q_len, d_model) out = self.proj(out) return out, next_cache class FeedForward(nn.Module): def __init__(self, cfg: TransformerConfig): super().__init__() # SwiGLU uses three projections, so 2/3 keeps parameter count comparable # to a conventional 4x two-projection MLP. hidden_dim = int((2.0 / 3.0) * cfg.d_model * cfg.mlp_ratio) hidden_dim = cfg.mlp_multiple_of * math.ceil(hidden_dim / cfg.mlp_multiple_of) self.hidden_dim = hidden_dim self.ffn_impl = cfg.ffn_impl self.gate_up = nn.Linear(cfg.d_model, 2 * hidden_dim, bias=cfg.bias) self.down = nn.Linear(hidden_dim, cfg.d_model, bias=cfg.bias) def __call__(self, x: mx.array) -> mx.array: gate_up = self.gate_up(x) if self.ffn_impl == "packed-metal": try: from .fast_swiglu import packed_swiglu except ImportError: from fast_swiglu import packed_swiglu return self.down(packed_swiglu(gate_up)) gate, up = mx.split(gate_up, [self.hidden_dim], axis=-1) return self.down(nn.silu(gate) * up) class TransformerBlock(nn.Module): def __init__(self, cfg: TransformerConfig): super().__init__() self.norm1 = nn.RMSNorm(cfg.d_model) self.attn = CausalSelfAttention(cfg) self.norm2 = nn.RMSNorm(cfg.d_model) self.ffn = FeedForward(cfg) def __call__( self, x: mx.array, cache: Optional[Tuple[mx.array, mx.array]] = None, ) -> Tuple[mx.array, Tuple[mx.array, mx.array]]: attn_out, new_cache = self.attn(self.norm1(x), cache=cache) x = x + attn_out x = x + self.ffn(self.norm2(x)) return x, new_cache class TransformerLM(nn.Module): def __init__(self, cfg: TransformerConfig): super().__init__() self.cfg = cfg self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model) self.blocks = [TransformerBlock(cfg) for _ in range(cfg.n_layers)] # MLX >= 0.32.2 supplies the memory-reduced Metal RMSNorm backward # for every nn.RMSNorm here, including block and Q/K norms. Keep the # native primitive so compilation/autodiff select that implementation. self.norm = nn.RMSNorm(cfg.d_model) def logits( self, input_ids: mx.array, caches: Optional[List[Tuple[mx.array, mx.array]]] = None, ) -> Tuple[mx.array, List[Tuple[mx.array, mx.array]]]: x = self.embed(input_ids) if caches is None: caches = [None] * len(self.blocks) new_caches = [] for block, block_cache in zip(self.blocks, caches): x, next_cache = block(x, cache=block_cache) new_caches.append(next_cache) x = self.norm(x) weight = self.embed.weight logits = mx.matmul(x, weight.transpose(1, 0)) return logits, new_caches def __call__( self, input_ids: mx.array, targets: Optional[mx.array] = None, ignore_index: int = -100, ) -> dict: logits, _ = self.logits(input_ids, caches=None) out = {"logits": logits} if targets is None: return out loss_logits = logits.astype(mx.float32) if self.cfg.loss_dtype == "float32" else logits mask = targets != ignore_index safe_targets = mx.where(mask, targets, mx.zeros_like(targets)) if self.cfg.ce_impl == "metal": try: from .fast_loss import cross_entropy except ImportError: from fast_loss import cross_entropy per_token = cross_entropy(logits, safe_targets) else: per_token = losses.cross_entropy(loss_logits, safe_targets, reduction="none") mask = mask.astype(per_token.dtype) denom = mx.maximum(mask.sum(), mx.array(1.0, dtype=per_token.dtype)) loss = (per_token * mask).sum() / denom out["loss"] = loss out["token_loss"] = per_token return out def step( self, token_ids: mx.array, caches: Optional[List[Tuple[mx.array, mx.array]]] = None, ) -> Tuple[mx.array, List[Tuple[mx.array, mx.array]]]: # token_ids shape: [B, 1] logits, new_caches = self.logits(token_ids, caches=caches) return logits[:, -1, :], new_caches