← Back to the Applied Science guide
From backprop to the transformer. Know the equations, know why they work, and know what breaks at scale.
This bank covers the depth round for neural nets. The first six questions are the classic core. Every Applied Science loop asks some of them. The last six cover transformers and large language models. Most loops now ask at least two of those.
Interviewers here do not want a definition. They want the equation, the reason behind it, and one real failure you have seen. Each answer below gives all three. Learn the math well enough to write it on a whiteboard without notes.
Can you apply the chain rule on a computational graph by hand? Do you know why reverse mode beats forward mode for training? Do you know what gets stored in memory?
Backpropagation is the chain rule applied in reverse order over a computational graph. The forward pass computes each node and caches the values it needs. The backward pass starts at the loss with gradient 1. Each node then multiplies the incoming gradient by its local derivative and passes it to its inputs.
For a node z = f(x, y) with upstream gradient ∂L/∂z, the rule is:
∂L/∂x = ∂L/∂z · ∂z/∂x ∂L/∂y = ∂L/∂z · ∂z/∂y
When a value feeds several nodes, its gradients add up. This is the multivariate chain rule. It is the most common bug in hand-written backprop.
Worked example. Take one neuron with a sigmoid and a squared loss. Let x = 2, w = 0.5, b = -0.5, target y = 1.
z = w·x + b = 0.5. Then a = σ(0.5) ≈ 0.622. Then L = ½(a - y)² ≈ 0.0714.∂L/∂a = a - y = -0.378.σ'(z) = a(1 - a) ≈ 0.235. So ∂L/∂z ≈ -0.378 × 0.235 ≈ -0.0889.∂L/∂w = ∂L/∂z · x ≈ -0.178 and ∂L/∂b = ∂L/∂z ≈ -0.0889.w ← 0.518. The output moves toward 1, as it should.Why reverse mode. A network maps millions of parameters to one scalar loss. Reverse mode gets the whole gradient in one backward pass. Its cost is about two to three times the forward pass. Forward mode would need one pass per input parameter. Reverse mode wins whenever outputs are few and inputs are many.
Matrix form. For a layer Y = XW with upstream gradient G = ∂L/∂Y:
∂L/∂W = Xᵀ G ∂L/∂X = G Wᵀ
Check shapes to derive these. If X is n×d and W is d×k, then G is n×k. The only way to get a d×k result is XᵀG.
import numpy as np
def linear_forward(X, W, b):
out = X @ W + b
cache = (X, W) # keep inputs for the backward pass
return out, cache
def linear_backward(dout, cache):
X, W = cache
dX = dout @ W.T # (n,k) @ (k,d) -> (n,d)
dW = X.T @ dout # (d,n) @ (n,k) -> (d,k)
db = dout.sum(axis=0) # b was broadcast over the batch, so sum
return dX, dW, db
Memory. The forward pass must keep activations for the backward pass. So training memory grows with depth times batch size times width. This is why activation checkpointing exists (see question 12).
(f(θ+ε) - f(θ-ε)) / 2ε with ε near 1e-5 in float64. Compare with relative error. Below 1e-6 is good.p - y, the predicted probabilities minus the one-hot label. That clean form is why we fuse them.Do you see the product of Jacobians behind the problem? Can you derive why Xavier and He init use the scales they use? Do you know which fix targets which cause?
Backprop through L layers multiplies L Jacobians together:
∂L/∂h₀ = ∂L/∂h_L · ∏ₗ (∂hₗ/∂hₗ₋₁) = ∂L/∂h_L · ∏ₗ diag(f'(zₗ)) Wₗ
If the typical gain of each factor is below 1, the product shrinks toward zero. Early layers stop learning. If the gain is above 1, the product blows up. Loss turns into NaN. The effect is exponential in depth, so a gain of 0.9 over 50 layers leaves 0.5 percent of the signal.
Causes.
Fix 1: initialization. The goal is to keep the variance of activations, and of gradients, the same from layer to layer. For y = Wx with fan-in n, zero-mean independent weights give Var(y) = n · Var(w) · Var(x). So we want n · Var(w) = 1.
Var(w) = 2 / (fan_in + fan_out). It balances the forward and backward passes. It assumes a linear-ish activation like tanh near zero.Var(w) = 2 / fan_in. ReLU zeros half its inputs, which halves the variance. The factor 2 puts it back. Use He for ReLU and its variants.1/√(2L). That keeps the residual stream from growing with depth.Fix 2: residual connections. With hₗ₊₁ = hₗ + F(hₗ), the Jacobian is I + ∂F/∂h. The identity term gives the gradient a direct path to every layer. Even if F's gradient is tiny, the signal does not vanish. This is the main reason we can train 100-layer ResNets and transformers.
Fix 3: normalization. BatchNorm and LayerNorm keep pre-activations near zero mean and unit variance. Activations then stay out of saturation, and the gain stays near 1.
Fix 4: non-saturating activations. ReLU has slope 1 for positive inputs. GELU and SiLU are smooth versions used in transformers. Leaky ReLU avoids dead units.
Fix 5: gradient clipping (for exploding). Clip by global norm: if ‖g‖ > c, set g ← g · c / ‖g‖. This keeps the direction and caps the step size. It is standard in RNNs and in LLM training, with c = 1.0 as a common value. Clipping does not fix vanishing.
Fix 6: gating. An LSTM has a cell state updated by addition: cₜ = fₜ ⊙ cₜ₋₁ + iₜ ⊙ gₜ. The gradient through the cell is mostly the forget gate fₜ. When f stays near 1, the signal flows across many steps. This is the same idea as a residual connection, applied through time. GRUs work the same way.
Can you write each update rule from memory? Do you know why Adam needs bias correction? Can you explain why AdamW decouples weight decay? Do you know how to set the learning rate schedule?
All of them take a step against the gradient gₜ = ∇L(θₜ) computed on a minibatch. They differ in how they shape that step.
SGD.
θₜ₊₁ = θₜ - η gₜ
Simple and memory-free. The noise from minibatches helps it escape sharp minima. But it zigzags in ravines where the curvature differs a lot between directions.
Momentum (heavy ball).
vₜ = μ vₜ₋₁ + gₜ
θₜ₊₁ = θₜ - η vₜ
The velocity is a running sum of past gradients with decay μ, often 0.9. Steps that agree add up. Steps that flip sign cancel out. In steady state the effective step is η / (1 - μ), ten times larger at μ = 0.9. Nesterov momentum takes the gradient at the look-ahead point. It gives a small, reliable gain.
Adam. It keeps a running mean and a running uncentered variance of each gradient coordinate:
mₜ = β₁ mₜ₋₁ + (1 - β₁) gₜ
vₜ = β₂ vₜ₋₁ + (1 - β₂) gₜ²
m̂ₜ = mₜ / (1 - β₁ᵗ) v̂ₜ = vₜ / (1 - β₂ᵗ)
θₜ₊₁ = θₜ - η m̂ₜ / (√v̂ₜ + ε)
Defaults are β₁ = 0.9, β₂ = 0.999, ε = 1e-8. LLMs often use β₂ = 0.95 for stability. Each parameter gets its own step size. Parameters with rare or small gradients get larger steps. The step size is roughly η in every coordinate, whatever the gradient scale.
Why bias correction. m and v start at zero. Early on they are biased toward zero. After one step, m₁ = (1 - β₁) g₁ = 0.1 g₁. The expected value of mₜ is (1 - β₁ᵗ) E[g] if g is stationary. Dividing by 1 - β₁ᵗ removes that bias. The v correction matters more. Without it, v is tiny early, so 1/√v is huge, and the first steps are far too big.
AdamW and decoupled weight decay. Classic L2 regularization adds λθ to the gradient. In SGD, that is the same as shrinking weights by a constant factor each step. In Adam it is not. The L2 term goes through the 1/√v scaling. So weights with large gradient history get almost no decay. Weights with small gradients get strong decay. That is not what we want. AdamW applies decay directly to the weights, outside the adaptive step:
θₜ₊₁ = θₜ - η ( m̂ₜ / (√v̂ₜ + ε) + λ θₜ )
Now every weight shrinks by the same fraction. AdamW generalizes better and makes λ easier to tune apart from η. It is the default for transformers, with λ near 0.1. Do not apply decay to biases or norm gains.
def adamw_step(p, g, m, v, t, lr=3e-4, b1=0.9, b2=0.95, eps=1e-8, wd=0.1):
m = b1 * m + (1 - b1) * g
v = b2 * v + (1 - b2) * g * g
m_hat = m / (1 - b1 ** t) # t starts at 1
v_hat = v / (1 - b2 ** t)
p = p - lr * (m_hat / (np.sqrt(v_hat) + eps) + wd * p)
return p, m, v
Learning rate schedules.
ηₜ = η_min + ½(η_max - η_min)(1 + cos(π t/T)). Most LLM runs use this, decaying to about 10 percent of peak.Do you know which axis each one normalizes? Do you know what changes at inference? Can you explain pre-norm vs post-norm and why RMSNorm took over?
Both apply the same recipe. Subtract a mean, divide by a standard deviation, then apply a learned scale γ and shift β:
y = γ · (x - μ) / √(σ² + ε) + β
The difference is which values the mean and variance come from.
Train vs inference for BatchNorm. In training, BN uses the current batch's mean and variance. It also updates a running average: μ_run ← 0.9 μ_run + 0.1 μ_batch. At inference, it uses those running stats. So a single example gives a fixed, repeatable output. The layer then folds into a plain affine transform. You can merge it into the previous conv weights for speed. LayerNorm does the same thing in training and at inference.
Why BN helps. It allows higher learning rates and makes training less sensitive to init. The original claim was that it reduces internal covariate shift. Later work showed the main effect is a smoother loss surface. The batch noise also acts as a mild regularizer.
Why transformers use LayerNorm.
RMSNorm. It drops the mean subtraction and the β shift:
y = γ · x / √(mean(x²) + ε)
It is cheaper and works just as well in practice. Re-centering turned out to matter little. Llama, Mistral and most recent LLMs use it.
Pre-norm vs post-norm.
x = LN(x + Sublayer(x)). The norm sits on the residual path. Gradients to early layers must pass through every LN. Deep post-norm models need careful warmup or they diverge. When they do train, final quality can be slightly better.x = x + Sublayer(LN(x)). The residual path is a clean identity from input to output. Gradients flow without change. Training is stable at depth and less sensitive to warmup. Add a final LN before the output head, because the residual stream is never normalized otherwise.One known weakness of pre-norm is that the residual stream grows with depth. Later layers then add relatively little. Some models add extra norms (sandwich norm or QK-norm) to fix stability at large scale.
model.eval(). BN then uses batch stats at inference.Do you know inverted dropout and why we scale? Can you name and justify the other tools? Do you know which ones matter for large models?
Dropout. During training, zero each activation with probability p. This stops units from co-adapting. Each step trains a random sub-network. At test time, you use the full network, which acts like an average over that ensemble.
Inverted dropout. The expected value of a dropped unit is (1 - p) · h. To keep the expected value the same, divide the kept units by 1 - p during training. Then inference needs no change at all. This is what every framework does.
def dropout(h, p, training):
if not training or p == 0:
return h
mask = (np.random.rand(*h.shape) > p) / (1 - p) # keep with prob 1-p, rescale
return h * mask # reuse mask in backward
Typical rates are 0.1 in transformers and 0.5 in old fully connected layers. Many large LLM pre-training runs use no dropout at all. They see each token about once, so they barely overfit. Dropout returns for fine-tuning on small data.
Weight decay. Add λ‖w‖²/2 to the loss, or shrink weights directly each step (see AdamW in question 3). It prefers small, spread-out weights. That is the same as a Gaussian prior on weights in a Bayesian view. In nets with normalization, weight decay also controls the effective learning rate. Scale-invariant weights shrink, so the same step changes them more.
Data augmentation. Make new training examples that keep the label. Crops, flips and color jitter for images. Mixup and CutMix blend images and labels. For text, back-translation or token masking. Augmentation encodes your known invariances. It often beats any other regularizer in vision.
Early stopping. Track validation loss and stop when it stops improving for some patience. Keep the best checkpoint. For a quadratic loss with gradient descent, early stopping acts like L2 regularization. The number of steps plays the role of 1/λ.
Label smoothing. Replace the one-hot target with (1 - ε) · onehot + ε / K, with ε near 0.1. The model can no longer push logits to infinity to reach zero loss. That improves calibration and often accuracy. It can hurt knowledge distillation, because it erases the relative information between wrong classes.
Others worth naming.
Can you compute shapes and parameter counts fast and correctly? Do you understand how receptive field grows? Do you know why 1×1 convs and pooling exist?
A conv layer slides a small kernel over the input. Each output is a dot product between the kernel and a local patch. Two ideas make it work for images. Local connectivity means each output sees only a small patch. Weight sharing means the same kernel runs everywhere. That gives translation equivariance and far fewer parameters than a dense layer.
Output size. For input size W, kernel K, padding P, stride S and dilation D:
out = floor( (W + 2P - D·(K - 1) - 1) / S ) + 1
With D = 1 this is floor((W + 2P - K) / S) + 1. Example: a 224 input, K = 7, S = 2, P = 3 gives floor(223/2) + 1 = 112. "Same" padding for odd K and stride 1 is P = (K - 1)/2.
Parameter count. A conv from C_in to C_out channels with a K×K kernel has:
params = C_out · (C_in · K · K + 1) # +1 for the bias
FLOPs ≈ 2 · H_out · W_out · C_out · C_in · K²
Example: 3×3, 64 to 128 channels gives 128 · (64·9 + 1) = 73,856 parameters. Note that parameters do not depend on the image size. FLOPs do.
Receptive field. This is the patch of input pixels that can affect one output unit. It grows layer by layer:
rₗ = rₗ₋₁ + (Kₗ - 1) · jₗ₋₁ jₗ = jₗ₋₁ · Sₗ (r₀ = 1, j₀ = 1)
Here j is the jump, the distance in input pixels between neighboring units. Three stacked 3×3 convs at stride 1 give a 7×7 receptive field. That matches one 7×7 conv but uses 3 · 9C² = 27C² parameters instead of 49C². It also adds two extra nonlinearities. This is the VGG insight. Strides and pooling multiply the jump, so receptive field grows fast in later layers. Dilation also grows it without extra parameters.
The effective receptive field is smaller than the theoretical one. Pixels at the center of the patch have far more paths to the output. The influence falls off roughly like a Gaussian.
Pooling. Max or average pooling downsamples the feature map. It cuts compute, grows receptive field and adds a little translation invariance. Modern nets often use strided convs instead. Global average pooling at the end turns any spatial size into one vector per channel. It replaces large dense heads.
1×1 convolution. It is a dense layer applied at each pixel across channels. Uses:
C_in · K² + C_in · C_out, versus C_in · C_out · K² for a full conv.Can you write scaled dot-product attention and explain every symbol? Do you know why we divide by √d_k? Can you draw the full block, count its cost, and explain the causal mask?
Attention lets each token build its output as a weighted mix of all other tokens. The weights depend on how well a query matches each key. Each token projects its vector x into three vectors: a query q = xW_Q, a key k = xW_K and a value v = xW_V.
Attention(Q, K, V) = softmax( Q Kᵀ / √d_k + M ) V
Here Q is n×d_k, K is n×d_k and V is n×d_v. The score matrix QKᵀ is n×n. Row i says how much token i should look at each token j. Softmax turns each row into weights that sum to 1. M is a mask of 0 and −∞.
Why divide by √d_k. Assume the entries of q and k are independent with mean 0 and variance 1. Then q·k = Σ qᵢkᵢ has mean 0 and variance d_k. With d_k = 128, scores have a standard deviation near 11. Softmax over such large values is nearly one-hot. Its gradient is then close to zero almost everywhere, so learning stalls. Dividing by √d_k brings the variance back to 1. Softmax stays in a range with useful gradients.
Multi-head attention. Split d_model into h heads, each of size d_k = d_model / h. Run attention in each head with its own projections. Concatenate the outputs and apply an output projection W_O. Different heads can track different relations: syntax, coreference, position. The total cost is the same as one big head, because each head is smaller.
Causal mask. A decoder must not see future tokens. Set M[i, j] = −∞ for j > i. After softmax, those weights become 0. This lets us train on every position of a sequence in one parallel pass, which is called teacher forcing. Encoders like BERT use no causal mask. They only mask padding.
def attention(Q, K, V, causal=False):
d_k = Q.shape[-1]
scores = Q @ K.swapaxes(-1, -2) / np.sqrt(d_k) # (..., n, n)
if causal:
n = scores.shape[-1]
mask = np.triu(np.ones((n, n), dtype=bool), k=1) # True above diagonal
scores = np.where(mask, -np.inf, scores)
scores = scores - scores.max(axis=-1, keepdims=True) # stable softmax
w = np.exp(scores)
w = w / w.sum(axis=-1, keepdims=True)
return w @ V
The block (pre-norm, modern form).
x = x + MultiHeadAttn( Norm(x) )
x = x + FFN( Norm(x) )
FFN(x) = W₂ · act(W₁ x) # hidden size often 4 · d_model
W₂(SiLU(W₁x) ⊙ W₃x), with a hidden size near 8/3 · d_model to keep the parameter count equal.Complexity. For sequence length n and model width d:
O(n d²).O(n² d).O(n d²) with a constant near 8 or more.O(n²) per head, if you store it.For short sequences the d² terms dominate. Past n near d (often a few thousand tokens), the n² term takes over. A rough rule for forward FLOPs per token is 2N + 2 · n_layers · n · d, where N is the parameter count.
FlashAttention. It computes exact attention in tiles that fit in on-chip SRAM. It uses an online softmax that keeps a running max and sum. It never writes the full n×n matrix to GPU memory. FLOPs stay O(n²d), but memory drops to O(n). Speed improves a lot, because attention is limited by memory bandwidth, not compute.
Do you know why transformers need position information? Can you compare absolute, relative and rotary schemes? Do you know how models handle sequences longer than they saw in training?
Self-attention has no sense of order. "Dog bites man" and "man bites dog" would give the same set of outputs. We must inject position somehow.
Sinusoidal (original transformer). Add a fixed vector to each token embedding:
PE(pos, 2i) = sin( pos / 10000^(2i/d) )
PE(pos, 2i+1) = cos( pos / 10000^(2i/d) )
Each pair of dimensions is a sine wave at its own frequency. The wavelengths form a geometric series from 2π to about 10000 · 2π. A shift by k positions is a fixed linear map (a rotation) on each pair. So the model can, in principle, learn relative offsets. It needs no parameters. In practice it does not extrapolate well to longer sequences.
Learned absolute. A trainable vector per position, up to a max length. Used in BERT and GPT-2. Simple and works well within range. It cannot handle positions past the max at all.
Relative bias. T5 adds a learned scalar bias to each attention score, based on the bucketed distance i - j. Position lives in the attention scores, not the embeddings.
RoPE (rotary position embedding). Rotate the query and key vectors by an angle that depends on position. Split each vector into 2-D pairs. Pair i at position m rotates by angle mθᵢ, with θᵢ = base^(-2i/d) and base = 10000 by default:
[q'₀] [cos mθ -sin mθ] [q₀]
[q'₁] = [sin mθ cos mθ] [q₁]
The key property is this. Rotations preserve dot products up to the angle difference. So
〈R(m) q, R(n) k〉 = 〈q, R(n - m) k〉
The attention score depends only on the content and the relative offset n - m. We get relative position for free, with no extra parameters. It applies only to Q and K, not V. It works with the KV cache, since each key is rotated once by its own position. Low-frequency pairs also give a mild decay of attention with distance. Llama, Mistral, Qwen and most open LLMs use RoPE.
def rope(x, pos, base=10000.0):
# x: (n, d) with d even. pos: (n,) integer positions.
d = x.shape[-1]
inv_freq = base ** (-np.arange(0, d, 2) / d) # (d/2,)
ang = pos[:, None] * inv_freq[None, :] # (n, d/2)
cos, sin = np.cos(ang), np.sin(ang)
x1, x2 = x[:, 0::2], x[:, 1::2] # pairs
out = np.empty_like(x)
out[:, 0::2] = x1 * cos - x2 * sin
out[:, 1::2] = x1 * sin + x2 * cos
return out
ALiBi. Use no position embedding at all. Add a linear penalty to each attention score: score - mₕ · (i - j), with a fixed slope mₕ per head. Slopes form a geometric series, so some heads look far and some look near. ALiBi extrapolates to longer inputs better than sinusoidal or vanilla RoPE. But its built-in recency bias can hurt tasks that need long-range retrieval.
Length extrapolation. A model trained at 4K tokens often fails at 16K. With RoPE, new positions produce rotation angles never seen in training, for the low-frequency pairs. Fixes:
Do you know the full LLM post-training pipeline? Can you write the LoRA update and count its parameters? Can you explain RLHF and derive why DPO skips the reward model?
A pre-trained LLM predicts text. It is not yet a helpful assistant. Post-training usually runs in stages: supervised fine-tuning, then preference tuning with RLHF or DPO.
Supervised fine-tuning (SFT). Train on prompt and response pairs written or approved by humans. Use the normal next-token cross-entropy, but only on the response tokens. Mask the prompt tokens out of the loss. A few thousand to a few hundred thousand high-quality examples teach format and tone. Quality beats quantity here.
Full fine-tuning vs PEFT. Full fine-tuning updates every weight. It gives the best quality on big shifts. But AdamW needs about 16 bytes per parameter for weights, gradients and optimizer states in mixed precision. A 7B model needs over 100 GB before activations. You also store a full copy per task. Parameter-efficient fine-tuning (PEFT) trains a small set of new weights and freezes the rest.
LoRA. The idea is that the weight change for a task has low intrinsic rank. Freeze W and learn a low-rank update:
W' = W + ΔW = W + (α / r) · B A
W ∈ R^(d×k), B ∈ R^(d×r), A ∈ R^(r×k), r « min(d, k)
dk to r(d + k). For d = k = 4096 and r = 16, that is 131K instead of 16.8M, a 128× cut.QLoRA stores the frozen base in 4-bit NF4 and trains LoRA adapters in bf16 on top. A 65B model then fits on one 48 GB GPU.
RLHF. Three steps.
L = -log σ( r(x, y_w) - r(x, y_l) ).max_π E[ r(x, y) ] - β · KL( π(y|x) ‖ π_ref(y|x) )
The KL term stops reward hacking. Without it, the policy finds outputs the reward model loves but humans do not, like long, flattering text. PPO needs four models in memory: policy, reference, reward and value. It is unstable and costly to tune.
DPO. The KL-constrained objective above has a closed-form optimal policy:
π*(y|x) ∝ π_ref(y|x) · exp( r(x, y) / β )
Solve for r. That gives r(x, y) = β log(π*(y|x) / π_ref(y|x)) + β log Z(x). Plug this into the Bradley-Terry loss. The partition function Z(x) cancels in the difference. What remains is a simple classification loss on the policy itself:
L_DPO = -E log σ( β [ log π(y_w|x)/π_ref(y_w|x) - log π(y_l|x)/π_ref(y_l|x) ] )
No reward model, no sampling, no value network. Just two models (policy and frozen reference) and a supervised loss on preference pairs. It is far simpler and more stable than PPO.
Trade-offs. DPO is offline. It learns only from fixed pairs, so it can overfit and drift toward odd outputs, like longer responses. Online RL can explore and use fresh samples. Many top labs still use online RL for reasoning, often with verifiable rewards like unit tests or math answers. GRPO is a popular PPO variant that drops the value network. It uses the mean reward of a group of samples as the baseline.
Can you compute KV cache memory for a real model? Do you know why decoding is memory bound? Can you explain sampling methods and speculative decoding?
The KV cache. An autoregressive model makes one token at a time. Each new token attends to all earlier tokens. The keys and values of earlier tokens never change, thanks to the causal mask. So we compute them once and store them. Each step then computes Q, K and V for only the new token. Without the cache, step t would redo all t tokens, for O(n²) total projection work.
Memory math.
KV bytes = 2 · n_layers · n_kv_heads · d_head · seq_len · batch · bytes_per_value
The 2 is for K and V. Example: a 7B Llama-style model with 32 layers, 32 KV heads, d_head 128, in fp16. Per token: 2 · 32 · 32 · 128 · 2 = 524,288 bytes, so 0.5 MB. A 4K context is 2 GB per sequence. A batch of 16 needs 32 GB, more than the 14 GB of weights.
With GQA and 8 KV heads, the cache shrinks 4× to 128 KB per token. This is why most new models use GQA. Other tools are KV quantization to 8 or 4 bits, and PagedAttention. Paging stores the cache in fixed blocks like OS pages, which kills fragmentation and lets requests share prefixes.
Prefill vs decode. Prefill processes the whole prompt in parallel. It is compute bound and sets time to first token. Decode makes one token per step. Each step must read all weights and the whole cache from GPU memory to do a small amount of math. So decode is memory-bandwidth bound. Larger batches raise throughput, because one weight read serves many sequences.
Decoding strategies.
pᵢ ∝ exp(zᵢ / T). T below 1 sharpens the distribution. T above 1 flattens it. T near 0 becomes greedy.def sample_top_p(logits, p=0.9, temperature=0.7, rng=np.random):
z = logits / temperature
probs = np.exp(z - z.max())
probs /= probs.sum()
order = np.argsort(-probs) # most likely first
cum = np.cumsum(probs[order])
keep = order[: np.searchsorted(cum, p) + 1] # smallest set reaching p
q = probs[keep] / probs[keep].sum()
return rng.choice(keep, p=q)
Speculative decoding. A small draft model proposes k tokens fast. The large target model checks all k in one forward pass, which costs about the same as one decode step. Accept each draft token with probability min(1, p_target / p_draft). At the first rejection, resample from the normalized residual max(0, p_target - p_draft) and stop. This scheme gives exactly the target model's distribution. The speedup is often 2 to 3×. It depends on how often the draft agrees with the target. Variants include Medusa (extra heads on the target itself) and EAGLE (draft from the target's hidden states).
Do you know the idea behind learned embeddings? Can you write the InfoNCE loss? Do you understand in-batch negatives, hard negatives and the temperature?
An embedding maps an item (a word, a sentence, an image, a user) to a dense vector. Similar items should land close together. We measure closeness with cosine similarity or a dot product. Embeddings power search, recommendation, clustering and retrieval for RAG.
word2vec. The idea is the distributional hypothesis: a word is known by the company it keeps. Skip-gram trains a word vector to predict its context words. A full softmax over the vocabulary is costly. So it uses negative sampling. Push up the score of a real (word, context) pair. Push down the scores of a few random words:
L = -log σ(u_c · v_w) - Σₖ log σ(-u_k · v_w), k ~ P(w)^0.75
The 0.75 power gives rare words more chances to be sampled. This is already a contrastive loss: one positive against several negatives.
InfoNCE. Given an anchor q, one positive k₊ and negatives k₋, normalize all to unit length and compute:
L = -log [ exp(sim(q, k₊) / τ) / Σᵢ exp(sim(q, kᵢ) / τ) ]
This is cross-entropy where the right "class" is the positive. Minimizing it maximizes a lower bound on the mutual information between the two views. The bound is at most log of the number of candidates. So more negatives give a tighter bound.
In-batch negatives. In a batch of N pairs (qᵢ, kᵢ), use every other kⱼ as a negative for qᵢ. This gives N − 1 negatives for free. The similarity matrix is N×N, and the labels are the diagonal. CLIP uses this in both directions, image to text and text to image. Bigger batches mean more negatives and better results. That is why CLIP used batches of 32K. MoCo keeps a queue of past keys from a slowly updated momentum encoder, to get many negatives with small batches.
def info_nce(q, k, tau=0.05):
# q, k: (N, d). Row i of q matches row i of k.
q = q / np.linalg.norm(q, axis=1, keepdims=True)
k = k / np.linalg.norm(k, axis=1, keepdims=True)
logits = q @ k.T / tau # (N, N)
logits -= logits.max(axis=1, keepdims=True)
log_prob = logits - np.log(np.exp(logits).sum(axis=1, keepdims=True))
return -np.mean(np.diag(log_prob)) # positives on the diagonal
Hard negatives. Random negatives get easy fast. The loss on them goes to zero and teaches nothing. Hard negatives are items that look similar but are wrong. For search, mine them with BM25 or with the current model's top results that are not labeled relevant. They give much stronger gradients. The risk is false negatives: items that are actually relevant but unlabeled. Filter by a score threshold, or drop the very top results.
Temperature τ. It sets how peaked the softmax is. Small τ (0.01 to 0.1) puts most of the gradient on the hardest negatives. It gives tight clusters but can be unstable and punish false negatives too hard. Large τ treats all negatives about the same and gives softer structure. CLIP learns τ as a parameter, with a cap.
Other ways to train embeddings.
max(0, d(a, p) - d(a, n) + margin). One negative at a time. Needs careful mining.Do you know the compute-optimal trade-off between model size and data? Can you estimate training FLOPs? Do you understand fp16 vs bf16, loss scaling and memory-saving tricks?
Scaling laws. Test loss falls as a smooth power law in model size N, data size D and compute C. This holds over many orders of magnitude. The Chinchilla form is:
L(N, D) = E + A / N^α + B / D^β (α ≈ 0.34, β ≈ 0.28)
E is the irreducible loss of the data. Training compute is about C ≈ 6 N D FLOPs. The 6 is 2 for the forward pass plus 4 for the backward pass, per parameter per token.
Chinchilla result. For a fixed compute budget, grow N and D at about the same rate, each roughly as C^0.5. That works out to about 20 training tokens per parameter. A 70B model wants about 1.4T tokens. Earlier work (Kaplan 2020) said to favor bigger models. So GPT-3 (175B on 300B tokens) was undertrained. Chinchilla, at 70B on 1.4T tokens, beat the 280B Gopher with the same compute.
Beyond compute-optimal. Chinchilla optimizes training cost only. Inference cost scales with N for every token served. So labs now overtrain small models far past 20 tokens per parameter. Llama 3 8B saw 15T tokens, close to 2000 per parameter. Loss keeps falling, just more slowly. You pay more in training to get a cheaper model to serve.
How to use scaling laws in practice. Train a sweep of small models. Fit the power law. Extrapolate to choose N, D and hyperparameters for the big run. Watch out: downstream task metrics can jump in ways the loss curve does not show, and data quality shifts the whole curve.
Mixed precision. Do the heavy matmuls in 16-bit and keep sensitive parts in fp32. You get about 2× less memory for activations and much faster math on tensor cores.
Loss scaling (fp16). Multiply the loss by a scale S, often starting near 2^16, before backward. All gradients shift up by S into fp16's range. Unscale before the optimizer step. With dynamic scaling, if any gradient is inf or NaN, skip the step and halve S. After a run of clean steps, double S. bf16 usually needs no loss scaling, because its range matches fp32. That is why bf16 is the default for LLM training on modern hardware.
What stays in fp32. A master copy of the weights, so tiny updates are not lost to rounding. With 7 mantissa bits, adding 1e-4 to a weight of 1.0 does nothing. Also the optimizer states, reductions like softmax and norm statistics, and the loss.
Memory budget with AdamW in mixed precision. Per parameter: 2 bytes for 16-bit weights, 2 for gradients, 4 for the fp32 master, 4 for m and 4 for v. That is 16 bytes per parameter. A 7B model needs 112 GB before activations. ZeRO and FSDP shard these states across GPUs to fit.
Activation checkpointing. Activations often dominate memory, and they grow with batch size and sequence length. Checkpointing saves only some activations, often each block's input. During backward, it reruns the forward pass for each block to rebuild the rest. Memory drops from O(L) to about O(√L) with optimal placement, or to one block's worth per layer in common setups. The cost is about one extra forward pass, so roughly 30 percent more compute.
import torch
from torch.utils.checkpoint import checkpoint
scaler = torch.cuda.amp.GradScaler() # only needed for fp16
for x, y in loader:
with torch.autocast("cuda", dtype=torch.float16):
h = x
for block in model.blocks:
h = checkpoint(block, h, use_reentrant=False) # recompute in backward
loss = loss_fn(model.head(h), y)
scaler.scale(loss).backward()
scaler.unscale_(opt)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(opt) # skips step on inf/NaN
scaler.update()
opt.zero_grad(set_to_none=True)