Train through a physics kernel with torch#

Write a physics step once, in hawk, and train a PyTorch model through it — with no second, hand-written backward pass.

Time: ~6 min · Runs on: CPU, GPU if visible (picked up automatically) · You need: a basic PyTorch training loop (nn.Module, loss.backward(), an optimizer) — this page teaches only what eagle and hawk add on top of that.

You will build: one kernel’s derivative checked against finite differences, then a one-parameter nn.Module whose forward pass runs that kernel.

Hide code cell source

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

import hawk
from hawk import Mutable, Param, Scalar, Terminated, Vector
from hawk.math import norm, vec

from eagle.frameworks import torch as eagle_torch

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.manual_seed(0)
print("running on", device.type)
running on cpu

The kernel: one time step of a ball in flight#

A ball under gravity and quadratic air drag, advanced by one semi-implicit Euler step. cd is the drag coefficient (per meter) we want to learn later; it is a per-sample Scalar plane, so a learned value stays on the device. dt and g are Params: plain numbers, held constant.

@hawk.kernel
def flight_step(
    r: Vector[3], v: Vector[3], cd: Scalar, dt: Param, g: Param,
    terminated: Terminated, r_next: Mutable[Vector[3]], v_next: Mutable[Vector[3]],
):
    a = vec(0.0, 0.0, -g) - cd * norm(v) * v
    v_new = v + dt * a
    v_next = v_new
    r_next = r + dt * v_new

Wrap it for torch#

eagle.frameworks.torch.function() turns a hawk kernel into a torch.autograd.Function: call step like any Python function, and if an input requires grad, it tracks the reverse-mode kernel hawk derives from flight_step for you — you never write that kernel yourself. wrt names the inputs worth deriving against; terminated is the early-termination mask, so it is not one.

step = eagle_torch.function(flight_step, wrt=("r", "v", "cd"))
step
KernelFunction('flight_step', inputs=['r', 'v', 'cd', 'dt', 'g', 'terminated'], outputs=['r_next', 'v_next'], wrt=['r', 'v', 'cd'])

A small batch to try it on#

8 balls, each with its own position, velocity and drag coefficient; 2 of the 8 are already marked terminated (landed).

n = 8
r = torch.randn(3, n, dtype=torch.float64, device=device, requires_grad=True)
v = torch.randn(3, n, dtype=torch.float64, device=device, requires_grad=True)
cd = torch.rand(n, dtype=torch.float64, device=device, requires_grad=True)
landed = torch.zeros(n, dtype=torch.bool, device=device)
landed[[2, 5]] = True
r.shape, v.shape, cd.shape
(torch.Size([3, 8]), torch.Size([3, 8]), torch.Size([8]))

A forward call, alone#

step is called exactly like flight_step reads: by position. The tensors cross the torch/hawk boundary zero-copy in both directions — nothing here is torch-specific yet, it is just a function call that happens to run a compiled kernel underneath.

r_next, v_next = step(r, v, cd, 0.05, 9.81, landed)
r_next.shape, v_next.shape, r_next.dtype
(torch.Size([3, 8]), torch.Size([3, 8]), torch.float64)

Check the derivative#

torch.autograd.gradcheck compares the derived reverse-mode kernel against finite differences, and check_forward_ad=True does the same for the derived forward-mode kernel, in the same call.

torch.autograd.gradcheck(
    lambda r, v, cd: step(r, v, cd, 0.05, 9.81, landed),
    (r, v, cd), check_forward_ad=True, check_batched_grad=False,
)
True

A terminated sample contributes zero gradient#

Call .backward() through step and read cd.grad: the two landed samples (indices 2 and 5) get exactly zero, because flight_step never touched them.

(r_next.sum() + v_next.sum()).backward()
print("d(loss)/d(cd) of the landed samples:", cd.grad[landed].tolist())
d(loss)/d(cd) of the landed samples: [0.0, 0.0]

Forward mode, too#

The forward-mode (JVP) rule gradcheck just exercised is also reachable on its own, through torch.autograd.forward_ad’s dual tensors: seed one direction on the inputs, run step once inside a dual_level(), and read the tangent off the output — no second graph, no .backward().

import torch.autograd.forward_ad as fwAD

