import math import numpy as np def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray: """Return a numerically stable softmax along one axis.""" shifted = x - np.max(x, axis=axis, keepdims=True) exponentials = np.exp(shifted) return exponentials / np.sum(exponentials, axis=axis, keepdims=True) def split_heads(projected: np.ndarray, num_heads: int) -> np.ndarray: """Convert [B, T, H * D_head] into [B, H, T, D_head].""" if projected.ndim != 3: raise ValueError("projected must be 3D") if num_heads <= 0: raise ValueError("num_heads must be positive") if projected.shape[-1] % num_heads != 0: raise ValueError("projection width must be divisible by num_heads") b = projected.shape[0] t = projected.shape[1] h = num_heads d_head = projected.shape[2] // h tmp = projected.reshape(b, t, h, d_head) multi_heads_attention = tmp.swapaxes(-2, -3) return multi_heads_attention def project_qkv( x: np.ndarray, w_q: np.ndarray, w_k: np.ndarray, w_v: np.ndarray, num_query_heads: int, num_kv_heads: int, ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: """Project X and return Q/K/V with explicit Head axes.""" if x.ndim != 3: raise ValueError("x must be 3D") if num_query_heads <= 0 or num_kv_heads <= 0: raise ValueError("num_query_heads and num_kv_heads must be positive") weights = (("w_q", w_q), ("w_k", w_k), ("w_v", w_v)) for name, weight in weights: if weight.ndim != 2: raise ValueError(f"{name} must be 2D") if weight.shape[0] != x.shape[-1]: raise ValueError(f"{name} input width must match x input width") projection_specs = ( ("w_q", w_q, num_query_heads), ("w_k", w_k, num_kv_heads), ("w_v", w_v, num_kv_heads), ) d_heads = [] for name, weight, num_heads in projection_specs: if weight.shape[1] % num_heads != 0: raise ValueError(f"{name} projection width must be divisible by num_heads") d_heads.append(weight.shape[1] // num_heads) if len(set(d_heads)) != 1: raise ValueError("Q/K/V D_head must match") q = x @ w_q k = x @ w_k v = x @ w_v q = split_heads(q, num_query_heads) k = split_heads(k, num_kv_heads) v = split_heads(v, num_kv_heads) return q, k, v def multi_head_causal_attention( q: np.ndarray, k: np.ndarray, v: np.ndarray, ) -> tuple[np.ndarray, np.ndarray]: """Return per-head output and attention weights without a KV cache.""" tensors = (("q", q), ("k", k), ("v", v)) for name, tensor in tensors: if tensor.ndim != 4: raise ValueError(f"{name} must be 4D [B, H, T, D_head]") if q.shape != k.shape or q.shape != v.shape: raise ValueError("q, k, and v must have the same [B, H, T, D_head] shape") if any(dimension <= 0 for dimension in q.shape): raise ValueError("B, H, T, and D_head must all be positive") _, _, num_tokens, head_dim = q.shape causal_mask = np.triu( np.full((num_tokens, num_tokens), -np.inf, dtype=np.float64), k=1, ) scores = (q @ np.swapaxes(k, -1, -2)) / math.sqrt(head_dim) weights = softmax(scores + causal_mask, axis=-1) output = weights @ v return output, weights def cached_multi_head_attention( q: np.ndarray, k: np.ndarray, v: np.ndarray, k_cache: np.ndarray | None = None, v_cache: np.ndarray | None = None, ) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]: """Run one decode step and return output, weights, K cache, and V cache. Current Q/K/V use [B, H, 1, D_head]. Existing caches are either both None or both [B, H, past_T, D_head]. M2-C asks you to append the current K/V and attend the current Query over the complete visible cache. """ raise NotImplementedError("M2-C TODO: append K/V and compute cached attention")