Hybrid orbit: analytic physics + a learned correction#
Every trajectory is a planar orbit around a point mass. The textbook two-body formula misses drag, so a small tanh network is trained in torch to predict exactly the missing piece of the acceleration. hawk traces ONE kernel that performs a whole RK4 step with that network evaluated inline, compiles it, and eagle runs the compiled kernel over the whole batch.
This notebook runs the --host arm (CPU threads); the same kernel also
runs as a replayed CUDA graph with --device on a machine with a GPU.
Equivalent command line:
python examples/two_body_hybrid.py --host --steps 50 --batch 256
Needs hawk and
eagle installed from
PyPI, plus pip install "raptor-core[demo]" for torch.
two_body_hybrid from examples/two_body_hybrid.py
The kernel is traced once, as ordinary Python — the MLP correction is evaluated inline, four times per RK4 step:
import inspect
print(inspect.getsource(demo.step_kernel))
def step_kernel(hidden: int = HIDDEN):
"""Trace the RK4 step. The MLP is evaluated inline, four times per step.
The weights arrive as ordinary traced arguments, so the kernel is a pure
function of ``(state, weights, dt)`` and carries no state of its own.
"""
import hawk
import hawk.math as hm
def hybrid_rhs(s, w1, b1, w2, b2):
"""Gravity plus the network's correction, as one traced expression."""
r2 = s[0] * s[0] + s[1] * s[1]
inv_r3 = hm.rsqrt(r2 * r2 * r2)
correction = w2 @ hm.tanh(w1 @ s + b1) + b2
return hm.vec(
s[2],
s[3],
-MU * s[0] * inv_r3 + correction[0],
-MU * s[1] * inv_r3 + correction[1],
)
@hawk.kernel
def two_body_hybrid_step(
x: hawk.Vector[4],
w1: hawk.Matrix[hidden, 4],
b1: hawk.Vector[hidden],
w2: hawk.Matrix[2, hidden],
b2: hawk.Vector[2],
dt: hawk.Param,
x_next: hawk.Mutable[hawk.Vector[4]],
):
k1 = hybrid_rhs(x, w1, b1, w2, b2)
k2 = hybrid_rhs(x + (0.5 * dt) * k1, w1, b1, w2, b2)
k3 = hybrid_rhs(x + (0.5 * dt) * k2, w1, b1, w2, b2)
k4 = hybrid_rhs(x + dt * k3, w1, b1, w2, b2)
x_next = x + (dt / 6.0) * (k1 + 2.0 * k2 + 2.0 * k3 + k4)
return two_body_hybrid_step
By hand: train the correction, then let hawk publish the compiled kernel into a manifest raptor validates.
import pathlib, tempfile
import numpy as np
import hawk.artifact
from raptor.schema import validate_manifest
import eagle.exec as eexec
from eagle import plan as eplan
state = demo.seed_states(batch=256, seed=0)
weights = demo.train_correction(seed=0)
bundle = hawk.build(
[demo.step_kernel()], pathlib.Path(tempfile.mkdtemp()),
mode="float32", targets=("host",))
validate_manifest(bundle.manifest)
print("exec_targets:", bundle.manifest["exec_targets"])
exec_targets: ['host']
Then hand it to eagle and run it — this is the one call that actually launches the kernel:
plugin = next(iter(hawk.artifact.plugins(bundle).values()))
plan = eplan.plan(plugin, structure=eexec.HostTeam)
planes = {k: np.repeat(weights[k].reshape(-1, 1).astype(np.float32), 256, 1)
for k in ("w1", "b1", "w2", "b2")}
x_next = np.zeros_like(state, dtype=np.float32)
bound = plan.bind(x=state.astype(np.float32), x_next=x_next, dt=0.01, **planes)
bound.launch()
print("x_next[:, 0]:", x_next[:, 0])
x_next[:, 0]: [-0.64382726 -0.7641566 0.77130795 -0.65005666]
The pieces above, assembled into one call — training, building, running and checking it, for reference:
results = demo.run_demo(steps=50, batch=256, target="host")
target : host (manifest exec_targets ['host'])
batch x steps : 256 x 50 at dt = 0.01
training loss : 1.230e-06
(a) max |compiled - torch eager| : 2.235e-07 (tolerance 1e-04)
(b) final-state error vs truth : corrected 6.695e-04 vs two-body only 2.405e-02
(c) wall per step (demo measurement, not a benchmark): compiled 1.130 ms, torch eager 0.397 ms
match is acceptance (a): the compiled rollout tracks a torch-eager
rollout of the same weights to a tight tolerance. hybrid_error vs.
two_body_error is acceptance (b): the learned correction should land
closer to the true (gravity + drag) trajectory than gravity alone. The wall
times are this run’s own measurement, not a committed benchmark — see the
repository README for the
measured numbers on a Quadro P2000 at the script’s default scale.