EPS = 1e-6
r0_, v0_, cd0_ = r.detach().clone(), v.detach().clone(), cd.detach().clone()
d_cd = torch.ones_like(cd0_)  # direction: d/d(cd), one unit per sample

with fwAD.dual_level():
    cd_dual = fwAD.make_dual(cd0_, d_cd)
    r_next_dual, v_next_dual = step(r0_, v0_, cd_dual, 0.05, 9.81, landed)
    d_r_next = fwAD.unpack_dual(r_next_dual).tangent

r_plus, _ = step(r0_, v0_, cd0_ + EPS * d_cd, 0.05, 9.81, landed)
r_minus, _ = step(r0_, v0_, cd0_ - EPS * d_cd, 0.05, 9.81, landed)
fd = (r_plus - r_minus) / (2 * EPS)
torch.testing.assert_close(d_r_next, fd, rtol=1e-4, atol=1e-4)
print("forward-mode matches a central finite difference:", True)
forward-mode matches a central finite difference: True

What just happened#

  • step read and wrote ordinary torch tensors; the derived reverse-mode and forward-mode kernels came from the SAME flight_step you wrote once.

  • A terminated sample contributes exactly zero gradient — it never ran the kernel body, so it has nothing to differentiate.

  • gradcheck and the finite-difference comparison above are both checking the SAME derived kernel two different ways; they agree to rtol=1e-4.

A small torch model#

Ballistic holds one learnable parameter, the logarithm of the drag coefficient (log, so training never sees a negative drag), and unrolls step over a flight. Each call marks the balls below ground as terminated; torch.where then freezes their state, so a ball stays where it landed.

import math


class Ballistic(torch.nn.Module):
    def __init__(self, cd_guess, steps, dt=0.05, g=9.81):
        super().__init__()
        log_cd0 = torch.tensor(math.log(cd_guess), dtype=torch.float64)
        self.log_cd = torch.nn.Parameter(log_cd0)
        self.steps, self.dt, self.g = steps, dt, g

    def forward(self, r0, v0):
        cd = self.log_cd.exp().expand(r0.shape[-1]).contiguous()
        r, v = r0, v0
        for _ in range(self.steps):
            landed = r[2] < 0.0
            r_next, v_next = step(r, v, cd, self.dt, self.g, landed)
            r, v = torch.where(landed, r, r_next), torch.where(landed, v, v_next)
        return r

Run it once#

One ball, thrown up and sideways, run through Ballistic for a handful of steps — one forward pass, no optimizer yet, just confirming the module works end to end.

r0 = torch.tensor([[0.0], [0.0], [1.0]], dtype=torch.float64, device=device)
v0 = torch.tensor([[10.0], [0.0], [8.0]], dtype=torch.float64, device=device)
model = Ballistic(cd_guess=0.01, steps=40).to(device)
final_r = model(r0, v0)
print("position after 40 steps:", final_r.squeeze().tolist())
position after 40 steps: [15.525589613797282, 0.0, -0.40993144828405503]

Try this#

Change cd_guess from 0.01 to 0.2 (much more drag) and re-run: the final z lands higher (less speed lost to gravity before drag slows the ball) — the SAME module, nothing else changed.

Next#

  • A multi-kernel step, declaratively — composing several kernels into one step, the pattern this page’s single kernel is one piece of.

  • The full training loop — fitting Ballistic’s drag coefficient to noisy observed trajectories with L-BFGS, 256 balls at once — is the family’s capstone example: hawk <https://amasat01.github.io/hawk/content/vocabulary/04_capstone.html>__’s vocabulary/04_capstone notebook builds on exactly this kernel and this module.

Going deeper (optional): what crosses, and on which stream#

On a GPU, step reads the tensors torch already holds and writes tensors torch allocated: no copies and no host round trips. Every launch rides torch’s current stream, so it is ordered after the work torch queued before it and before the work torch queues after it, with no synchronization. On the CPU, the same tensors cross into hawk’s host runtime the same way.

Roadmap: JAX

The same three pieces — a forward kernel, its derived reverse-mode kernel and a DLPack crossing — are the template for a JAX jax.custom_vjp bridge. It is foreseen, not available today: torch is the framework this bridge supports.