Dual numbers and forward-mode autodiff
Overview
Section titled “Overview”| Module | M08.1 · build · Python · Pass 2 · 2 to 3 h |
| You build | python/tinyllm/autograd/dual.py: Dual (arithmetic, powers, @, indexing), dual_exp, dual_log, dual_tanh, dual_erf, derivative, jvp |
| Contract | course/contracts/py/tinyllm/autograd/dual.pyi |
| Tests | course/tests/M08.1/ (what they check: section 4) |
| Needs | M01.3 activation derivatives (the tests compare against them) · M04.2 numeric JVP · reading: M02.1 Taylor series (or --ref-deps) |
| Used by | M08.2 the forward-mode oracle for reverse mode · L0.2 the derivative oracle for every elementwise op |
| Milestone | MS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains) |
| Optional depth | Baydin, Pearlmutter, Radul, and Siskind, Automatic Differentiation in Machine Learning: a Survey (JMLR 2018), sections 2 and 3.1; Griewank and Walther, Evaluating Derivatives (SIAM, 2nd ed.), ch. 3 |
Key Takeaways
Section titled “Key Takeaways”- A dual number with turns Taylor’s theorem into an identity: , so evaluating a function on returns its value and its exact derivative together (
test_hand_example). - One rule per primitive is enough; the chain rule happens by itself when operations compose (
test_arithmetic_rules,test_chain_rule_composition). - With a vector tangent, one forward pass computes a Jacobian-vector product ; a full Jacobian costs one pass per input (
test_jvp_matches_numeric_jvp,test_jvp_is_linear_in_v). - Dual derivatives have no step size and match closed-form derivatives to rounding, which makes them the oracle for every elementwise op in your autograd (
test_matches_activation_derivatives).
How to work this chapter
Section titled “How to work this chapter”ol start M08.1 # stubs dual.py into your repo, contract alongsideol tests M08.1 # read the test catalog first: rung R0, you write no tests hereol check M08.1 # exit code is the verdictol check M08.1 --ref-deps # only if your M01.3 or M04.2 is not passing yetol diff M08.1 # after passing: your code against the reference1. Why now
Section titled “1. Why now”This pass replaces the tracer’s count table with a bigram trained by your own autograd (L0.1 to L0.5), and an autograd engine is a pile of derivative rules: one per op, each easy to get subtly wrong. So far you have two ways to check a derivative, and neither is good enough as an oracle. The closed forms of M01.3 cover only the functions you differentiated by hand. The central differences of M04.1 and M04.2 work for anything, but carry a step-size error near and cannot tell a correct rule from one that is off by . Dual numbers give a third way: write the forward computation once, run it on a number that carries its own derivative, and read off the exact derivative of the composition. L0.2 uses this to check each elementwise op’s backward, and M08.2 checks reverse mode against it.
2. Principles
Section titled “2. Principles”| Symbol | Meaning | Type / shape |
|---|---|---|
| a formal symbol with (and ) | ||
| a dual number: value , tangent | Dual(val, eps) | |
| a differentiable function and its derivative | ||
| an input point | float64[n] | |
| a tangent direction | float64[n] | |
| the Jacobian of , | float64[m, n] | |
| the Jacobian-vector product (JVP) | float64[m] | |
| a constant matrix | float64[m, n] | |
| unit roundoff of float64, about | scalar |
The algebra. Dual numbers add and multiply like polynomials in , then drop every :
The tangent of a product is the product rule. Division follows by multiplying with the conjugate , since :
the quotient rule. A plain number is the dual number : a constant has tangent 0.
Why it computes derivatives. Taylor’s theorem (M02.1) expands a smooth around :
Put . Every term from on contains , so
No step size, no truncation: the tangent is the derivative times the input tangent, and floating point adds only the rounding of each operation.
One rule per primitive. Read off each function’s dual rule from its derivative:
| Function | Dual rule |
|---|---|
| , constant | |
| , constant | |
| , both dual | for , |
The chain rule comes free. If , then feeding that into gives : the derivative of . Every program built from the primitives is differentiated by running it.
Vectors and Jacobian-vector products. Give every input its own tangent: with . The multivariable Taylor expansion has the same shape, , so the output tangent is the JVP. A linear map passes tangents through itself: . To build the whole Jacobian you run once per basis vector , getting column each time: forward mode costs one pass per input. It is ideal for few inputs and many outputs; a loss with millions of parameters and one output is the opposite case, and the reason M08.2 builds reverse mode.
Arrays of dual numbers. Store a vector of dual numbers as two arrays of one shape, val and eps, and apply every rule elementwise. Indexing and slicing take the same entries from both arrays; @ maps the tangent by the same matrix.
numpy on the left. np.float64(2.0) * d asks numpy first. numpy’s scalar and array types try to treat d as an element of an object array, and the tangent is lost (or you get an object array back). Setting the class attribute __array_ufunc__ = None tells numpy to return NotImplemented for any operation with a Dual, and Python then calls Dual.__rmul__. Every reflected method (__radd__, __rsub__, __rmul__, __rtruediv__, __rpow__, __rmatmul__) must keep the operand order: has tangent , and has tangent .
Exact versus approximate. A central difference has truncation error about and rounding error about , balanced near at an error near (M01.1). A dual derivative is as accurate as evaluating itself, about relative. That is why the tests compare dual derivatives with M01.3’s closed forms at a relative tolerance of , but with M04.2’s numeric JVP only at .
3. Worked example by hand
Section titled “3. Worked example by hand”at . Seed the input with tangent 1: .
| step | value | tangent | rule |
|---|---|---|---|
| 1 | 1 | seed | |
| exp | |||
| product | |||
| constant |
So and . By hand, , which is at 1. Same number, from one evaluation.
A quotient, at . Numerator , denominator ; by the quotient rule the tangent is . The function is 2 there, and its derivative .
A JVP. at along : , . Then and , so : the first column of .
The first example is the first test case in section 4, test_hand_example; the quotient is a row of test_arithmetic_rules.
4. The interface
Section titled “4. The interface”class Dual: __array_ufunc__ = None def __init__(self, val, eps=0.0) -> None # floats, or float64 arrays of one shape # + - * / ** @ with Dual, numbers, or arrays on either side; unary -; d[idx]def dual_exp(x) -> Dual; def dual_log(x) -> Dual; def dual_tanh(x) -> Dual; def dual_erf(x) -> Dualdef derivative(f: Callable[[Dual], Any], x: float) -> floatdef jvp(f: Callable[[Dual], Any], x: ArrayLike, v: ArrayLike) -> tuple[NDArray, NDArray]derivative seeds Dual(x, 1.0) and returns the output’s tangent (0.0 when f ignores its input). jvp seeds Dual(x, v) and returns (f(x), J v) as float64 arrays. numpy has no erf; use math.erf, elementwise for arrays (np.vectorize).
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
test_hand_example | unit | , for | you and the test agree on what a dual number carries |
test_arithmetic_rules | unit | 13 primitives and reflected forms against known derivatives | each rule is right on its own |
test_pow_rules | unit | , , , , | three different rules behind ** |
test_matches_activation_derivatives | differential | sigmoid, tanh, SiLU, both GELUs, softplus at 1001 points in against M01.3 | the exact oracle L0.2 relies on |
test_jvp_matches_numeric_jvp | differential | of a map against M04.2’s jvp_numeric | vector forward mode is right |
test_numpy_operand_on_the_left | boundary | np.float64(2) * d, 1 - d, array + d stay Dual | numpy scalars appear everywhere in real code |
test_constant_function_has_zero_derivative | boundary | a constant gives 0, jvp gives zeros, no nested Dual | functions that ignore their input |
test_chain_rule_composition | property | at 40 points | composition is the chain rule |
test_jvp_is_linear_in_v | property | linearity in ; basis vectors give the Jacobian’s columns | the cost model of forward mode |
test_vector_duals_index_and_matmul | unit | slices, W @ x, x @ W.T move val and tangent together | layers are matrices |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
| 1. keeping one term of the product rule | instead of in the worked example | test_hand_example (mutant s01) |
| 2. reflected operators that swap the order | gets slope ; gets the wrong sign | test_arithmetic_rules (mutants s03, s04, m01) |
3. no __array_ufunc__ = None | np.float64(2) * d silently loses the tangent | test_numpy_operand_on_the_left (mutant s08) |
4. W @ x that maps the value but not the tangent | JVPs of layers are wrong while scalar tests pass | test_vector_duals_index_and_matmul (mutant s12) |
| 5. a sign in the quotient rule, , | wrong slopes for division, tanh, and GELU | test_arithmetic_rules (mutants s02, s07, s14), test_matches_activation_derivatives (mutants s05, s06) |
| 6. instead of ; without | power and exponential slopes off | test_pow_rules (mutants s10, s11) |
| 7. seeding the input tangent with 0 | every derivative is 0 | test_constant_function_has_zero_derivative (mutant s09) |
| 8. indexing the value but not the tangent | slices carry another entry’s derivative | test_vector_duals_index_and_matmul (mutant s13) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Direction | Module | How it uses this |
|---|---|---|
| Back | M01.3 | its closed-form derivatives are what the dual derivatives must match |
| Back | M04.2 | its jvp_numeric approximates the same JVP by central differences |
| Back | M02.1 | Taylor’s theorem is why ; erf as a series |
| Forward | M08.2 | reverse mode is checked against Dual on the same expressions |
| Forward | L0.2 | each elementwise op’s backward is checked against its dual derivative |
If you skip this module, ol check M08.2 stops with M08.2 needs M08.1: build it, or rerun with --ref-deps.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
Dual | PyTorch forward-mode AD | dual tensors (fwAD.make_dual, unpack_dual) and torch.func.jvp over every op | torch/autograd/forward_ad.py |
jvp | JAX jax.jvp | forward mode by tracing, composable with vmap and with reverse mode (forward-over-reverse Hessian-vector products, M08.4) | jax/_src/interpreters/ad.py |
| arrays of duals | Julia ForwardDiff.jl | chunked tangents: up to 12 directions per pass, so a gradient costs passes | src/dual.jl |