Example: checking a gradient against finite differences

Example: checking a gradient against finite differences#

One task, the compact version of what Gradients, backward teaches step by step: hawk’s vjp differentiates a kernel’s own intermediate representation (IR, the typed graph hawk traces a kernel’s body into) symbolically, with no runtime tape (a recorded log of operations some other autodiff systems replay backward; hawk keeps none). This notebook checks a reverse-mode gradient the ordinary way anyone checks a hand-derived one — against a central finite difference of the compiled primal (the original, undifferentiated kernel) itself.

Time: ~3 min · Runs on: CPU · You need: Gradients, backward

Hide code cell source

import sys, pathlib
sys.path.insert(0, str(pathlib.Path.cwd().parent / "_shared"))
from nb_helpers import plot_style
plot_style()
import numpy as np

import hawk
from hawk import Mutable, Scalar, Vector
from hawk.math import dot


@hawk.kernel
def energy(v: Vector[3], out: Mutable[Scalar]):
    out = 0.5 * dot(v, v)


energy
<Kernel energy slots=2>

Derive its gradient and build both kernels:

import pathlib, tempfile

from hawk import Kernel
from hawk.diff import vjp

energy_vjp = Kernel("energy_vjp", vjp(energy, wrt=("v",)))
work_dir = pathlib.Path(tempfile.mkdtemp())
hawk.build([energy, energy_vjp], work_dir, targets=("host",));
def run_energy(v):
    out = np.zeros(v.shape[1])
    hawk.run(hawk.load(work_dir, "energy"), v=v, out=out)
    return out


def run_grad(v):
    n = v.shape[1]
    bar_out, bar_v = np.ones(n), np.zeros((3, n))
    hawk.run(hawk.load(work_dir, "energy_vjp"), v=v, bar_out=bar_out, bar_v=bar_v)
    return bar_v
rng = np.random.default_rng(42)
n = 2000
v = rng.normal(size=(3, n))
analytic = run_grad(v)

h = 1e-6
numeric = np.zeros_like(v)
for axis in range(3):
    plus, minus = v.copy(), v.copy()
    plus[axis] += h
    minus[axis] -= h
    numeric[axis] = (run_energy(plus) - run_energy(minus)) / (2 * h)
max_abs_error = np.max(np.abs(analytic - numeric))
msg = f"{n} samples x 3 components, max|analytic - central diff| = {max_abs_error:.3e}"
print(msg)
np.testing.assert_allclose(analytic, numeric, atol=1e-5)
2000 samples x 3 components, max|analytic - central diff| = 9.426e-10

What just happened#

The gradient hawk.diff.vjp derived symbolically agrees with a numerical central difference of the compiled primal to the precision a step h = 1e-6 can resolve — the standard sanity check for any new differentiation rule, run here against the real compiled artifact rather than a hand derivation.