Skip to content

The op library and gradcheckall

ModuleL0.2 · build · Python · Pass 2 · 5 to 7 h
You buildpython/tinyllm/autograd/functional.py, imported as F: 25 ops (activations, reductions, shape, selection, the softmax family, dropout, matmul) and gradcheck_all()
Contractcourse/contracts/py/tinyllm/autograd/functional.pyi
Testscourse/tests/L0.2/ (what they check: section 4), golden values from torch 2.14 in course/fixtures/L0.2/ops_torch.npz
NeedsL0.1 Tensor and from_op · M01.3 activations and their derivatives · M09.2 stable softmax, log-softmax, logsumexp · M08.3 the softmax VJPs and unbroadcast · M08.1 dual numbers · M04.1 gradcheck · M06.3 PCG32 (or --ref-deps)
Used byL0.4 layers are written in F · L0.5 BigramLogits is an F.embedding · L0.6 the checkpoint tests compute their loss with F · later: L2.2, L3.1, L3.2, L3.3, L3.4, L3.6, L4.1, L4.2, L4.3, L5.3, L5.4, L5.5, L6.1, L6.2, L6.3, L6.5, L6.6, L7.1, L7.2, L7.3, L7.5, L7.6, L7.8, L7.9
MilestoneMS-L0 (step 1: {tinyllm} gradcheck --suite all)
Optional depthGriewank and Walther, Evaluating Derivatives (2nd ed.), ch. 3 and 4; the PyTorch derivatives.yaml file, read as a table of VJPs
  • Every op is one numpy forward plus one closed-form VJP joined by from_op; nothing in the library calls another op’s backward (test_gradcheck_each_op, test_matches_torch).
  • Reductions put the reduced axes back before broadcasting the gradient, a maximum shares its gradient among ties, and a mean divides by the reduced count only (test_max_ties_share_the_gradient, test_var_correction).
  • Selection ops route gradient by addition: a token or index used twice gets both gradients, and an out-of-range id is an error where it enters (test_embedding_repeated_ids_accumulate, test_embedding_rejects_bad_ids).
  • The softmax family’s VJPs are computed from the stable outputs, so logits of 1000 and fully masked rows give finite gradients (test_logsumexp_no_overflow, test_softmax_fully_masked_row).
  • gradcheck_all checks every op against central differences with a random output weighting, plus forward-mode duals for the elementwise ops, and it fails on a planted 0.1% or 1e-7 error (test_gradcheck_all_catches_a_wrong_op).
Terminal window
ol start L0.2 # stubs functional.py into your repo
ol tests L0.2 # read the test catalog first: rung R0
ol check L0.2 # exit code is the verdict
ol check L0.2 --ref-deps # only if a math dependency is not passing yet
ol diff L0.2 # after passing: your code against the reference

Your CLI gains its first Pass 2 verb, {tinyllm} gradcheck --suite all (spec/cli-roles.md, fixed by MS-L0): call F.gradcheck_all(rtol=1e-5), print one line per op, and end with {"suite": "all", "checks": N, "failed": K, "max_rel_err": E, "worst": "<op>"}; exit 0 only when failed is 0.


L0.1 gave you a Tensor that knows + - * / @ ** and indexing. A model needs more: a nonlinearity between layers, a softmax over the vocabulary, an embedding lookup for token ids, a mean for LayerNorm, a dropout mask. Written inline in each model, each of those is a chance to get a VJP wrong silently, because a wrong gradient still trains, just badly. This module puts every op in one library with one VJP each, checks each against torch and against numerical differentiation, and gives your CLI the gradcheck verb MS-L0 runs first. From here on, model code (L0.4, L2, L3, L5) composes these ops and never writes a backward pass.

