Skip to content

Oracles, golden and differential tests, gradcheck as a test (R5)

Modulecraft.05 · practice · Python · Pass 5 · 3 to 4 h
You buildprimers/craft.05/kernels.py (a kata: four forward passes of Parts 5 to 7 and their hand-written backward passes, in numpy) and primers/craft.05/test_oracles.py, your oracle tests for it
Contractthe rules in the kata’s docstring (ol check craft.05 writes the kata with stubs on its first run)
Testscourse/tests/craft.05/: the grade of your oracle tests by planted faults (section 4), with torch’s values and gradients in course/fixtures/craft.05/oracles.npz (course/oracle/craft.05/oracles_torch.py)
Needsreading: craft.03 how tests are graded · craft.04 laws and models · M04.1 central differences (your gradcheck, which these tests may use) · L7.1 RMSNorm and L7.2 SwiGLU, two of the kata’s ops
Used byno call site (a practice): rung R5 grades your oracle tests in L7.1 to L7.9, and the same three oracles prove your C kernels (L9) and Rust engine (L10.1)
MilestoneMS-P5 (the Pass 5 gate requires a passing craft.05)
Optional depthBarr et al., “The Oracle Problem in Software Testing: A Survey” (IEEE TSE, 2015); McKeeman, “Differential Testing for Software” (1998); the PyTorch torch.autograd.gradcheck documentation
  • A test needs an oracle: something other than the code under test that knows the right answer. Golden files (torch’s numbers), a second implementation (differential), and a mathematical identity (central differences for a gradient) are the three you will use for every kernel from here on.
  • A hand-written backward is where the subtle bugs live: a transposed weight gradient has the right shape on a square matrix and the wrong values (test_hand_example_transposed_gradient). Shapes are not an oracle.
  • Gradcheck is a test: central differences of ∑f(x)⊙g\sum f(x) \odot g for a random upstream gg check every input’s gradient, in float64, with a tolerance that float64 rounding cannot break.
  • Tolerances are part of the oracle: too tight and the correct kata fails (test_your_tests_accept_the_reference); too loose and a lost 1/d1/\sqrt{d} survives.
  • The grade: your tests must catch at least 80% of ten planted faults, and always the transposed weight gradient (test_planted_faults_are_caught).
Terminal window
ol check craft.05 # first run: writes primers/craft.05/kernels.py (stubs) and fails
# write primers/craft.05/test_oracles.py (section 4), then the kata; run your tests yourself:
TINYLLM_FIXTURES=$PWD/.ol/openlearn/course/fixtures uv run --no-project --with pytest --with numpy \
python -m pytest -q primers/craft.05
ol check craft.05 # grades your oracles against the planted faults

(ol check sets TINYLLM_FIXTURES itself; when you run pytest by hand, point it at the course’s fixtures/ directory of your openlearn checkout.) Write the tests first, against the stub: they must fail. ol check never shows a planted fault’s code; a survivor prints only its one-line description.


Everything you built in Parts 5 to 7 got its gradients from your autograd (L0.1, L0.2): you wrote each op’s VJP once and composed them. The C kernels of Part 9 and the Rust engine of Part 10 have no autograd and no torch to lean on. FlashAttention’s backward, a fused RMSNorm, a SwiGLU kernel: each is a forward and a backward written by hand, and each one is checked against your Python (P6). That only works if you know how to test a numerical kernel with nothing but an oracle. The course tests of L7.1 to L7.9 did it for you with frozen helpers; from this chapter on, rung R5 grades the oracle tests you write for them.

SymbolMeaningType / shape
ffthe op under test, x↦yx \mapsto yfunction
ggan upstream gradient, the shape of yyarray
L=∑f(x)⊙gL = \sum f(x) \odot ga scalar whose gradient with respect to xx is the VJP of ff at ggfloat
ϵ\epsilonthe central-difference step, 10−610^{-6} in float64float
eie_ithe unit vector of element iiarray

A golden test compares with numbers an independent implementation produced and a maintainer recorded: here torch’s float64 forward values and autograd gradients for fixed inputs (oracles.npz). It catches every bug that changes those numbers, including the ones that keep the shape. Its weakness is coverage: only the recorded inputs. Read the file through TINYLLM_FIXTURES, never through a path into the course tree.

