Transformer Mathematics: Attention Scaling, RoPE & Latent Projections
Attention variance scaling, SO(2) Givens rotation (RoPE), and Multi-Head Latent Attention (MLA).
🔍 Inspect Architecture: Attention mechanism💡 1. Core Intuition & Concepts
Modern transformer architectures rely on precise geometric principles: scaling dot-products by 1/√d_k preserves variance and prevents softmax saturation, Rotary Position Embeddings (RoPE) encode relative token distance via complex plane Givens rotations, and Multi-Head Latent Attention (MLA) uses low-rank projections to compress the KV-cache by 93%.
📐 2. Mathematical Formulations & Derivations
• Represent a 2D query vector in complex form q = q_0 + i q_1 and key vector k = k_0 + i k_1.
• Applying rotation at token position m corresponds to multiplication by complex exponential: q_m = q * e^(i m theta), and at position n: k_n = k * e^(i n theta).
• Compute the inner product as the real part of the complex product with conjugate: <q_m, k_n> = Re(q_m * conj(k_n)).
• Substitute the rotated forms: Re((q * e^(i m theta)) * (conj(k) * e^(-i n theta))) = Re((q * conj(k)) * e^(i (m - n) theta)).
• The resulting dot product depends exclusively on the relative token displacement (m - n) and the base frequency theta. Q.E.D.
⚙️ 3. Step-by-Step Computational Mechanism
💻 4. Code from Scratch (python)
import numpy as np
# Rotary Position Embedding (RoPE) 2D Complex Rotation from Scratch
def apply_rope_2d(x: np.ndarray, seq_len: int, dim: int, base: float = 10000.0):
dim_pairs = dim // 2
inv_freq = 1.0 / (base ** (np.arange(0, dim_pairs) * 2.0 / dim))
positions = np.arange(seq_len)
angles = np.outer(positions, inv_freq)
cos_vals = np.cos(angles)
sin_vals = np.sin(angles)
x_pairs = x.reshape(seq_len, dim_pairs, 2)
x0, x1 = x_pairs[:, :, 0], x_pairs[:, :, 1]
rot_x0 = x0 * cos_vals - x1 * sin_vals
rot_x1 = x0 * sin_vals + x1 * cos_vals
return np.stack([rot_x0, rot_x1], axis=-1).reshape(seq_len, dim)
if __name__ == "__main__":
seq_len, dim = 4, 8
q = np.random.randn(seq_len, dim)
q_rot = apply_rope_2d(q, seq_len, dim)
print("RoPE rotated Query shape:", q_rot.shape)
print("✓ RoPE 2D Rotation applied successfully.")🧠 5. Comprehension Checkpoint
In the catalog
Finished this chapter?