SymbolMeaningType / shape
xxan op’s inputTensor, any shape
y=f(x)y = f(x)its outputTensor
yˉ\bar{y}, xˉ\bar{x}upstream gradient and the gradient handed back (L0.1)shapes of yy, xx
AAthe reduced axes of a reductiontuple of ints
n=∏a∈Asan = \prod_{a \in A} s_anumber of elements reduced into each outputint
softmax(x)i=exi/∑jexj\mathrm{softmax}(x)_i = e^{x_i} / \sum_j e^{x_j}along one axissame shape as xx
⟨u,v⟩=∑iuivi\langle u, v \rangle = \sum_i u_i v_iinner product along the softmax axis
ϕ\phi, Φ\Phistandard normal density and CDF (gelu)functions
ppdropout probabilityfloat in [0,1][0, 1]
uku_kthe kk-th uniform draw of a PCG32 (M06.3)float in [0,1)[0, 1)
wwrandom output weights inside gradcheck_allshape of yy
ϵ\epsiloncentral-difference step, 10−610^{-6} (M04.1)float

Elementwise ops. For y=f(x)y = f(x) applied entry by entry, xˉ=yˉ⊙f′(x)\bar{x} = \bar{y} \odot f'(x). The derivatives are M01.3’s: exp⁡′=exp⁡\exp' = \exp (reuse the output), log⁡′=1/x\log' = 1/x, tanh⁡′=1−tanh⁡2\tanh' = 1 - \tanh^2, σ′=σ(1−σ)\sigma' = \sigma(1 - \sigma), relu′(x)=[x>0]\mathrm{relu}'(x) = [x > 0] (0 at the kink, as torch), silu(x)=xσ(x)\mathrm{silu}(x) = x\sigma(x), and two gelu forms: exact xΦ(x)x\Phi(x) and the tanh approximation x2(1+tanh⁡(2/π(x+0.044715x3)))\tfrac{x}{2}(1 + \tanh(\sqrt{2/\pi}(x + 0.044715 x^3))). approximate="none" or "tanh" picks one; anything else is an error, because HF configs name exactly these two.

Reductions put axes back. sum(x, axis=A) adds every input element into one output with weight 1, so xˉ\bar{x} is yˉ\bar{y} copied back over AA. Without keepdims the output lost those axes; np.expand_dims(g, A) restores them as size 1, and broadcasting copies. mean is sum times 1/n1/n with nn the size of the reduced axes only. max routes the gradient to the entries equal to the maximum; with ties torch (amax) splits it equally, so the shares still add up to yˉ\bar{y}. var with correction cc is ∑(x−μ)2/(n−c)\sum (x - \mu)^2 / (n - c): c=0c = 0 is the population variance LayerNorm uses, c=1c = 1 the sample variance. Its gradient is yˉ⋅2(xi−μ)/(n−c)\bar{y} \cdot 2(x_i - \mu)/(n - c); the terms through μ\mu vanish because ∑i(xi−μ)=0\sum_i (x_i - \mu) = 0.

Shape ops move gradients back. reshape reshapes yˉ\bar{y} to the input shape. transpose(a, b) is its own inverse. permute(dims) is undone by the inverse permutation, argsort(dims), not by applying dims again. concat splits yˉ\bar{y} at the cumulative sizes of its inputs, and stack takes slice ii along the stacking axis for input ii.

Selection ops add. where(cond, a, b) sends yˉ\bar{y} to aa where cond holds and to bb elsewhere, then unbroadcasts each. masked_fill(x, mask, v) sends nothing to the filled positions: an attention mask must not train the scores it hid. gather(x, idx, axis) (torch’s) and embedding(W, ids) read entries, possibly the same one twice, so their VJPs are scatter-adds (np.add.at). Both check their indices: numpy would read -1 as the last row and never complain.

The softmax family, stably. Forward values come from M09.2 (subtract the max first). The VJPs come from M08.3 and use the outputs:

xˉ=y⊙(yˉ−⟨yˉ,y⟩)(softmax),xˉ=yˉ−softmax(x) ∑iyˉi(log_softmax),\bar{x} = y \odot (\bar{y} - \langle \bar{y}, y \rangle) \quad (\mathrm{softmax}), \qquad \bar{x} = \bar{y} - \mathrm{softmax}(x)\,\textstyle\sum_i \bar{y}_i \quad (\mathrm{log\_softmax}),