For any differentiable ff and any gg, the VJP is the gradient of L(x)=∑f(x)⊙gL(x) = \sum f(x) \odot g:

∂L∂xi≈L(x+ϵei)−L(x−ϵei)2ϵ,error O(ϵ2)+O(u/ϵ).\frac{\partial L}{\partial x_i} \approx \frac{L(x + \epsilon e_i) - L(x - \epsilon e_i)}{2\epsilon}, \qquad \text{error } O(\epsilon^2) + O(u / \epsilon).

In float64 (u≈1.1×10−16u \approx 1.1 \times 10^{-16}) the step ϵ=10−6\epsilon = 10^{-6} balances truncation and rounding near 10−1010^{-10}, so a tolerance of 10−610^{-6} relative is safe for a correct backward and far below any real bug. A random gg (not all ones) matters: with g=1g = 1 a backward that sums where it should weigh can still pass. Check every input of the op (x, W, and b for a linear layer), with shapes where a transposed gradient would not fit and with square ones where it would.

Two implementations of the same contract must agree: the causal mask against attention over each query’s prefix, chunked against whole, C against Python, Rust against Python. A differential test needs no recorded numbers, so it covers any input you can generate. When both sides come from you, one side must also be checked against an oracle, or two equally wrong implementations pass together.

Exact comparison fails for correct code as soon as two implementations sum in a different order. Use the dtype’s rounding: float64 differences near 10−1510^{-15} relative for a few operations, more for long reductions (K\sqrt{K} growth, M09.3). Too loose a tolerance hides real bugs: a missing 1/d1/\sqrt{d} is a factor of 2 for d=4d = 4, but a dropped term of size 10−310^{-3} needs a tolerance below that.

y=xW⊤y = x W^\top with x=(1,2)x = (1, 2) and W=(1031)W = \begin{pmatrix} 1 & 0 \\ 3 & 1 \end{pmatrix}: y=(1⋅1+2⋅0, 1⋅3+2⋅1)=(1,5)y = (1 \cdot 1 + 2 \cdot 0,\ 1 \cdot 3 + 2 \cdot 1) = (1, 5). With upstream g=(1,0)g = (1, 0):

  • ∂L/∂W=g⊤x=(10)(1,2)=(1200)\partial L / \partial W = g^\top x = \begin{pmatrix} 1 \\ 0 \end{pmatrix} (1, 2) = \begin{pmatrix} 1 & 2 \\ 0 & 0 \end{pmatrix}: only row 0 of WW made y0y_0.
  • The transposed bug x⊤g=(1020)x^\top g = \begin{pmatrix} 1 & 0 \\ 2 & 0 \end{pmatrix} has the same 2×22 \times 2 shape and is wrong.
  • Central difference on W0,1W_{0,1}: L(W0,1±ϵ)=1±2ϵL(W_{0,1} \pm \epsilon) = 1 \pm 2\epsilon, so the quotient is 22, matching the correct entry and not the transposed one (0).

A shape check passes both; the golden file and the gradcheck reject the bug. This is test_hand_example_transposed_gradient.

The kata, primers/craft.05/kernels.py (its docstring is the spec), numpy float64:

def linear(x, W, b): ... # x @ W.T + b
def linear_backward(x, W, gy): ... # (gx, gW, gb)
def rmsnorm(x, w, eps=1e-6): ... # x / sqrt(mean(x^2) + eps) * w (L7.1)
def rmsnorm_backward(x, w, gy, eps=1e-6): ... # (gx, gw)
def swiglu(a, b): ... # silu(a) * b (L7.2)
def swiglu_backward(a, b, gh): ... # (ga, gb)
def attention(q, k, v, causal=False): ... # softmax(q k^T / sqrt(d)) v (L5.1)
def attention_backward(q, k, v, go, causal=False): ... # (gq, gk, gv)

Your tests, primers/craft.05/test_oracles.py: at least six test functions, importing the kata as kernels (and, if you like, your own tinyllm.num.gradcheck from M04.1); no torch, no autograd, no random. Each oracle catches some planted faults:

