Appendix D
Solutions to Exercises
Worked solutions for every published exercise, grouped by chapter.
6 Probability Theory
The density at the mean is . The probability of landing within 0.05 of 0 is , where is the standard Gaussian CDF. A density measures probability per unit length. Squeezing the distribution into a narrow range must raise its height so that the total area stays 1. A probability is an area, bounded by the total area of 1. For a narrow interval, , which is small however tall is.
With a 20% base rate, flagged essays are generated and human, so
The detector’s error rates are properties of the detector, but the posterior also depends on the prior. When generated essays are rare, even a small false-positive rate applied to the large human majority produces more false alarms than true detections. The same arithmetic governs any classifier used to find a rare class.
Expand the variance of the sum into variances and covariances:
For i.i.d. draws every covariance is zero, which leaves . With correlation , each of the covariance terms is , so
However large becomes, the variance never falls below . With and , it is 0.128 instead of the i.i.d. 0.031: four times larger. A minibatch of near-duplicates behaves like a much smaller batch. The simulation below, checked in the tests, builds correlated draws from a shared component:
def variance_of_correlated_mean(n, rho, trials, rng):
"""Empirical Var of the mean of n unit-variance draws, pairwise correlation rho."""
shared = rng.standard_normal((trials, 1)) # what every draw has in common
own = rng.standard_normal((trials, n))
x = np.sqrt(rho) * shared + np.sqrt(1 - rho) * own # Var 1, Cov(x_i, x_j) = rho
return x.mean(axis=1).var()
Up to a constant, the average negative log-likelihood is
Setting gives . Setting gives . The tests confirm that the finite-difference gradient vanishes at these values.
For the bias, write . Summing the squares, the cross terms combine into . Take expectations, using :
The sample mean sits closer to the data than the true mean does, so deviations from it are
too small on average. Dividing by instead removes the bias (NumPy’s
ddof=1). With , the maximum-likelihood variance averages 0.8 of the truth:
def average_mle_variance(n, trials, rng):
"""Average maximum-likelihood variance of many size-n samples from N(0, 1)."""
x = rng.standard_normal((trials, n))
return np.mean(np.mean((x - x.mean(axis=1, keepdims=True)) ** 2, axis=1))
With fixed, the Gaussian negative log-likelihood is plus terms that do not depend on the prediction. Minimizing its average is minimizing mean squared error. For Laplace noise,
so the loss is mean absolute error:
def laplace_regression_nll(y, prediction, scale=1.0):
"""y ~ Laplace(prediction, scale): mean absolute error / scale + log(2 scale)."""
return np.mean(np.abs(y - prediction) / scale + np.log(2 * scale))
For a constant prediction, squared error is minimized by the mean of the targets and absolute error by the median. A few large outliers drag the mean but barely move the median, so MSE chases outliers and MAE resists them. Each loss is the right choice exactly when its noise model matches the data.
For a discrete distribution, the gradient passes through the finite sum. Then use :
With the left side is the gradient of 1, so . Hence for any constant .
For with , the score is . Use the moments , , , with odd moments zero:
-
Score function: . Then , so .
-
With baseline 2: . Then , so .
-
Reparameterized: , so .
def score_function_gradient_with_baseline(f, mean, std, n, rng, baseline):
"""Subtracting a constant from f leaves the expected gradient unchanged."""
x = mean + std * rng.standard_normal(n)
return np.mean((f(x) - baseline) * (x - mean) / std ** 2)
The tests check all three means and standard deviations against these values.
Let and be the Gumbel CDF and density. Index wins when for every . Condition on ; the other noises are independent:
The term in the sum comes from itself. Write and substitute , so that :
Scaling after the noise, , picks the same index for any , so temperature must act on the logits first: samples . The tests compare empirical frequencies from 200,000 draws with the softmax probabilities at and .
The CDF is on , so :
def sample_triangle(n, rng):
"""Density 2x on [0, 1]: F(x) = x^2, so x = sqrt(u)."""
return np.sqrt(rng.random(n))
The mean is , and . With 200,000 samples, both agree with these values to within 1%.
7 Information Theory
Surprisal is bits, or nats:
-
A fair coin landing heads, : 1 bit, or 0.693 nats.
-
A fair die showing six, : 2.585 bits, or 1.792 nats.
-
One token from a uniform 128,000-token vocabulary: bits, or 11.76 nats.
def surprisal_bits(probability):
return float(-np.log2(probability))
A language model with a cross-entropy of 2 nats per token is therefore doing far better than uniform guessing, which would cost 11.76 nats per token.
With ,
Gibbs' inequality makes the left side nonnegative, so , with equality exactly when . The tests check the identity for 200 random distributions.
Restrict the sum to outcomes with , and let with . Jensen’s inequality for the concave logarithm gives
Equality in Jensen’s step requires to be constant where . Equality in the last step requires to put all its mass there. Together they force .
Group the sum over samples by value. The value appears times, so
and splits it into a constant and a divergence. The tests confirm the equality numerically:
def empirical_distribution(samples, categories):
"""The fraction of samples equal to each category: p_hat."""
return np.bincount(samples, minlength=categories) / len(samples)
def average_nll(samples, q):
"""The maximum-likelihood objective: -(1/N) sum_i log q(x_i)."""
return float(-np.mean(np.log(q[samples])))
def cross_entropy_of_empirical(samples, q):
"""The same number, as the cross-entropy H(p_hat, q)."""
return float(cross_entropy(empirical_distribution(samples, len(q)), q))
Over all distributions, Gibbs' inequality says the minimum is at : a model that only memorizes training frequencies, assigning zero probability to anything unseen. Real models are kept from that point by their parameterization, which cannot represent arbitrary tables and must share structure across inputs. They are also held back by regularization and early stopping, and by the softmax, which never outputs an exact zero. The gap between memorizing and generalizing toward is the subject of the next part of the book.
The log-ratio of the densities is
Under , and . Taking expectations gives (7.6). With equal parameters the result is . With equal variances it reduces to : quadratic in the distance between the means. The tests compare the formula with numerical integration for three pairs of Gaussians.
Summing over the support of , . So and
For nonnegativity, the logarithm is concave, so it lies below its tangent line at 1: for every . Hence , with equality only at . Adding the zero-mean term cancels much of 's fluctuation, because and move together.
, and only the second term depends on :
The mean enters only through , which is minimized by . Setting the derivative in to zero gives . For the two-mode target, the mean is 0 and the variance is , so . The grid search in the tests finds the same values, to within its step size.
The reverse direction is . The target’s log-density sits inside an expectation over the fitted distribution. For a mixture, that expectation has no closed form, and the objective has one local minimum per mode. Which one an optimizer finds depends on where it starts.
For , the marginals are and :
-
As a KL divergence from the product of the marginals: 0.1258 nats.
-
From entropies, : the same 0.1258 nats.
-
nats and nats, so again nats, or 0.18 bits.
When always, for example with the diagonal table , knowing removes all uncertainty about . Then nats. For the product of the marginals, the joint equals the product, and the mutual information is exactly 0. All three cases are checked in the tests.
8 Hypothesis Testing
The accuracy is . The plug-in standard error is . A 95% normal interval is , or 0.572 to 0.668. The standard error is the standard deviation we would expect for the accuracy estimate over repeated evaluation sets of the same size, not the model’s per-example error rate.
The median has no simple Bernoulli standard-error formula, but it is still a statistic of rows. The bootstrap approximates its sampling distribution by resampling rows with replacement and recomputing the median. For an LLM benchmark, the row should be the independent evaluation unit: a task, prompt, conversation, or judged comparison, not an individual token from inside one answer.
def median_interval(scores, rng):
"""Bootstrap a robust median score instead of an accuracy."""
return bootstrap_ci(np.asarray(scores), np.median, draws=2000, rng=rng)
The five discordant rows contain four wins for B and one win for A, so the accuracy gap is . Under the null, the five discordant rows are fair coin flips. Outcomes at least this extreme have zero or one B wins, or four or five B wins, so the p-value is . The exact permutation test and McNemar’s test coincide here:
def paired_demo():
"""A tiny benchmark where B fixes four A errors and breaks one A success."""
a_correct = np.array([1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 1, 1])
b_correct = np.array([1, 0, 1, 1, 1, 1, 1, 1, 1, 0, 1, 1])
return {
"delta": float(b_correct.mean() - a_correct.mean()),
"permutation_p": exact_permutation_p_value(a_correct, b_correct),
"mcnemar_p": mcnemar_exact_p_value(a_correct, b_correct),
}
Rearrange to get
, then round up. The worst case is . With
, . The implementation in the chapter
uses exactly that formula and ceil, because collecting a fraction of an example is impossible.
9 Learning from Data
Population risk is , the average loss on future examples from the deployment distribution. Empirical risk replaces that expectation with a sample mean over the training set. The estimate can be biased if the training set was collected from a different population, filtered by a previous model, deduplicated incorrectly, or contaminated with test examples. More examples reduce sampling noise but do not fix a mismatched sampling process.
Start with . Expanding gives . The gradient is . Setting it to zero and multiplying by gives . If is invertible, solve for ; otherwise use least squares or regularization.
For one example, and . The BCE derivatives are and . Multiplying gives . The chain rule then gives . Averaging rows stacks those row vectors into , whose shape matches because maps per-example errors back to feature weights.
The tested snippet builds a noisy cubic training set, fits three polynomial feature maps, and measures validation loss against the clean signal. Degree 1 underfits because it cannot bend. Degree 11 overfits because it uses extra powers to chase noise. Degree 3 best matches the data generating function in this example.
def polynomial_losses(seed=9):
rng = np.random.default_rng(seed)
x_train = np.linspace(-1.0, 1.0, 12, dtype=np.float32)
y_clean = 0.5 + x_train - 1.5 * x_train ** 2 + 0.7 * x_train ** 3
y_train = y_clean + rng.normal(0.0, 0.08, size=len(x_train)).astype(np.float32)
x_val = np.linspace(-1.0, 1.0, 200, dtype=np.float32)
y_val = 0.5 + x_val - 1.5 * x_val ** 2 + 0.7 * x_val ** 3
losses = {}
for degree in (1, 3, 11):
X_train = polynomial_features(x_train, degree)
X_val = polynomial_features(x_val, degree)
w, b = normal_equation(X_train, y_train)
losses[degree] = float(np.mean((predict_linear(X_val, w, b) - y_val) ** 2))
return losses
10 Automatic Differentiation
The graph has leaves and . It computes , , , then
adds those three values and applies log. The node is shared because it feeds both
and . During the backward pass, must receive the contribution
from the multiply path and the square path. Missing either contribution gives the derivative of
a different program.
For , a small change gives . Multiplying by the output adjoint gives and . For , use the trace identity: = . Thus has the same shape as , and has the same shape as .
Forward mode propagates one chosen input direction to all outputs, so computing a full gradient with respect to many parameters would require many sweeps. Reverse mode starts from the scalar loss adjoint and computes all parameter adjoints in one backward sweep, so it is the natural choice for training neural networks. Forward mode is attractive when the input dimension is small and the output dimension is large, or when only a few directional derivatives are needed.
The tested helper constructs the graph, calls backward, and returns gradients for all three
inputs. The chapter test flattens , , and , recomputes the scalar loss
under small central-difference perturbations, and checks the concatenated autodiff gradient.
def tiny_network_loss_and_grads(X, W, b):
x = Tensor(X)
w = Tensor(W)
bias = Tensor(b)
loss = ((x @ w + bias).relu().exp().log()).sum()
loss.backward()
return float(loss.data), x.grad, w.grad, bias.grad
11 Activation Functions
For sigmoid, . As , ; as , . Either way the product goes to 0. Since and , tanh also saturates. ReLU has slope 0 for negative inputs, so a unit that is always negative never updates from the loss signal. Leaky ReLU replaces that 0 slope by , leaving a path for gradients.
For ,
, whose largest value is at .
Using gives , largest
at 0 with value 1. SiLU is a product, so
. Softplus has derivative
. The test test_elementwise_activation_derivatives_match_finite_differences
checks these formulas against central differences.
The exact GELU is . Since , the product rule gives . The measured approximation error is computed, not guessed:
def gelu_tanh_max_error(limit=8.0, points=200_001):
x = np.linspace(-limit, limit, points)
return float(np.max(np.abs(gelu_exact(x) - gelu_tanh(x))))
On the grid used by the tests, the maximum absolute error is , and the test asserts that exact measured value.
Ignoring biases, the plain MLP has weights. The gated block has two input projections and one output projection, . Setting gives .
def gated_hidden_width(model_width, mlp_multiplier=4):
"""Hidden width h with 3 d h parameters matching a d -> 4d -> d MLP."""
return mlp_multiplier * 2 * model_width / 3
The backward pass in scratch.activations.gated_ffn_backward applies the product rule to the
gate and is checked with finite differences for , , , and
.
12 Softmax & Cross-Entropy
For every class, , so the factor cancels. Temperature changes the gaps before this normalization:
def probabilities_at_temperatures(logits, temperatures):
return [softmax(logits, temperature=T) for T in temperatures]
The test checks that for the largest probability is highest at , lower at , and closer to uniform at .
Let and . For , quotient rule gives . For , only the denominator changes, giving . Therefore . A row sum is , matching shift invariance.
Using log-softmax, . Since , this is . Differentiating gives . With temperature, all logits inside softmax are , so the chain rule multiplies the gradient by .
The implementation computes log-softmax by shifting logits, then returns the mean loss and the gradient with respect to logits:
def loss_only(logits, labels, smoothing=0.0, z_loss=0.0):
loss, _ = softmax_cross_entropy(logits, labels,
label_smoothing=smoothing,
z_loss=z_loss)
return loss
The tests check three cases: the plain one-hot gradient , a temperature and label-smoothing gradient, and the z-loss gradient added to cross-entropy.
13 Loss Functions & Divergences
For Gaussian noise with fixed , . Only the squared residual depends on , so maximum likelihood minimizes MSE. For Laplace noise, , so the prediction-dependent part is MAE. The scale changes the gradient size but not the optimum when fixed.
With , MSE has gradient , MAE has away from zero, and Huber has for and outside. The cap is visible in the tested helper:
def outlier_gradients(residual, delta=1.0):
prediction = np.array([residual], dtype=np.float64)
target = np.array([0.0])
_, mse_grad = (residual ** 2, np.array([2 * residual]))
_, huber_grad = huber_loss(prediction, target, delta)
return mse_grad[0], huber_grad[0]
For BCE, substitute and simplify separately for and to get . Differentiating gives . The stable implementation stays finite even for logits :
def stable_bce_example():
logits = np.array([1000.0, -1000.0])
labels = np.array([1.0, 0.0])
return bce_with_logits(logits, labels)[0]
Focal loss multiplies BCE by . When , the multiplier is near zero, so easy examples contribute little. When is small, the multiplier is near one, so hard examples keep a cross-entropy-like signal. Forward KL, , averages over the target and punishes missing target mass. Reverse KL, , averages over the model and is more willing to focus on one mode. Jensen-Shannon is symmetric because it averages and with the same midpoint .
For softened student probabilities and fixed teacher , the cross-entropy gradient is . Multiplying the loss by makes it . The tests use this wrapper:
def scaled_and_unscaled_distillation(student, teacher, temperature):
_, scaled = distillation_loss(student, teacher, temperature, scale=True)
_, unscaled = distillation_loss(student, teacher, temperature, scale=False)
return scaled, unscaled
They check that the scaled gradient is times the unscaled one for the same softened distributions, and that the analytic gradient matches finite differences.
14 Neural Networks from Scratch
The arrays have shapes and . The parameters are , , , and . The count is , checked by the helper:
def parameter_count():
params = initialize(seed=0)
return sum(value.size for value in params.values())
The softmax derivative is . Therefore
For the mean loss, , so every row gradient is divided by . All later matrix products are linear in that upstream gradient, so dividing again would make the gradient too small by another factor of . The tests check this gradient against finite differences.
Because the terms are independent and zero mean, . Variances of independent sums add, so . Choosing keeps the preactivation scale comparable to the input scale. A ReLU makes about half of a symmetric preactivation zero, so the second moment is roughly halved; He initialization compensates with .
The tested implementation uses float64 copies for gradient checking and float32 for the training run. Its SGD loop lowers the synthetic-data loss and reaches high training accuracy. On a tiny set with permuted labels, it can fit the update examples better than it agrees with the original clean labels:
def memorization_gap(seed=7):
inputs, labels = synthetic_data(seed=seed, examples_per_class=4)
rng = np.random.default_rng(seed)
noisy = rng.permutation(labels)
params, _ = train_sgd(inputs, noisy, epochs=400, learning_rate=0.12, seed=seed)
_, cache = forward(params, inputs, noisy)
train_accuracy = np.mean(cache["P"].argmax(axis=1) == noisy)
_, clean_cache = forward(params, inputs, labels)
clean_accuracy = np.mean(clean_cache["P"].argmax(axis=1) == labels)
return float(train_accuracy), float(clean_accuracy)
That gap is overfitting: optimization succeeded on the training objective, but the fitted rule matched noise instead of the data-generating pattern.
15 Optimizers & Schedules
For eigenvalue , the error multiplier is . Stability requires for every eigenvalue, so . With , the direction goes to zero in one step, but the direction is multiplied by each step. The best fixed-rate worst-case factor is:
def optimal_gd_rate(condition_number):
return (condition_number - 1) / (condition_number + 1)
Unrolling the recurrence gives
The same geometric sum gives . Without correction, both moment estimates are biased toward zero at small . Dividing by and makes the constant-gradient estimates equal to and .
def adam_first_step():
params = {"w": np.array([2.0, -3.0])}
grads = {"w": np.array([0.5, -0.25])}
state = adam_state(params)
before = params["w"].copy()
adamw_step(params, grads, state, lr=0.01)
return before - params["w"]
With L2 regularization, Adam sees and then divides by the adaptive , so the decay part is coordinate-scaled like any other gradient. AdamW instead applies separately, then takes the adaptive gradient step. For clipping, , so every tensor is scaled by :
def clipped_demo_norm():
grads = {"a": np.array([3.0, 4.0]), "b": np.array([12.0])}
clipped, before = clip_by_global_norm(grads, max_norm=5.0)
after = np.sqrt(sum(np.sum(g * g) for g in clipped.values()))
return before, float(after)
For a matrix , the nearest orthogonal polar factor in Frobenius norm is . The Newton-Schulz iteration approximates that factor using only matrix multiplies, which is cheaper than an SVD in a training step. The tests compare the chapter implementation with NumPy’s SVD on tall and wide matrices. Muon is restricted here to 2-D hidden weights because biases, gains, embeddings, and output heads do not have the same interior matrix geometry.
16 Normalization, Residuals & Precision
BatchNorm estimates and from the current minibatch during training, but inference may use a different batch size, or one example at a time, so it uses running statistics accumulated during training. For , BatchNorm reduces over the axis for each feature. LayerNorm reduces over the axis inside each row, so it uses the current example’s own statistics in both training and inference.
Let and . A change in one input coordinate affects the normalized output directly, through the row mean, and through the row variance. Collecting those terms gives
The first mean removes the component caused by mean subtraction. The second removes the component caused by changing the row’s variance. The tests check this formula with finite differences.
For and ,
The helper checks the same invariance numerically:
def rms_scale_invariance(x, weight, factor):
y1, _ = rms_norm_forward(x, weight)
y2, _ = rms_norm_forward(factor * x, weight)
return np.max(np.abs(y1 - y2))
For a residual block , differentiating gives an identity term plus the gradient through . Even if the learned branch is small or badly scaled, the identity term passes gradient backward.
Inverted dropout samples a Bernoulli keep mask and divides kept activations by the keep probability, so its expectation equals the input. The tested helper estimates that mean:
def dropout_mean(seed=0):
rng = np.random.default_rng(seed)
x = np.ones(20_000)
y, _ = inverted_dropout(x, 0.25, rng)
return float(y.mean())
The precision demo casts large float32 values to fp16 and rounded bfloat16, showing fp16 overflow while bfloat16 remains finite. It also multiplies a tiny gradient by a scale before the fp16 cast and divides after, recovering a nonzero unscaled gradient. Float32 master weights then accumulate updates that would be lost if every step were rounded to a low-precision copy.
17 Tokenization & Embeddings
Characters make unknown text easy and keep the alphabet small, but they make the sequence long. Words give short, readable sequences for common text, but unhappiness and the emoji may need unknown-token handling unless both were in the vocabulary. Byte-level subwords keep exact coverage because every string is bytes first, and frequent chunks such as un or ness can be merged; their drawback is that rare text may still split into many small pieces.
Stack the one-hot rows into , so the embedding output is . For a scalar loss,
Column of is 1 exactly for positions whose id is , so row of is the sum of those upstream rows. The checked helper shows the repeated id receiving both contributions:
def repeated_index_gradient():
"""A repeated token id receives the sum of both upstream gradients."""
indices = np.array([1, 3, 1])
grad_output = np.array([[1.0, 0.0], [0.0, 2.0], [3.0, 4.0]])
return embedding_backward(indices, grad_output, vocab_size=5)
The trainer in Section 17.2 counts pairs after each rewrite, so later merges can combine pieces that did not exist at the start. A compact way to check the result is to compare corpus length before and after BPE:
def corpus_token_lengths(texts=TINY_CORPUS, vocab_size=266):
"""Byte count before BPE and token count after BPE on a tiny corpus."""
_vocab, merges = train_byte_bpe(texts, vocab_size)
before = sum(len(utf8_bytes(text)) for text in texts)
after = sum(len(encode(text, merges)) for text in texts)
return before, after, 256 + len(merges)
The tests assert that the tiny corpus has 27 UTF-8 bytes, becomes 10 BPE tokens after ten merges, and still round-trips. The emoji round-trips because all 256 single-byte tokens remain in the vocabulary even if no emoji byte sequence appeared during training.
Use np.add.at, because plain grad_weight[ids] += grad_output can lose updates when an id repeats. A complete implementation is the embedding_backward function in Section 17.4. It creates a zero table with the vocabulary size and embedding width, then scatter-adds each upstream row into the selected token row. The chapter tests compare it with a Python loop over flattened ids.
18 Language Modeling
The factorization writes a joint probability as a product of next-token conditionals. To sample, start with a beginning context, draw from , append it, then draw from . Repeating this procedure samples from the product distribution defined by the model. Greedy decoding uses the same conditionals but takes the most likely token instead of sampling.
The row total is , and there are three possible next tokens. Add-one smoothing gives
The entries sum to . In general, the numerator adds to each of cells, and the denominator adds to the row total.
def smoothed_tiny_loss(alpha=1.0):
ids, vocab = word_ids(SYNTHETIC_TEXT)
previous, target = make_bigrams(ids)
probs = smoothed_bigram_probs(ids, len(vocab), alpha)
return average_nll_from_probs(probs, previous, target)
For one example, , , and . The softmax-cross-entropy derivative is . Since is row of , only that row receives the gradient:
A batch averages those row updates over examples. The tests check the implemented gradient against finite differences and then verify that gradient descent lowers the loss:
def trained_tiny_losses():
ids, vocab = word_ids(SYNTHETIC_TEXT)
previous, target = make_bigrams(ids)
_W, losses = train_neural_bigram(previous, target, len(vocab))
return np.array([losses[0], losses[-1]])
For ids a b c d e and width , the contexts are (a,b,c) and (b,c,d), with targets d and e. The chapter code constructs this sliding window and feeds the gathered embeddings through an MLP. Any dependency farther back than the width is invisible: with width , the prediction after a b c d cannot depend on a except through parameters learned from other examples. Attention removes that fixed cutoff by letting the current position read earlier hidden states directly.
19 Scaled Dot-Product Attention
The query is the thing asking for information. The keys are searchable addresses, and the values are the content stored at those addresses. Attention compares the query with every key, turns the scores into weights, and returns the weighted average of the values. It is "soft" because every unmasked value can contribute and because the weights are differentiable.
For one coordinate, independence gives and . Different coordinates are independent, so variances add:
Dividing by divides the variance by , leaving variance near one. The tested helper estimates the same ratio numerically:
def variance_ratios(widths=(4, 16, 64)):
"""Return Var(q dot k) / d_k for several widths."""
return np.array([dot_product_variance(width) / width for width in widths])
The softmax Jacobian for one row is . Multiplying by the upstream row gives
The scalar is the row sum of . Applying this row by row gives (19.4).
The implementation in Section 19.3 follows the chain rule in reverse: , then row-softmax, then the scaled score matrix. Masked score gradients are set to zero. The chapter tests check all three inputs with finite differences under a causal mask. This tiny helper exposes the weights for a causal self-attention example:
def causal_attention_weights():
"""Attention weights for a tiny self-attention problem with a causal mask."""
Q = np.array([[[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]]])
K = Q.copy()
V = np.eye(3)[None, :, :]
_output, cache = attention_forward(Q, K, V, causal_mask(3), True)
return cache[3][0]
20 Multi-Head Attention
One head produces one attention distribution per query position. Several heads can produce several distributions, each after a different learned projection of the same hidden states. That lets the layer retrieve different mixtures of values for different relations, then combine them with . If the total width stays fixed, this is not about adding more output dimensions; it is about giving the model several scoring subspaces.
Starting from , each projection keeps shape . Splitting into heads gives with . Per-head attention returns . Combining heads transposes and reshapes that back to , and the output projection keeps . The code path is the pair of split_heads and combine_heads in Section 20.1.
Each of , , , and has parameters, so the total is . For self-attention, the four dense projections cost multiply-adds. Scores cost , and multiplying weights by values costs another , giving . The helper computes the same budget for a tiny setting:
def tiny_budget(model_width=8, length=5, batch=2):
"""Parameters and dominant self-attention FLOPs for a tiny setting."""
flops = self_attention_flops(batch, length, model_width)
return parameter_count(model_width), flops
Reverse the forward pass. Backpropagate through , split the concatenated gradient into heads, call the single-head attention backward for each head, combine query/key/value head gradients, and then backpropagate through , , and . For self-attention the same input fed all three projections, so add those three input gradients before comparing with finite differences. The tests also verify that a causal mask is shared across heads:
def first_head_causal_weights():
"""Return attention weights from the first head of a tiny masked MHA."""
X = np.arange(12, dtype=np.float64).reshape(1, 3, 4) / 10
eye = np.eye(4)
params = (eye, eye, eye, eye)
_out, cache = multi_head_attention_forward(
X, X, X, params, 2, causal_mask(3), True
)
return cache[5][3][0, 0]
21 Positional Encoding & RoPE
Let . After applying the same permutation to queries and keys, the scores are . Rowwise softmax preserves that row and column permutation, so the weights are . Multiplying by leaves , proving the output is merely reordered. A language model must distinguish orders, so it needs an extra positional signal.
def permutation_error(x, wq, wk, wv, permutation):
q, k, v = x @ wq, x @ wk, x @ wv
original = attention(q, k, v)
xp = x[permutation]
permuted = attention(xp @ wq, xp @ wk, xp @ wv)
return np.max(np.abs(permuted - original[permutation]))
For one pair, because a rotation matrix is orthogonal. Angles add, so . Summing over all independent two-dimensional pairs gives (21.4). If both positions are shifted by , the relative angle becomes , so the dot product is unchanged.
def shifted_rope_dot(q, k, m, n, shift):
left = apply_rope(q, m) @ apply_rope(k, n)
shifted = apply_rope(q, m + shift) @ apply_rope(k, n + shift)
return left, shifted
The last query has , so the row is . A larger makes distant past keys pay a larger negative bias before softmax, concentrating attention more strongly on recent keys unless the content score overcomes it.
def last_query_alibi(length, slope):
return alibi_bias(length, slope)[-1]
Use positions shaped so the cosine and sine tables broadcast across batch and heads. The even coordinates receive , and the odd coordinates receive . The inverse uses the negative angle. The chapter tests check inverse recovery and the shared-shift dot-product identity on seeded random arrays.
def apply_rope(x, positions, base=10_000.0, inverse=False):
"""Rotate each adjacent 2-D pair by positions * theta_i."""
x = np.asarray(x)
if x.shape[-1] % 2:
raise ValueError("RoPE needs an even last dimension")
positions = np.asarray(positions, dtype=x.dtype)
if positions.ndim == 1 and x.ndim > 2:
shape = [1] * (x.ndim - 1)
shape[1] = positions.size
positions = positions.reshape(shape)
theta = rope_frequencies(x.shape[-1], base).astype(x.dtype)
angles = positions[..., None] * theta
if inverse:
angles = -angles
cos, sin = np.cos(angles), np.sin(angles)
y = np.empty_like(x)
even, odd = x[..., 0::2], x[..., 1::2]
y[..., 0::2] = even * cos - odd * sin
y[..., 1::2] = even * sin + odd * cos
return y
22 The Transformer Block
Attention is the communication sublayer: each token reads earlier tokens through causal weights and writes a proposed update. The feed-forward network is the local computation sublayer: it applies the same nonlinear map to each token independently. Pre-norm leaves the residual stream itself on an identity path, so gradients can move from deep layers to shallow layers without first passing through attention or the MLP.
Let and . Since ,
The direct path gives . The path through contributes . Adding them gives (22.3). With a learned gain, first set , and the gain gradient is over batch and time.
Attention has query, key, value, and output matrices, each , for . SwiGLU has gate and up matrices and a down matrix , for . Thus
If , the block has parameters. If , the feed-forward part has and the block is near , matching the classic attention-plus-MLP budget.
def block_parameter_count(d_model, hidden_dim):
attention = 4 * d_model * d_model
swiglu = 3 * d_model * hidden_dim
norms = 2 * d_model
return attention + swiglu + norms
The tests build a block with one batch item, a short sequence, two heads, and float64 parameters.
They run the forward pass, backpropagate a fixed upstream array, flatten all parameter gradients,
and compare both input and parameter gradients with central differences using
scratch.gradcheck.check_gradient.
def block_scalar_loss(x, params, n_heads, upstream):
y, _ = transformer_block_forward(x, params, n_heads)
return float(np.sum(y * upstream))
23 Training a GPT from Scratch
For start and length , the input row is . The target row is . Because the transformer returns logits at every input position, each position predicts the next character for its prefix, giving supervised next-token examples from one contiguous slice.
For the output head , each row of receives the classifier gradient
. The hidden state also receives
. Backpropagating to the embedding lookup adds the upstream gradient
for each input occurrence into that token’s row. If a character appears multiple times in the
batch, np.add.at accumulates all of those sparse contributions.
Warmup makes the first few Adam steps smaller while the moment estimates are still settling. AdamW decouples shrinkage from the adaptive gradient: the code multiplies each parameter by and then applies the Adam direction. If weight decay were added to the gradient, Adam’s per-coordinate normalization would change the meaning of the decay term.
def adamw_step(params, grads, state, lr, weight_decay=0.01,
beta1=0.9, beta2=0.999, eps=1e-8):
state["t"] = state.get("t", 0) + 1
t = state["t"]
for path, param, grad in tree_items(params, grads):
slot = state.setdefault(path, {
"m": np.zeros_like(param),
"v": np.zeros_like(param),
})
slot["m"] = beta1 * slot["m"] + (1.0 - beta1) * grad
slot["v"] = beta2 * slot["v"] + (1.0 - beta2) * (grad * grad)
m_hat = slot["m"] / (1.0 - beta1 ** t)
v_hat = slot["v"] / (1.0 - beta2 ** t)
param *= 1.0 - lr * weight_decay
param -= lr * m_hat / (np.sqrt(v_hat) + eps)
The tests create one fixed batch, run several AdamW updates, and assert that the last loss is
smaller than the first. They also initialize a tiny model and call sample_text twice with the
same seed, temperature, and top-k setting; both calls must produce the same string.
def sample_text(params, prompt, stoi, itos, n_heads, steps, seed=0,
temperature=1.0, top_k=None, max_context=64):
rng = np.random.default_rng(seed)
ids = [stoi[ch] for ch in prompt]
for _ in range(steps):
context = np.array([ids[-max_context:]], dtype=np.int64)
logits = gpt_logits(params, context, n_heads)[0, -1] / temperature
if top_k is not None and top_k < logits.size:
keep = np.argpartition(logits, -top_k)[-top_k:]
masked = np.full_like(logits, -np.inf)
masked[keep] = logits[keep]
logits = masked
probs = softmax(logits)
ids.append(int(sample_categorical(probs[None, :], rng)[0]))
return "".join(itos[i] for i in ids)
24 KV Cache & Grouped-Query Attention
For the newest token , full recomputation evaluates the same projections for every that earlier decode steps already cached. The causal mask must never let a token read a future position. Under that mask, replacing old projection work with cache reads changes only the order of computation, not the attention sum. The tests confirm this with a random NumPy attention layer.
Each cached token stores one key and one value per layer and KV head. Each vector has numbers, and each number has bytes, so one sequence costs bytes.
def memory_example_mib():
bytes_used = kv_cache_bytes(
layers=32,
num_kv_heads=8,
head_dim=128,
tokens=4096,
bytes_per_value=2,
)
return bytes_used // (1024 ** 2)
For the chapter’s dimensions this returns 512 MiB, and the test also asserts the exact byte count .
is multi-query attention: all query heads share one KV head. assigns one KV head to each query head. Then the group map is , no KV vector is repeated across different query heads, and the GQA equation is the ordinary multi-head equation. The test compares the implementation with a direct multi-head computation in that case.
A compact implementation is:
def last_row_with_sink():
return sliding_window_mask(tokens=6, window=3, sinks=1)[5].astype(int)
For , , and one sink, the final row is [1, 0, 0, 1, 1, 1]. The last token
may read token 0 as the sink and tokens 3, 4, and 5 as the local causal window.
25 Multi-Head Latent Attention
MLA caches for each previous token. When a query attends, the layer up-projects that latent to head-specific keys and values, or absorbs the key projection into the query for score computation. MHA stores both key and value vectors for every head; MLA stores one latent vector, so its cache scales with .
With , associativity gives
The key up-projection moves from every cached token to the current query. The tests compute both sides for random tensors and assert equality.
RoPE inserts position-specific rotations, producing . The term involving changes for each cached position, so there is no single absorbed query that works for all . A decoupled design keeps a small RoPE key channel outside the absorbed latent content path.
The table values are computed here:
def cache_table_mib():
rows = cache_size_table(
layers=32,
tokens=4096,
bytes_per_value=2,
heads=32,
kv_heads=8,
head_dim=128,
latent_dim=512,
)
return {name: bytes_used // (1024 ** 2) for name, bytes_used in rows.items()}
The result is {MHA: 2048, GQA: 512, MLA: 128} in MiB for the tested dimensions. To make the
sparse toy causal, set every index score for future keys to negative infinity before taking the
top-k set, then run the same masked softmax.
26 Online Softmax & FlashAttention
Subtracting the maximum makes the largest exponent , so exponentials avoid overflow while all softmax ratios stay unchanged. Online softmax needs the maximum because every stored partial sum is measured relative to it. If a later block has a larger maximum, the old partial sum must be rescaled to the new reference before adding the new block.
For old scores, multiply and divide by :
The new block gives . The numerator uses the same weights with values attached:
A dense implementation has one score for each query-key pair, so it stores score elements. The tiled version keeps a maximum, normalizer, and output numerator per query row, plus the current tile, so persistent row state is linear in .
def memory_elements_for_4096():
return attention_memory_elements(4096)
For , the tested counts are naive score elements and online-state entries.
The implementation in scratch/flash_attention.py already accepts causal=True. For each key
block it builds key positions, compares them with query positions, and sets future scores to
negative infinity before the online update. The tests compare tiled causal attention with dense
causal attention for random arrays, including uneven block sizes, so the mask and rescaling are
checked together.
27 Mixture of Experts
The selected mass is . Renormalizing gives weights and . They sum to one because every selected probability is divided by the same selected total: . The unselected expert contributes no expert output on this token.
Use Cauchy’s inequality on :
Thus the fixed-point Switch loss is at least . Equality requires all to be equal, so . A collapsed router has one and the rest zero, so its loss is . The tests assert the uniform value and the collapsed value for a small router.
Expert 0 receives the first three assignments, but its capacity is two. The third assignment to expert 0 is dropped; the assignment to expert 1 is kept. A dropped assignment contributes zero to the weighted sum in (27.3). This is why capacity protects the batch shape and communication budget but can lose information if the router is badly imbalanced.
A reference implementation routes one token at a time and accumulates the selected expert outputs directly:
def dense_reference_moe(x, w1, b1, w2, b2, logits, k):
"""Token-by-token reference used to test the sparse implementation."""
experts, weights, _ = top_k_router(logits, k)
out_dim = w2.shape[-1]
y = np.zeros((x.shape[0], out_dim), dtype=x.dtype)
for token in range(x.shape[0]):
for slot, expert in enumerate(experts[token]):
hidden = np.maximum(x[token] @ w1[expert] + b1[expert], 0)
y[token] += weights[token, slot] * (hidden @ w2[expert] + b2[expert])
return y
The chapter tests generate seeded random tensors, run this reference, and compare it with
sparse_moe_forward to float32 tolerance. Shape checks would miss swapped experts, missing
renormalization, and scatter-add mistakes; numerical equality catches those errors.
28 Linear Attention & State-Space Models
The feature map should make nonnegative for every query-key pair used in the denominator. Then the normalized linear-attention formula is a weighted average: each value receives weight , and the weights sum to one. If scores can be negative, the denominator can cancel or change sign, so the output is no longer an average of values.
Substitute into the numerator:
The sum is exactly , which updates by adding the new outer product. The denominator is the same factorization without , giving . The tests assert that this recurrent computation equals the explicit lower-triangular parallel form.
With , only the first row of matters. Let the current prediction error be . The update with adds half the error to that row, so the next prediction error is . After repeated updates the error is multiplied by each time. The test checks this decay and the closed form after repeated writes.
A one-token decoding update only needs the current token and the cached state:
def decode_one(q_t, k_t, v_t, state, normalizer, feature_map, eps=1e-8):
"""One-token update for linear-attention decoding."""
qt, kt = feature_map(q_t), feature_map(k_t)
state = state + np.outer(kt, v_t)
normalizer = normalizer + kt
y_t = (qt @ state) / max(qt @ normalizer, eps)
return y_t, state, normalizer
Running this function over a prefix produces the same outputs as recurrent_linear_attention in
the tests. The cached state contains scalars: the matrix
and vector . It does not grow with prefix length.
29 Scaling Laws & Pretraining Recipes
The forward pass touches each parameter for each token at roughly one multiply-add, counted as FLOPs, so it costs . Backpropagation computes activation and weight gradients and is about twice the forward cost, . The total is therefore . If doubles while is fixed, compute doubles.
Take logarithms of :
Thus a least-squares line fit with input and target has intercept and slope . The tested implementation recovers both values on synthetic data:
def fit_power_law(x, y):
"""Fit y = coefficient * x ** (-exponent) in log space."""
x = np.asarray(x, dtype=np.float64)
y = np.asarray(y, dtype=np.float64)
slope, intercept = np.polyfit(np.log(x), np.log(y), deg=1)
return float(np.exp(intercept)), float(-slope)
Let and . The part of the loss that depends on is
Differentiating and setting the result to zero gives . Multiplying by and substituting gives (29.4). At the optimum, the marginal benefit of spending compute on more parameters matches the marginal benefit of spending it on more tokens.
A grid search is short and useful for checking the closed form:
def best_on_grid(compute, parameter_grid, token_grid):
"""Search a small grid for the lowest Chinchilla loss under 6ND <= compute."""
best = None
for parameters in parameter_grid:
for tokens in token_grid:
if 6 * parameters * tokens > compute:
continue
loss = chinchilla_loss(parameters, tokens)
if best is None or loss < best[0]:
best = (loss, parameters, tokens)
return best
The tests verify that every feasible grid point has loss at least as large as the returned one. They also verify the closed-form stationarity condition, so the grid search is a sanity check, not the source of the formula.
30 Contrastive & Metric Learning
The target probability is . It is 0.982 at and 0.690 at .
def positive_probability(gap, temperature):
logits = np.array([[gap, 0.0]]) / temperature
return float(softmax_rows(logits)[0, 0])
The smaller temperature makes the same score gap look larger. In (30.5), it also divides the gradient by , so it changes update scale.
For one row,
Differentiating with respect to gives from the first term and from the second. Averaging rows gives (30.5). The tests finite-difference the implementation.
At zero logits, each pair contributes before the mean. With three positives and six negatives, the mean gradient is .
def one_positive_three_negative_bias():
similarity = np.zeros((3, 3))
_, _, grad_bias = siglip_loss_and_grad(similarity, bias=0.0)
return grad_bias
The positive sign means gradient descent decreases the bias. That raises the effective bar for calling a pair positive, compensating for there being more negatives.
The implementation forms a full similarity matrix, trains with the symmetric CLIP loss, and then sorts each row. With the seeded synthetic data used in the tests, recall@1 starts at or below 0.10, then reaches at least 0.70; recall@5 reaches at least 0.90.
from scratch.contrastive_learning import (make_synthetic_pairs,
retrieval_recall_at_k,
train_linear_pair)
images, texts = make_synthetic_pairs(n=32, seed=3)
image_z, text_z = train_linear_pair(images, texts, steps=220, lr=0.7, seed=5)
print(retrieval_recall_at_k(image_z, text_z, k=1))
print(retrieval_recall_at_k(image_z, text_z, k=5))
The threshold, not an exact printed number, is the claim: the tests assert both recalls.
31 Vision Transformers
The grid is patches on each side, so the image has patch tokens and 197 tokens after prepending a class token. At , the grid is , giving 784 patch tokens.
def token_counts():
grid = 224 // 16
return grid, grid * grid, sequence_length(224, 224, 16, True)
The tests also assert that sequence_length(448, 448, 16) returns 784.
The four patches are the two-by-two blocks in raster order:
def patch_order_example():
image = np.arange(16, dtype=np.float32).reshape(1, 4, 4, 1)
return patchify(image, 2)[0].astype(int).tolist()
They evaluate to , , , and . The test checks this exact order.
Let a flattened patch be and the embedding weight be . Reshape each column of into a kernel. A stride- convolution places that kernel on exactly the pixels of one patch and computes the same dot product . Because stride equals patch size, windows do not overlap. The test compares the two arrays.
For images and patches, the patch grid is , so there are 16 patch tokens. Class-token mode sends 17 tokens into the encoder and reads token 0. Mean-pooling mode sends 16 tokens and averages them after the encoder. The classification head is shared, so both modes return logits shaped . The test runs both modes and checks finite logits on a seeded synthetic batch.
32 Vision-Language Models
The first image is closest to the first prompt and the second image is closest to the second
prompt, so the predicted prompt indices are [0, 1].
def toy_zero_shot_label():
images = np.array([[1.0, 0.0], [0.0, 1.0]])
prompts = np.array([[0.9, 0.1], [0.1, 0.9], [-1.0, 0.0]])
return np.argmax(clip_zero_shot(images, prompts), axis=1).tolist()
No classifier head is needed because the text embeddings are the class weights. Changing labels means encoding different prompts and recomputing similarities.
Substitute into (32.3). Since , the residual branch becomes . Therefore for any image tokens . The tests assert exact equality at initialization and a changed output when the gate is nonzero.
The patch grid is by , so the image has visual tokens before merging. A 2 by 2 merge halves both grid axes, giving tokens.
def dynamic_counts_example():
return dynamic_token_counts(336, 672, patch_size=14, merge=2)
The code returns ((24, 48), 1152, 288), and the test asserts those values.
perceiver_resampler broadcasts the learned queries across the batch, uses them as cross-attention
queries, and attends over however many image tokens the encoder produced. The output length is
therefore the number of learned queries, not the image-token length. With four queries, both a
short sequence and a long sequence return shape . The bottleneck is
that all visual evidence must be compressed into those four output tokens before the LLM sees it.
33 Supervised Fine-Tuning & LoRA
Role markers are part of the token sequence. If training writes <|user|> and inference writes User:, the first tokens after every turn come from a distribution the model did not practice. In a user-assistant example, all tokens are context for later positions, but the loss mask is 1 only on assistant targets. User, system, padding, and boundary tokens have mask 0. The test test_chat_template_keeps_role_markers_and_masks_assistant_targets checks this exact split.
For one position, . Multiplying by and by the normalizer in (33.1) gives (33.2). If , every component of is zero, so changing that position’s logits cannot change the loss.
def softmax(logits, axis=-1):
"""Stable softmax."""
shifted = logits - np.max(logits, axis=axis, keepdims=True)
exp = np.exp(shifted)
return exp / np.sum(exp, axis=axis, keepdims=True)
def masked_cross_entropy(logits, targets, train_mask):
"""Mean next-token cross-entropy over masked positions, with gradient."""
logits = np.asarray(logits)
targets = np.asarray(targets, dtype=np.int64)
mask = np.asarray(train_mask, dtype=bool)
if not np.any(mask):
raise ValueError("at least one position must be trainable")
probabilities = softmax(logits, axis=-1)
rows = np.arange(targets.shape[0])
losses = -np.log(probabilities[rows, targets])
normalizer = np.sum(mask)
loss = np.sum(np.where(mask, losses, 0.0)) / normalizer
grad_logits = probabilities.copy()
grad_logits[rows, targets] -= 1.0
grad_logits *= mask[:, None] / normalizer
return loss, grad_logits
Both documents fit in one length-4 row after inserting boundaries:
The inputs are , and the targets are . The mask is : token 1 may predict token 2 inside document 0, but token 2 should not predict the boundary and the boundary should not predict the next document. The chapter test uses the same rule on a larger packed batch.
Let and be the upstream gradient after the scale and the multiply are accounted for. The differential is
Restoring the scale gives (33.4). For and , the code computes base weights and LoRA weights, so the trainable matrix parameters are reduced by . The tests gradient-check both factors and assert those numbers.
def lora_gradients(x, a, b, alpha, grad_y):
"""Backpropagate through the LoRA update for a loss on Y."""
scale = alpha / a.shape[0]
grad_update = x.T @ grad_y
grad_a = scale * b.T @ grad_update
grad_b = scale * grad_update @ a.T
return grad_a, grad_b
def lora_mse_loss_and_grads(x, w, a, b, alpha, target):
"""Tiny objective used by the chapter tests."""
y = lora_output(x, w, a, b, alpha)
diff = y - target
loss = 0.5 * np.mean(diff * diff)
grad_y = diff / diff.size
grad_a, grad_b = lora_gradients(x, a, b, alpha, grad_y)
return loss, grad_a, grad_b
def lora_parameter_savings(d_in, d_out, rank):
base = d_in * d_out
trainable = rank * (d_in + d_out)
return base, trainable, base / trainable
34 Reinforcement Learning Foundations
The terminal state’s value is 0. State 1 gives reward 1 and then terminates, so . State 0 gives reward 0 and moves to state 1, so . The same computation is the Bellman solve in the tested tiny_chain helper.
Use the score-function identity on the trajectory distribution:
The trajectory log-probability is environment terms plus . Environment terms have zero gradient. Rewards before time are fixed before action , so their expected score term is zero; replacing by gives (34.4). The bandit test gradient-checks the resulting categorical gradient.
Condition on . Since does not depend on the sampled action,
For importance sampling, use one copy of each behavior probability in expectation:
The chapter test builds a logged batch with exactly those behavior proportions and checks the estimate.
Write the first few TD errors:
Intermediate value terms cancel, so equals the -step reward sum plus . Separating the first term of (34.9) gives
def gae_recursive(rewards, values, gamma, lam):
"""Generalized advantage estimates by the backward recursion."""
deltas = td_errors(rewards, values, gamma)
advantages = np.zeros_like(deltas)
running = 0.0
for t in range(len(deltas) - 1, -1, -1):
running = deltas[t] + gamma * lam * running
advantages[t] = running
return advantages
35 Reward Models, PPO & RLHF
The reward margin is . The Bradley-Terry probability is . The loss for the preferred response winning is . A larger positive margin makes the preference more likely and the loss smaller; a negative margin would mean the model currently scores the loser above the winner.
For ,
Thus
Gradient descent subtracts this vector. Since , the update moves toward , raising relative to . The test gradient-checks this expression on synthetic preference pairs.
def bradley_terry_loss_and_grad(weights, winners, losers):
"""Loss -log sigmoid(r_w - r_l) for a linear reward model."""
weights = np.asarray(weights, dtype=np.float64)
features = np.asarray(winners) - np.asarray(losers)
margins = features @ weights
loss = np.mean(np.logaddexp(0.0, -margins))
sigmoid = 1.0 / (1.0 + np.exp(-margins))
grad = ((sigmoid - 1.0)[:, None] * features).mean(axis=0)
return float(loss), grad
def train_reward_model(winners, losers, steps=300, lr=0.5):
weights = np.zeros(winners.shape[1], dtype=np.float64)
for _ in range(steps):
_, grad = bradley_terry_loss_and_grad(weights, winners, losers)
weights -= lr * grad
return weights
For , the unclipped term increases with . The minimum in (35.5) is active as while , and becomes the constant when . So the gradient is up to the upper clip and zero above it.
For , lowering improves the objective. The clipped constant is selected when , so the gradient is zero below the lower clip and otherwise. The chapter tests check both zero-gradient blocked regions and finite-difference the active region.
def ppo_clipped_objective_and_grad(
logits,
old_logits,
actions,
advantages,
clip_eps=0.2,
):
"""Mean PPO clipped surrogate and its gradient for one categorical state."""
logits = np.asarray(logits, dtype=np.float64)
old_logits = np.asarray(old_logits, dtype=np.float64)
actions = np.asarray(actions, dtype=np.int64)
advantages = np.asarray(advantages, dtype=np.float64)
probs = softmax(logits)
old_probs = softmax(old_logits)
ratios = probs[actions] / old_probs[actions]
clipped = np.clip(ratios, 1.0 - clip_eps, 1.0 + clip_eps)
objective_terms = np.minimum(ratios * advantages, clipped * advantages)
active = np.where(
advantages >= 0.0,
ratios <= 1.0 + clip_eps,
ratios >= 1.0 - clip_eps,
)
grad = np.zeros_like(logits)
for action, ratio, advantage, is_active in zip(actions, ratios, advantages, active):
if not is_active:
continue
grad_logp = -probs.copy()
grad_logp[action] += 1.0
grad += advantage * ratio * grad_logp
return float(np.mean(objective_terms)), grad / len(actions)
With rewards , action 1 is best. The tested run starts from uniform logits, repeatedly samples a batch from the old policy, forms advantages by subtracting the batch mean reward, and applies clipped PPO ascent. The final probability of action 1 is above 0.9, and the test also confirms it increased from the first recorded policy.
For values and returns , the squared value loss is . Its gradient is . Both numbers are asserted in the test.
def categorical_entropy(probs):
probs = np.asarray(probs, dtype=np.float64)
return float(-np.sum(probs * np.log(probs)))
def value_loss_and_grad(values, returns):
values = np.asarray(values, dtype=np.float64)
returns = np.asarray(returns, dtype=np.float64)
diff = values - returns
return float(0.5 * np.mean(diff * diff)), diff / diff.size
def toy_ppo_run(action_rewards, steps=80, batch_size=96, lr=0.35, seed=0):
"""PPO on a one-state categorical policy with known action rewards."""
rng = np.random.default_rng(seed)
rewards = np.asarray(action_rewards, dtype=np.float64)
logits = np.zeros_like(rewards)
history = []
for _ in range(steps):
old_logits = logits.copy()
old_probs = softmax(old_logits)
actions = rng.choice(len(rewards), size=batch_size, p=old_probs)
batch_rewards = rewards[actions]
advantages = batch_rewards - np.mean(batch_rewards)
for _ in range(4):
_, grad = ppo_clipped_objective_and_grad(
logits,
old_logits,
actions,
advantages,
)
logits += lr * grad
history.append(softmax(logits))
return logits, np.array(history)
36 Direct Preference Optimization
Let and . Then
The last sum is . By Section 7.4, it is nonnegative and equals zero only when . The tests compare this rewritten form with the original objective for several categorical policies.
For two answers to the same prompt,
def dpo_delta(logp_w, logp_l, logref_w, logref_l, beta):
return beta * ((logp_w - logref_w) - (logp_l - logref_l))
The tests construct a policy from the closed-form optimum, recover the implicit reward including , and check that the preference logit equals the true reward difference.
For , use :
Then multiply by . The multiplier is near one for a badly ranked pair and near zero for an already confident pair.
def preference_weight(delta):
return sigmoid(-delta)
The test checks the analytic gradient against finite differences.
Use soft Bradley—Terry targets for each unordered pair and binary cross-entropy on the DPO logit. The derivative for pair is for item and its negative for item . Gradient descent then drives toward for every pair, so the learned policy has the same normalized form as (36.2).
def toy_distance():
reference = np.array([0.50, 0.30, 0.20])
rewards = np.array([0.0, 0.7, -0.4])
policy, optimum = fit_toy_dpo(reference, rewards, beta=0.6)
return float(np.max(np.abs(policy - optimum)))
The tests assert that this error is below and gradient-check the expected loss.
37 GRPO & Verifiable Rewards
For , the mean is . The variance is , so the standard deviation is . The advantages are therefore approximately .
def centered_group_example():
rewards = np.array([[1.0, 0.0, 0.0]])
return group_advantages(rewards)[0]
For , every reward equals the group mean and the standard deviation is zero, so the implementation returns all zeros. A prompt with no within-group ranking gives no direction for the policy update. The tests check zero mean and unit standard deviation for the mixed group.
With samples from , let . Then . Hence
def exact_k3_kl(policy, reference):
logp = np.log(policy)
logref = np.log(reference)
return float(np.sum(policy * k3_kl_to_reference(logp, logref)))
The tests compare this sum with the shared KL implementation and also check that each sample’s value is nonnegative.
Sequence-level normalization gives the two completions equal weights . Token-level normalization divides by the total of tokens, giving weights .
def example_weights():
lengths = np.array([[2.0, 6.0]])
sequence = normalization_weights(lengths, "sequence")
token = normalization_weights(lengths, "token")
return sequence, token
With sequence-level weighting, a short answer and a long answer can have equal influence. With token-level weighting, the long answer contributes more token gradients. This is why length policy and loss normalization cannot be separated.
The implementation treats each candidate answer as one categorical sequence. Each step freezes the current logits as , computes group-relative advantages from verifier rewards, and applies the clipped loss plus a small reference KL penalty. The gradient is checked by finite differences in the tests.
def toy_correct_probabilities():
policy, rewards = toy_grpo_run()
correct = rewards.astype(bool)
return policy[correct]
The tests assert that both correct answers exceed probability and that each wrong answer falls below .
38 Distillation & Reasoning Models
Temperature divides the logits before softmax. With , is sharply peaked on the first class. With larger , the same ordering remains, but probability moves from the top class into the lower classes.
def softened_teacher(teacher_logits, temperature):
return softmax(np.asarray(teacher_logits) / temperature)
The useful information is in the relative sizes of the wrong classes. A target that says one wrong class is much more plausible than another gives the student a smoother learning signal than a one-hot label.
For , the usual softmax-cross-entropy derivative with respect to is . By the chain rule,
Multiplying the objective by gives . Since the difference between softened distributions shrinks roughly like , the multiplier keeps gradients from vanishing as temperature rises.
def scaled_and_unscaled_gradients(student_logits, teacher_logits, temperature):
_, scaled = distillation_loss_and_grad(student_logits, teacher_logits,
temperature, scale_t2=True)
_, unscaled = distillation_loss_and_grad(student_logits, teacher_logits,
temperature, scale_t2=False)
return scaled, unscaled
The tests assert that the scaled gradient is exactly times the unscaled gradient and check it against finite differences.
For majority vote with and , sum the cases with three, four, or five correct samples:
For best-of- with a perfect selector,
def vote_values():
return majority_vote_accuracy(0.6, 5), best_of_n_accuracy(0.6, 5)
The second number is larger because it assumes a verifier can find one correct answer among the samples; majority voting has no such selector.
The implementation computes softened distributions, the -scaled KL loss, and its analytic gradient. It also computes best-of- directly and majority vote by summing the binomial tail from strict majority to .
def distillation_loss_and_grad(student_logits, teacher_logits, temperature=1.0,
scale_t2=True):
"""KL teacher_T || student_T, with optional T^2 multiplier."""
student_logits = np.asarray(student_logits, dtype=np.float64)
teacher_logits = np.asarray(teacher_logits, dtype=np.float64)
teacher = softmax(teacher_logits / temperature)
log_student = log_softmax(student_logits / temperature)
loss = -np.sum(teacher * log_student)
grad = (softmax(student_logits / temperature) - teacher) / temperature
if scale_t2:
loss *= temperature ** 2
grad *= temperature ** 2
return float(loss), grad
The tests gradient-check the distillation loss, verify the scaling relation, and assert the exact binomial values from Exercise 38.3.
39 Decoding & Speculative Sampling
Top-k with keeps the first two tokens and renormalizes to . Top-p with threshold 0.70 also keeps the first two, because is not enough and crosses the threshold. Min-p with keeps tokens with probability at least , so it keeps the first three and gives . The two-token filters are more peaked.
For each token, . Summing over tokens gives , so the positive and negative mismatch masses are equal. Also , hence . Multiplying the residual distribution by this rejection probability leaves , which added to equals .
The draft accepts at least one token with probability , at least two with probability , and all three with probability . Therefore . The test simulates 200,000 independent verification steps and asserts a mean within 0.01 of 1.64.
The implementation treats the prefix as enough state for this tiny grammar: start allows [, [ or , allows a digit, a digit allows ] or ,, and a closed bracket allows nothing.
def tiny_json_number_mask(prefix, vocab):
"""Allowed tokens for a tiny grammar: '[' digit (',' digit)* ']'."""
if not prefix:
return np.array([token == "[" for token in vocab])
if prefix[-1] in {"[", ","}:
return np.array([token.isdigit() for token in vocab])
if prefix[-1].isdigit():
return np.array([token in {"]", ","} for token in vocab])
return np.zeros(len(vocab), dtype=bool)
def grammar_step(logits, prefix, vocab):
return softmax(mask_logits(logits, tiny_json_number_mask(prefix, vocab)))
The tests check that after [3 only ] and , have nonzero probability, and that digit-only masks put zero probability on non-digits.
40 Quantization & Serving
For 3 signed bits, . The scale is . Rounding gives integer codes . Dequantization multiplies by , giving . The chapter test asserts these codes and values.
The affine quantizer is . Requiring to map to gives , hence . The value must be rounded because integer hardware stores an integer zero point, and clipped because rounding can move it outside the available unsigned code range. With and 2 bits, the tested scale is and .
If has large off-diagonal entries, an error in one coordinate can be offset by changing a correlated coordinate. Independent rounding ignores those cross terms. GPTQ quantizes one coordinate, computes its error, and updates later coordinates with a factor from before they are quantized.
def gptq_quantize_vector(w, hessian, bits=2, damping=1e-8):
"""Quantize coordinates and compensate later ones with H^{-1}."""
w = np.asarray(w, dtype=np.float64)
hessian = np.asarray(hessian, dtype=np.float64)
inv_h = np.linalg.inv(hessian + damping * np.eye(len(w)))
work = w.copy()
quantized = np.zeros_like(work)
qmax = 2 ** (bits - 1) - 1
scale = np.max(np.abs(w)) / qmax
for i in range(len(w)):
qi = np.clip(np.round(work[i] / scale), -qmax, qmax) * scale
error = work[i] - qi
quantized[i] = qi
if i + 1 < len(w):
work[i + 1:] -= error * inv_h[i + 1:, i] / inv_h[i, i]
return quantized
def reconstruction_loss(w, q, hessian):
error = np.asarray(w) - np.asarray(q)
return float(error @ hessian @ error)
The test uses correlated calibration inputs and 2-bit weights; the GPTQ reconstruction loss is less than one quarter of round-to-nearest loss.
For and two bytes per parameter, one token reads bytes and performs about FLOPs, so the arithmetic intensity is FLOP/byte. At 3 TB/s, the bandwidth roof is tokens/s before KV-cache traffic and overhead. Int4 halves the weight bytes, so the weight-only intensity doubles and the bandwidth roof doubles, provided the kernel can consume packed int4 efficiently.
41 Training at Scale
Mixed-precision Adam stores bf16 weights, bf16 gradients, fp32 master weights, and two fp32 moments, so one million parameters use bytes. The activation calculator gives bytes when it saves one tensor per layer. With three checkpoint segments, the tested calculator saves segment boundaries plus one segment interior during recomputation, giving 1792 bytes.
A ring all-reduce is reduce-scatter plus all-gather. Each phase sends of the tensor per rank, so together they send . For and , the formulas give bytes for ordinary data parallelism, for ZeRO-1, for ZeRO-2, and for ZeRO-3. The test asserts these exact values.
Write and . After the elementwise activation, . Matrix multiplication by the row-split second weight gives . The implementation generalizes this sum to any number of ranks.
def gelu(x):
return 0.5 * x * (1.0 + np.tanh(np.sqrt(2 / np.pi) *
(x + 0.044715 * x ** 3)))
def mlp(x, w1, b1, w2, b2):
return gelu(x @ w1 + b1) @ w2 + b2
def tensor_parallel_mlp(x, w1, b1, w2, b2, ranks):
"""Column-parallel first layer, row-parallel second layer."""
w1_parts = np.array_split(w1, ranks, axis=1)
b1_parts = np.array_split(b1, ranks)
w2_parts = np.array_split(w2, ranks, axis=0)
partials = []
for w1_i, b1_i, w2_i in zip(w1_parts, b1_parts, w2_parts):
partials.append(gelu(x @ w1_i + b1_i) @ w2_i)
return np.sum(partials, axis=0) + b2
The bubble fraction is . Source rank 0 sends two tokens to rank 0’s experts and one token to rank 1’s experts; source rank 1 sends all three tokens to rank 1. The count matrix is therefore . Block FP8 scaling helps because the small first block gets its own scale instead of sharing a scale with values around 20; the test checks that this local scaling has lower MSE than one global block.
42 Tool Use & Agent Loops
The required argument list is just expression, and its schema says the value must be a string.
The integer 3 is rejected because the calculator parses text arithmetic; accepting mixed types
would make the tool contract ambiguous.
def required_argument_names(schema):
return tuple(schema.get("required", ()))
The first action is lookup({"key": "paris_population_millions"}), whose observation is
{"ok": True, "content": "2.1"}. The second action is
calculate({"expression": "2.1 + 2"}), whose observation is {"ok": True, "content": "4.1"}.
The next model message is the final answer 4.1 million, so the stop reason is final.
def answer_population_question():
return react_loop("What is the Paris population plus two?", ScriptedModel())
The request is a JSON-RPC object with method tools/call and params naming the tool and arguments:
{"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": {"name": "calculate", "arguments": {"expression": "(8 - 3) / 2"}}}.
The result is {"ok": True, "content": "2.5"}.
def mcp_calculate(expression):
client = MCPClient(MCPServer())
return client.request("tools/call", {
"name": "calculate",
"arguments": {"expression": expression},
})["result"]
The wrapper below returns the lookup result exactly as an observation string. It does not parse the text for commands, call a tool named inside the text, or promote the text to a developer message. A real host would also label the channel in the prompt and keep tool permissions narrow.
def observation_as_data(key):
"""Return lookup text without treating it as a developer instruction."""
result = call_tool("lookup", {"key": key})
if not result["ok"]:
return "missing"
return result["content"]
43 Retrieval, Memory, Planning & Evaluation
The chunks are one two three four five six and five six seven eight nine ten. The overlap is
five six, so a fact that crosses the first boundary is still visible in a later retrieved chunk.
def overlapping_chunks(text):
return chunk_words(text, size=6, overlap=2)
There are two-sample subsets. Since samples are wrong, subsets fail completely. The estimator is .
def failed_subset_ratio(n, c, k):
return 1.0 - pass_at_k(n, c, k)
The tests also enumerate all correctness patterns and check that the estimator’s expectation is .
def exact_unbiased_value(n, k, p):
return expected_pass_at_k(n, k, p)
If the judge always favors the first position, one order alone confounds quality with placement.
Let judge(a, b) return the score difference in that displayed order. Evaluating both orders and
using (judge(a, b) - judge(b, a)) / 2 cancels a constant first-position bonus. For lengths 4 and 2
with a +1 first-position bonus, the debiased difference is 2.
def debiased_pairwise_score(judge, answer_a, answer_b):
first = judge(answer_a, answer_b)
second = judge(answer_b, answer_a)
return (first - second) / 2
The prompt should label retrieved text as context and tell the model to use only that context for the answer. This does not make the context true or safe, but it prevents retrieved text from being silently promoted to a developer instruction.
def tiny_rag_prompt(question, documents):
return retrieval_prompt(question, documents, k=1)
44 Capstone: An LLM End to End
With , 16 query heads, and 4 key-value heads (), a token at position passes through:
-
An ID, a scalar, which selects a row of the embedding table: shape (2048,).
-
RMSNorm: (2048,).
-
The query projection: (16, 128). The key and value projections: (4, 128) each. RoPE rotates the queries and keys.
-
Attention: each query head scores the cached keys of its group, giving (16, t) scores and weights. The weighted values are (16, 128), concatenated to (2048,) and projected back to (2048,).
-
The residual add, then RMSNorm: (2048,). SwiGLU’s two branches: (5461,) each. The down projection: (2048,), then the residual add.
-
After 16 blocks and a final RMSNorm, the tied output head gives logits of shape (32768,).
The query and output projections are each, which contributes . Keys and values map to each, which contributes . SwiGLU has two input matrices of shape and one output matrix of shape , for . With , , so the block has . The tests compare the formula with an explicit sum over all seven weight matrices, using integer widths; the two agree to within 0.1%.
Substituting gives , so and :
def allocate(flops, tokens_per_parameter=20):
"""Split a compute budget C = 6ND with D = 20N: N = sqrt(C / 120)."""
parameters = math.sqrt(flops / (6 * tokens_per_parameter))
return parameters, tokens_per_parameter * parameters
A budget of FLOPs buys about 2.89 billion parameters trained on 57.7 billion tokens. Both grow as , so ten times the compute buys about 3.2 times the parameters and 3.2 times the data.
Subtract the bf16 weights, bytes, from the device memory. Then divide by the KV cache of one sequence, bytes:
def concurrent_sequences(memory_bytes, context, kv_heads):
"""Sequences whose KV cache fits beside bf16 weights in a memory budget."""
head_width = EXAMPLE["width"] // EXAMPLE["heads"]
weights = 2 * transformer_parameters(**{**EXAMPLE, "kv_heads": kv_heads})
per_sequence = kv_cache_bytes(EXAMPLE["layers"], kv_heads, head_width, context)
return int((memory_bytes - weights) // per_sequence)
With 4 key-value heads, 292 sequences of 8,192 tokens fit. With 16 heads, only 72 fit: the cache per sequence is four times larger, and the weights grow slightly. Grouped-query attention therefore quadruples the batch a device can serve, which is its main purpose.
A Notation & Shapes
, , , and . Gradients take the shape of their variable, so and . The parameters number . The batch size does not appear: the same weights serve every example.
Transpose the weight and keep the bias:
def from_pytorch_linear(weight, bias):
"""nn.Linear stores weight as (d_out, d_in) and computes x @ weight.T + bias."""
return weight.T, bias
Since weight is , its gradient is the transpose of :
, with shape .
The tests confirm that the two layouts give identical outputs and gradients.
appears in every output of row : , so , and outputs of other rows do not depend on it. By the chain rule,
To check it, fix a random and treat
as the loss. Its exact gradient is affine_backward(G, X, W)[0], which check_gradient compares
with central differences. The book’s tests run this check for , , and
.
Output depends only on , so the Jacobian is diagonal: . The product is :
def elementwise_vjp(grad_y, x, derivative):
return grad_y * derivative(x) # O(n): the Jacobian is never built
def elementwise_vjp_dense(grad_y, x, derivative):
jacobian = np.diag(derivative(x)) # (n, n), zero off the diagonal
return jacobian.T @ grad_y # O(n^2) memory for the same answer
The dense Jacobian stores numbers, all but of them zero. For one layer of a small model with a million activations, that is numbers, or 4 TB in float32. The direct product stores .
B NumPy for Deep Learning
Right-align the shapes and treat missing leading axes as 1:
-
A + B: (4, 1, 3) with (1, 5, 1) gives (4, 5, 3). -
A + C: (4, 1, 3) with (1, 1, 3) gives (4, 1, 3). -
B + C: (5, 1) with (1, 3) gives (5, 3). -
A * D: (4, 1, 3) with (1, 4, 3) gives (4, 4, 3).
D + B aligns (4, 3) with (5, 1). The last axes are compatible (3 against 1), but the next pair
is 4 against 5. Neither is 1, so NumPy raises a ValueError. np.broadcast_shapes answers
these questions without allocating anything.
X.sum(axis=1) has shape (N,), and broadcasting aligns it with the last axis of X:
keepdimsdef normalize_rows_wrong(X):
return X / np.sum(X, axis=1) # (N, d) / (N,): aligns N with the d axis
-
If N ≠ d and d ≠ 1, the shapes (N, d) and (N,) are incompatible and NumPy raises an error.
-
If N = d, it runs and computes , dividing column by the sum of row . The rows no longer sum to one.
-
If d = 1, (N, 1) and (N,) broadcast to (N, N): a matrix where a column was expected.
Write X.sum(axis=1, keepdims=True), or equivalently X.sum(axis=1)[:, None].
Each output depends on the bias through , so is 1 when and 0 otherwise. The chain rule sums over every output:
In general, broadcasting is a linear map that copies each input entry to several output positions. For a linear map, the gradient with respect to the input is the adjoint applied to the upstream gradient: . Expand the inner product by grouping the output positions that copy the same input entry:
Here is the input entry that output copies. The inner sum adds
over exactly the positions that copied . Those are the leading axes that
broadcasting added and the size-1 axes it stretched, which is precisely what unbroadcast
sums. So unbroadcast is , and the book’s tests check this adjoint
identity with random and .
Factor out of the sum and take logarithms:
With , every term lies in , and the maximizing term equals 1. The sum therefore lies between 1 and . Its logarithm lies between 0 and , which gives both bounds. For the gradient,
The tests confirm this with check_gradient, comparing against exp(log_softmax(z)).
Expand both sides to third order:
Subtracting cancels and . The two third-order terms average to for some between them, by the intermediate value theorem. Dividing by gives (B.6).
For , with and , set to get . At the optimum both terms scale as :
In float64, and the error is about . In float32, and the error is about . Figure B.3 shows both floors.
The lookup table[ids] is the matrix product , where has a
single 1 per row, in the column of that row’s token. The gradient with respect to the table is
therefore . The loop and np.add.at compute the same sum one row at a
time:
def embedding_backward_loop(upstream, ids, vocabulary_size):
gradient = np.zeros((vocabulary_size, upstream.shape[-1]), dtype=upstream.dtype)
for token, row in zip(ids.reshape(-1), upstream.reshape(-1, upstream.shape[-1])):
gradient[token] += row
return gradient
def embedding_backward_one_hot(upstream, ids, vocabulary_size):
one_hot = np.eye(vocabulary_size, dtype=upstream.dtype)[ids.reshape(-1)] # (n, V)
return one_hot.T @ upstream.reshape(-1, upstream.shape[-1]) # (V, d)
def embedding_backward_buggy(upstream, ids, vocabulary_size):
gradient = np.zeros((vocabulary_size, upstream.shape[-1]), dtype=upstream.dtype)
# Repeated ids collide: only the last write to each row survives.
gradient[ids.reshape(-1)] += upstream.reshape(-1, upstream.shape[-1])
return gradient
gradient[ids] += upstream means gradient[ids] = gradient[ids] + upstream. The right-hand side
gathers a (possibly repeated) row for each ID and adds its upstream row. The assignment then
writes those rows back in order. A repeated ID is written several times, and only the last
write survives. Rows of tokens that appear at most once are correct: once-used tokens get their
single contribution, and unused tokens stay zero. The tests check all of this with a sequence
in which token 1 appears three times.
def split_heads_wrong(X, heads):
B, T, width = X.shape
# Reinterprets memory in order, mixing tokens and heads.
return X.reshape(B, heads, T, width // heads)
In memory, each batch entry of a C-ordered (B, T, H·d_h) array stores token 0’s whole
feature vector, then token 1’s, and so on. A reshape reads those numbers in order.
reshape(B, H, T, d_h) gives head 0 the first numbers. With , three
heads, and , that is all 12 features of token 0 and the first 8 of token 1: a
mixture of tokens. reshape(B, T, H, d_h) only splits the last axis, so head of token
is the contiguous chunk X[b, t, h*d_h:(h+1)*d_h]. The transpose then moves the head
axis forward by changing strides, without mixing anything. Any random input exposes the
difference.
def running_sum(value, steps, round_to):
total = round_to(np.float32(0))
for _ in range(steps):
total = round_to(total + round_to(np.float32(value)))
return float(total)
def bfloat16_vs_float32(value=1e-3, steps=10_000):
bfloat16_total = running_sum(value, steps, round_to_bfloat16)
float32_total = running_sum(value, steps, np.float32)
return bfloat16_total, float32_total
The bfloat16 total stops at 0.5, the float16 total at 4.0, and the float32 total reaches about 10.0004. A sum stops growing once the addend is less than half the gap between neighbouring numbers at the current total, because rounding then returns the old total.
-
bfloat16 has 7 fraction bits. On the gap is , and half of it, 0.00195, exceeds 0.001. Just below 0.5 the gap is . Each addition there is slightly more than half a gap, so it rounds up by a whole gap. The total overshoots and reaches 0.5 after about 383 steps instead of 500.
-
float16 has 10 fraction bits, so the same thing happens eight times higher. At 4 the gap is .
-
float32 has a gap of about near 10, so each addition survives with a tiny rounding error. That error accumulates to the final 0.0004.
This is why mixed-precision training keeps master weights, optimizer state, and reductions in float32 even when matrix products run in 16 bits.
C Matrix Calculus Cookbook
For each row , . Therefore , and the VJP sums upstream gradients over the broadcasted batch axis: .
def broadcast_add_gradient(grad, x, bias):
del x
return grad.sum(axis=0, keepdims=True).reshape(bias.shape)
For ,
which is . Similarly,
which is .
def matmul_shapes(a, b, grad):
grad_a, grad_b = matmul_vjp(grad, a, b)
return grad_a.shape, grad_b.shape
For , . Then
def softmax_jacobian_times_vector(logits, vector):
y = softmax_forward(logits)
return softmax_vjp(vector, y)
The attention VJP is the reverse of the forward decomposition: , , . The query gradient is . The test file checks this function against central differences for , , and .
def attention_query_gradient(q, k, v, grad):
grad_q, _, _ = attention_vjp(grad, q, k, v)
return grad_q