and for ℓ=logsumexp(x)\ell = \mathrm{logsumexp}(x), xˉ=ℓˉ softmax(x)\bar{x} = \bar{\ell}\, \mathrm{softmax}(x). None of these ever computes exe^{x} of a raw logit, so a logit of 1000 is as safe as one of 1. A row masked entirely to −∞-\infty has softmax zeros (M09.2’s rule) and therefore a zero gradient.

Inverted dropout. During training each element is kept with probability 1−p1 - p and scaled by 1/(1−p)1/(1 - p), so the expected output equals the input and evaluation needs no rescaling: in eval mode, or with p=0p = 0, dropout returns its input object and draws nothing. The mask uses one uniform per element in C order, keepk=[uk≥p]\text{keep}_k = [u_k \ge p], from the PCG32 you pass in. The draw count is fixed (exactly x.size draws per call), which is what lets a resumed run (L0.6) reproduce the same masks. p=1p = 1 drops everything (no division by zero).

gradcheck_all: a test that can fail. For each op, gradcheck_all draws float64 inputs from PCG32(0), keeps them away from kinks (relu at 0, max ties), and compares the analytic gradient of f(x)=∑y⊙wf(x) = \sum y \odot w with central differences (f(x+ϵei)−f(x−ϵei))/2ϵ\big(f(x + \epsilon e_i) - f(x - \epsilon e_i)\big)/2\epsilon through M04.1’s gradcheck at rtol. The weights ww are random on purpose: ∑isoftmax(x)i=1\sum_i \mathrm{softmax}(x)_i = 1 for every xx, so the gradient of the plain sum is zero and any softmax VJP, right or wrong, would pass. Central differences resolve relative errors near 10−510^{-5}; a smaller error in an elementwise derivative is caught by a second check, forward-mode dual numbers (M08.1), which are exact to rounding and flag a gap above 10−910^{-9}. Ops are looked up in the module when the check runs, so a monkeypatched op is the one checked.

Softmax and its VJP. Take x=[0,ln⁡2,ln⁡3]x = [0, \ln 2, \ln 3]. Then ex=[1,2,3]e^{x} = [1, 2, 3], the sum is 6, and y=[1/6,2/6,3/6]y = [1/6, 2/6, 3/6]. With the upstream gradient yˉ=[1,0,0]\bar{y} = [1, 0, 0] (the loss looks only at the first output):

  1. ⟨yˉ,y⟩=1⋅1/6=1/6\langle \bar{y}, y \rangle = 1 \cdot 1/6 = 1/6.
  2. yˉ−1/6=[5/6,−1/6,−1/6]\bar{y} - 1/6 = [5/6, -1/6, -1/6].
  3. Multiply by yy: xˉ=[5/36,−2/36,−3/36]\bar{x} = [5/36, -2/36, -3/36].

The entries sum to zero: raising every logit by the same amount leaves softmax unchanged, so no direction along (1,1,1)(1, 1, 1) can change the loss. Check xˉ0\bar{x}_0 by perturbation: y0=ex0/(ex0+5)y_0 = e^{x_0}/(e^{x_0} + 5), whose derivative at x0=0x_0 = 0 is 5/365/36.

A tie in max. max⁡([1,3,3])=3\max([1, 3, 3]) = 3 is reached twice. torch’s amax gives each maximal entry half: xˉ=[0,0.5,0.5]\bar{x} = [0, 0.5, 0.5] for yˉ=1\bar{y} = 1.

Variance. For [1,2,3,4][1, 2, 3, 4]: μ=2.5\mu = 2.5, squared deviations 2.25+0.25+0.25+2.25=52.25 + 0.25 + 0.25 + 2.25 = 5, so var is 5/4=1.255/4 = 1.25 with correction 0 and 5/35/3 with correction 1.

These are test_hand_example_softmax_backward, test_max_ties_share_the_gradient, and test_var_correction.