OracleFor
golden: every forward and every gradient equals oracles.npz (linear, linear_sq, rmsnorm, swiglu, attn, attn_causal)all four ops, square and non-square
gradcheck: central differences of ∑f⊙g\sum f \odot g for every input, random ggevery backward, on inputs you generate
differential: causal attention equals attention over each query’s prefixthe mask
stability: scores of 300 give finite outputs (the mean of the values)the softmax’s max subtraction

The check (ol check craft.05, course/tests/craft.05/check) runs six tests:

TestKINDChecks
test_hand_example_transposed_gradientunitsection 3 on the course’s kata
test_your_tests_are_oracle_testsunitthe file exists, at least 6 tests, reads the fixture through TINYLLM_FIXTURES, has a gradcheck, imports nothing forbidden
test_your_kata_matches_torchgoldenyour kata’s values and gradients equal torch’s
test_your_kata_passes_your_testsunityour oracles hold for your kata
test_your_tests_accept_the_referenceconformanceyour oracles hold for the course’s kata
test_planted_faults_are_caughtfaultagainst each of the 10 planted faults your tests fail; score at least 0.80 and s01 always caught

Each run copies your test file next to one version of the kata in a scratch directory and runs pytest in its own process group with a timeout; your tree is never touched.

PitfallSymptomCaught by
A transposed weight gradient, x⊤gx^\top g for g⊤xg^\top xright shape on a square WW, wrong valuesa golden test on the square linear case or a gradcheck with a square WW (mutant s01)
gW⊤g W^\top for the input gradientfits only a square WWthe golden test or a gradcheck with a non-square WW (mutant s02)
Averaging the bias gradient over the batchevery bias learns NN times too slowlythe golden test (mutant s03)
RMSNorm’s backward treating the rms as a constantthe gradient misses the −n⋅mean(un)-n \cdot \mathrm{mean}(u n) terma gradcheck of RMSNorm (mutant s04)
SwiGLU’s gate derivative as σ(a)\sigma(a) alonewrong by aσ(1−σ)a \sigma (1 - \sigma), small near 0a gradcheck with large aa (mutant s05)
A softmax backward without the row sumgq and gk off by a term that vanishes only when gg is uniformgradcheck with a random gg (mutant s06)
Losing 1/d1/\sqrt{d} in the query gradientgq too large by d\sqrt{d}the golden test (mutant s07)
gSqg_S q for gS⊤qg_S^\top q in the key gradientfits square scores onlythe causal golden case or a square gradcheck (mutant s08)
A softmax without the max subtractioninf / inf at scores near 700 in float64the stability test (mutant s09)
A causal mask that hides the diagonalquery 0 sees nothing: NaNthe causal golden case, the prefix differential (mutant s10)
g = ones in a gradchecksum-instead-of-weigh bugs passtest_planted_faults_are_caught reports the survivors
A tolerance tighter than float64 roundingthe correct kata failstest_your_tests_accept_the_reference
Comparing shapes onlyevery transposed bug on square inputs survivestest_planted_faults_are_caught (s01 is required)
DirectionModuleHow it uses this
Backcraft.03mutation grading, baselines A and B, required faults
Backcraft.04the model property is a differential oracle
BackM04.1central differences, the gradcheck these tests may use
ForwardL7.1 to L7.9your rung R5 suites: golden, differential, gradcheck
ForwardL9.2, L9.3, L9.6C kernels checked against your Python with the same three oracles
ForwardL10.1the Rust engine’s logits against your Python and HF’s fixture
Forwardcraft.06benchmarks: once the numbers are right, make them fast
Your pieceProduction equivalentWhat it addsWhere to look
your gradchecktorch.autograd.gradcheck, gradgradcheckcomplex inputs, fast mode (a random projection instead of every element), second derivativestorch/autograd/gradcheck.py
golden filesPyTorch OpInfoone description per op drives forward, backward, dtype, and device teststorch/testing/_internal/common_methods_invocations.py
differential testsFlashAttention’s teststhe fused kernel against a float64 reference, tolerance set from the reference’s own float32 errorflash-attention/tests/test_flash_attn.py
tolerancestorch.testing.assert_closeper-dtype defaults, like the course’s _lib.closePyTorch testing docs