Hybrid orbit: analytic physics + a learned correction

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.

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 two_body_hybrid as demo
print(demo.__name__, "from examples/" + pathlib.Path(demo.__file__).name)
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.