# python/tinyllm/autograd/functional.py (from tinyllm.autograd import functional as F)
exp, log, tanh, sigmoid, relu, silu; gelu(x, approximate="none" | "tanh")
sum(x, axis=None, keepdims=False); mean(...); max(...); var(x, axis=None, keepdims=False, correction=0)
reshape(x, shape); transpose(x, a, b); permute(x, dims); concat(xs, axis=0); stack(xs, axis=0)
where(cond, a, b); gather(x, idx, axis); embedding(weight, ids); masked_fill(x, mask, value)
softmax(x, axis=-1); log_softmax(x, axis=-1); logsumexp(x, axis=-1, keepdims=False)
dropout(x, p, training, rng); matmul(a, b)
def gradcheck_all(rtol: float = 1e-5) -> dict[str, GradcheckReport]

sum, max, and mean shadow Python builtins inside the module; use builtins.max where you need the builtin. Float32 inputs give float32 outputs and gradients (cast derivative arrays and masks to the input’s dtype).

TestKINDChecksWhy it matters downstream
test_hand_example_softmax_backwardunitsection 3: y=[1/6,2/6,3/6]y = [1/6, 2/6, 3/6], xˉ=[5/36,−2/36,−3/36]\bar{x} = [5/36, -2/36, -3/36]you and the test agree on the softmax VJP
test_max_ties_share_the_gradientboundarymax([1, 3, 3]) gives [0,0.5,0.5][0, 0.5, 0.5]ReLU-max pooling and torch parity
test_matches_torchgoldenvalue and gradients equal torch 2.14 on 45 casesmodels ported from torch (L5 to L7)
test_gradcheck_each_opgradcheckevery VJP against the frozen central differences, float64a wrong VJP trains badly, not visibly
test_float32_in_float32_outboundaryfloat32 in, float32 out for twelve ops, and a float32 gradient through dropoutPython agrees with the float32 C kernels (L9)
test_relu_derivative_at_zero_is_zeroboundaryrelu′(0)=0\mathrm{relu}'(0) = 0torch’s convention
test_embedding_repeated_ids_accumulateboundarya repeated token gets both rows of gradientevery embedding table
test_gather_repeated_indices_accumulateboundaryan index read twice gets gradient 2the cross-entropy gather in L0.3
test_embedding_rejects_bad_idsboundary-1, n, and float ids are a ValueErrortokenizer bugs fail where they enter
test_dropout_eval_is_identityuniteval mode returns the input object, draws nothingevaluation is deterministic
test_dropout_mask_and_scalestatisticalabout 75% kept at p=0.25p = 0.25, survivors scaled by 4/34/3the expected output is the input
test_dropout_draw_accountingunitthe mask is u≥pu \ge p for the same-seed uniforms, exactly x.size drawsa resumed run replays the masks (L0.6)
test_dropout_p_boundsboundaryp=1p = 1 gives zeros; p∉[0,1]p \notin [0, 1] is an errorconfig errors fail early
test_softmax_fully_masked_rowboundaryan all-masked row gives zeros and a zero gradientpadded rows in attention (L5.2)
test_logsumexp_no_overflowboundarylogsumexp([1000, 1000]) and its gradient [0.5,0.5][0.5, 0.5]trained logits are large
test_var_correctionunit1.25 and 5/35/3; n≤cn \le c is an errorLayerNorm’s population variance (L0.4)
test_where_and_masked_fill_route_gradientsunitgradients go to the chosen operand; filled positions get noneattention masks (L5)
test_gelu_rejects_unknown_approximationboundaryonly "none" and "tanh"HF configs name exactly these
test_ops_build_no_graph_under_no_gradpropertyevery op returns a constant under no_gradevaluation memory
test_gradcheck_all_reports_every_opunitone ok report per op, 33 namesMS-L0’s gradcheck --suite all
test_gradcheck_all_catches_a_wrong_opboundarya 0.1% softmax error and a 10−710^{-7} sigmoid error are both flaggeda checker that always passes is worse than none
PitfallSymptomCaught by
1. out[ids] = g in embedding or gatherfrequent tokens train slower than rare onestest_embedding_repeated_ids_accumulate (mutant s06), test_gather_repeated_indices_accumulate (mutant s07)
2. trusting ids-1 silently trains the last row of the tabletest_embedding_rejects_bad_ids (mutant s08)
3. ties in max: every maximal entry gets the full gradient, or only the firstgradients double, or disagree with torchtest_max_ties_share_the_gradient (mutants s03, s04)
4. the logsumexp VJP as exp(x) / exp(y)inf / inf = nan at logit 1000test_logsumexp_no_overflow (mutant s13)
5. feeding the softmax VJP the input xx instead of the output yywrong gradients everywhere, nan on masked rowstest_softmax_fully_masked_row, test_hand_example_softmax_backward (mutant s01)
6. a gradcheck_all that cannot fail (no dual check, or a loose tolerance)MS-L0 step 1 passes with a broken optest_gradcheck_all_catches_a_wrong_op (mutants s19, s20)
7. float64 leaking from a mask or a derivative arraya float32 model becomes float64, twice the memorytest_float32_in_float32_out (mutant m04)

