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.
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#
stepread and wrote ordinary torch tensors; the derived reverse-mode and forward-mode kernels came from the SAMEflight_stepyou wrote once.A
terminatedsample contributes exactly zero gradient — it never ran the kernel body, so it has nothing to differentiate.gradcheckand the finite-difference comparison above are both checking the SAME derived kernel two different ways; they agree tortol=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>__’svocabulary/04_capstonenotebook 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.