Autodiff against torch#
hawk differentiates the traced kernel itself — reverse-mode (vjp) and
forward-mode (jvp) — rather than executing a separate tape. The primal is
one RK4 two-body step; hawk publishes the primal and both derivative
kernels in a single hawk.build call. The --host arm below checks both
derivatives against a float64 finite-difference Jacobian (and, since torch
is installed here, against torch.autograd.grad / torch.func.jvp on CPU
too) — no GPU needed.
python examples/autodiff_vs_torch.py --host
Needs hawk and
eagle installed from
PyPI, plus pip install "raptor-core[demo]" for the optional torch
cross-check.
autodiff_vs_torch from examples/autodiff_vs_torch.py
The primal step kernel hawk differentiates, both forward- and reverse-mode, from a single traced definition:
import inspect
print(inspect.getsource(demo.step_kernel))
def step_kernel():
"""Trace the primal RK4 step -- the same physics as
``early_termination.py``'s ``et_step``, with no termination test."""
import hawk
import hawk.math as hm
def rhs(s):
r2 = s[0] * s[0] + s[1] * s[1]
inv_r3 = hm.rsqrt(r2 * r2 * r2)
return hm.vec(s[2], s[3], -MU * s[0] * inv_r3, -MU * s[1] * inv_r3)
@hawk.kernel
def rk4_step(x: hawk.Vector[4], dt: hawk.Param,
x_next: hawk.Mutable[hawk.Vector[4]]):
k1 = rhs(x)
k2 = rhs(x + (0.5 * dt) * k1)
k3 = rhs(x + (0.5 * dt) * k2)
k4 = rhs(x + dt * k3)
x_next = x + (dt / 6.0) * (k1 + 2.0 * k2 + 2.0 * k3 + k4)
return rk4_step
By hand: seed a batch, then let hawk publish the primal kernel AND
its derived vjp/jvp kernels, all from one hawk.build call,
into a manifest raptor validates.
import pathlib, tempfile
import numpy as np
import hawk.artifact
from hawk import Kernel
from hawk.diff import ADJOINT_PREFIX, jvp, vjp
from raptor.schema import validate_manifest
import eagle.exec as eexec
from eagle import plan as eplan
state = demo.seed_states(1024, seed=0)
primal = demo.step_kernel()
wrt = ("x",) # differentiate w.r.t. the kernel's one Vector[4] input
derived = [Kernel("rk4_step_vjp", vjp(primal, wrt=wrt)),
Kernel("rk4_step_jvp", jvp(primal, wrt=wrt))]
bundle = hawk.build(
[primal] + derived, pathlib.Path(tempfile.mkdtemp()),
mode="float32", targets=("host",))
validate_manifest(bundle.manifest)
print("exec_targets:", bundle.manifest["exec_targets"])
exec_targets: ['host']
Then hand the vjp kernel to eagle and run it — this is the one call
that actually computes the gradient:
plugin = hawk.artifact.plugins(bundle)["rk4_step_vjp"]
plan = eplan.plan(plugin, structure=eexec.HostTeam)
x32 = np.ascontiguousarray(state, dtype=np.float32)
seed = np.ones_like(x32) # cotangent: d(loss)/d(x_next) = 1
bar_x = np.zeros_like(x32)
plan.bind(x=x32, dt=0.01,
**{f"{ADJOINT_PREFIX}_x_next": seed, f"{ADJOINT_PREFIX}_x": bar_x}).launch()
print("d(x_next)/d(x), sample 0:", bar_x[:, 0])
d(x_next)/d(x), sample 0: [1.0000026 0.98851097 1.0100002 1.0099425 ]
The pieces above, assembled into one call — building, running and checking both derivatives, for reference:
results = demo.run_demo(batch=1024, target="host")
target : host (manifest exec_targets ['host'])
batch x dt : 1024 x 0.01
vjp vs float64 FD (8 samples): max relative error = 5.471e-08 (tolerance 1e-05)
jvp vs float64 FD (8 samples): max relative error = 4.129e-08 (tolerance 1e-05)
vjp vs torch.autograd.grad (CPU): max relative error = 1.800e-07 (tolerance 1e-05)
jvp vs torch.func.jvp (CPU) : max relative error = 8.575e-08 (tolerance 1e-05)
Both derivative kernels agree with the float64 finite difference to well
under the declared tolerance, and — since this environment’s torch build is
CPU-only but still installed — also against torch.autograd.grad and
torch.func.jvp run on CPU. The --device arm additionally times hawk’s
derivative kernels against torch’s and feeds every plane through
eagle.to_cupy as in the zero-copy interop
example; see the repository
README for the measured
per-call speedups on a Quadro P2000, which this host-only run does not
reproduce.