Autodiff against torch

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.

Hide code cell source

import sys, pathlib, warnings

# This build's torch pulls in a deprecated pynvml that warns on the
# first torch.cuda touch; the warning's own traceback cites this
# machine's install path, which the docs' path-leak gate refuses --
# silence this one, known-irrelevant warning rather than the path leak.
warnings.filterwarnings(
    "ignore", message="The pynvml package is deprecated", category=FutureWarning
)

# These examples live in the repository's own examples/ directory, not in
# the installed raptor package — add it to sys.path so the module below
# imports straight from the checkout.
sys.path.insert(0, str(pathlib.Path("../../../examples").resolve()))
import autodiff_vs_torch as demo
print(demo.__name__, "from examples/" + pathlib.Path(demo.__file__).name)
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.