Transformer Math Lab: Test the Identities Behind the Diagrams
Run the reference calculations
Download transformer_math_lab.py. It requires Python 3 and NumPy. It does not download a model or train a network. Run it in an environment with NumPy installed:
python transformer_math_lab.py
The program uses a fixed random seed, float64 reference arithmetic, and assertions. It prints the maximum absolute error for each comparison and fails if an invariant is violated. Inspect the recorded run for the tested environment and results. These are correctness experiments, not performance benchmarks or trained-model evaluations.
1. Attention and causal caching
Compute dense causal attention over a short sequence. Independently process one query at a time with only its permitted prefix. The outputs should agree. Perturb future key/value rows and confirm earlier causal outputs stay unchanged. Together, these checks catch the two common mask errors: wrong triangle direction and wrong prefix offset.
The lab uses fixed Q/K/V arrays to isolate attention. A complete multi-layer decoder should additionally compare cached and uncached logits with its position handling, normalization, and cache append logic enabled. Passing the small attention test alone does not certify a whole model implementation.
2. Merge softmax summaries
Direct stable softmax subtracts the largest score in the full row. Online softmax instead processes chunks using a running maximum, exponential sum, and weighted value sum. The lab splits the same scores into uneven chunks, includes an empty/masked chunk, and checks the merged output against direct attention.
Using scores near 1000 ensures naive exponentiation would overflow, while stable versions remain finite. Averaging independently normalized chunk outputs is included as a conceptual counterexample in the Ring Attention chapter; it omits the relative partition masses.
3. Kernel attention as a recurrence
For positive feature vectors $\phi(q)$ and $\phi(k)$, compute each causal output directly from all pairwise kernel weights. Independently update $S_t=S_{t-1}+\phi(k_t)v_t^\top$ and $z_t=z_{t-1}+\phi(k_t)$, then read $S_t^\top\phi(q_t)/(z_t^\top\phi(q_t))$.
The identity should hold to numerical tolerance. This proves agreement for the chosen kernel features; it does not establish equality with ordinary softmax attention. Replacing an exponential kernel with another kernel is a model change, even if the recurrent implementation of the new kernel is exact.
4. Writes, decay, and noncommuting updates
With a unit key and $\beta=1$, a delta write should make the state read that key as the target value. An orthogonal query should retain its previous read. The gated example checks the order “decay, read, correct,” including a hand-computable result of 4 for the scalar example in Gated DeltaNet.
For KDA, compare the direct update with its affine matrix expression, then verify that swapping featurewise decay and the delta projection changes a deliberately chosen state. An equality test and a counterexample serve different purposes: one checks implementation consistency; the other guards against an invalid algebraic simplification.
5. RoPE depends on a relative offset
Construct two-dimensional rotations explicitly. Verify $(R_iq)^\top(R_jk)=q^\top R_{j-i}k$ and invariance when both positions receive the same shift. These tests catch sign and orientation mistakes without requiring a trained head.
The test covers one rotation plane. A production implementation must also match dimension pairing, frequency layout, broadcasting, and cached position indices across all heads and planes.
6. Absorb the MLA projections
Generate a latent matrix $C$, explicit keys $K=CU_K^\top$, and values $V=CU_V^\top$. Compare ordinary attention on those explicit arrays with scores computed using the absorbed query $U_K^\top q$, followed by a latent weighted sum and the value up-projection.
Include the decoupled rotary score term in both paths, with a single shared softmax. This checks the projection identity explained in MLA. It does not claim that an arbitrary pretrained MHA layer can be compressed to that latent rank without changing its function.
7. Recover the exact target distribution
For the three-token example in speculative decoding, calculate accepted proposal mass $\min(p,q)$ and the rejection contribution $[p-q]_+$. Their sum must equal $p$. This deterministic probability check is more precise than inferring exactness from a finite histogram of random samples.
An end-to-end speculative implementation also needs prefix-dependent distributions, ordered acceptance, stopping rules, and correct cache rollback. The one-step proof is necessary to understand it, but does not test all that state management.
Turn each experiment into a debugging question
Try changing the causal mask, dropping the linear-attention denominator, reading the pre-decay state in the gated update, or applying separate softmaxes to MLA score components. Predict which check should fail before running it. Restore the reference after each change so one intentional error does not obscure another.
For neural-network work, repeat the relevant comparisons in the intended dtype and then test gradients where training uses the operation. A float64 identity check establishes the algebra; mixed-precision tolerances and kernel behavior are additional engineering questions.