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
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.