| Forward | L2.2 | Registered call site uses this module. | | Forward | L3.1 | Registered call site uses this module. | | Forward | L3.2 | Registered call site uses this module. | | Forward | L3.3 | Registered call site uses this module. | | Forward | L3.4 | Registered call site uses this module. | | Forward | L3.6 | Registered call site uses this module. | | Forward | L4.1 | Registered call site uses this module. | | Forward | L4.2 | Registered call site uses this module. | | Forward | L4.3 | Registered call site uses this module. | | Forward | L5.3 | Registered call site uses this module. | | Forward | L5.4 | Registered call site uses this module. | | Forward | L5.5 | Registered call site uses this module. | | Forward | L6.1 | Registered call site uses this module. | | Forward | L6.2 | Registered call site uses this module. | | Forward | L6.3 | Registered call site uses this module. | | Forward | L6.5 | Registered call site uses this module. | | Forward | L6.6 | Registered call site uses this module. | | Forward | L7.1 | Registered call site uses this module. | | Forward | L7.2 | Registered call site uses this module. | | Forward | L7.3 | Registered call site uses this module. | | Forward | L7.5 | Registered call site uses this module. | | Forward | L7.6 | Registered call site uses this module. | | Forward | L7.8 | Registered call site uses this module. | | Forward | L7.9 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackL0.1every op is from_op(forward, parents, vjp) on its Tensor
BackM01.3the activation functions and their derivatives
BackM09.2stable softmax, log_softmax, logsumexp forwards
BackM08.3softmax_vjp, log_softmax_vjp, unbroadcast
BackM08.1derivative on dual numbers, the second check in gradcheck_all
BackM04.1gradcheck, the central-difference check
BackM06.3PCG32: dropout’s uniforms and gradcheck_all’s inputs
ForwardL0.4Linear, LayerNorm, Embedding, Dropout are compositions of F ops
ForwardL0.5BigramLogits.forward is F.embedding; the trainer reshapes with F.reshape
ForwardL0.6the checkpoint tests’ training loss is F.mean of a squared error

If you skip this module, ol check L0.4 stops with L0.4 needs L0.2: build it, or rerun with --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
one VJP per opPyTorch derivatives.yamla generated backward for about 2,000 ops, double-backward, forward-mode formulastools/autograd/derivatives.yaml
gradcheck_alltorch.autograd.gradcheck, gradgradcheckcomplex inputs, sparse layouts, second derivatives, fast mode with random projectionstorch/autograd/gradcheck.py
dropout with an explicit generatortorch.nn.functional.dropout, JAX random.bernoulliPhilox counter-based RNG on GPU, fused into matmul epiloguesaten/src/ATen/native/Dropout.cpp
stable softmax VJPfused softmax and log-softmax kernelsone pass over the row, online max (L9.2)aten/src/ATen/native/SoftMax.cpp