Early termination: a batch that finishes at different times

Early termination: a batch that finishes at different times#

Every trajectory is a planar two-body orbit that eventually crosses a collision or an escape radius and then stops changing. The step kernel maintains a device-side live count of samples still going; on the GPU arm, a conditional graph node guards the whole step with that count so a captured CUDA graph’s later replays do no work once every sample is done. The --host arm below runs the identical kernel through eagle’s host execution structure, with the guard evaluated eagerly — no GPU needed.

python examples/early_termination.py --host

Needs hawk and eagle installed from PyPI.

Hide code cell source

import sys, pathlib

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

The step kernel below is what runs once per sample, per step — the live count and the guard it drives are declared vocabulary, not hidden machinery:

import inspect
print(inspect.getsource(demo.step_kernel))
def step_kernel():
    """Trace the gated RK4 step.

    ``alive``/``alive_next`` are explicit 0/1 planes (``T.Terminated`` is
    read-only, so a terminated sample's own flag has to be carried this way).
    ``age``/``age_next`` ping-pong an increment-while-alive counter
    (``age_next = age + alive``): once a sample's ``alive`` reaches 0, its age
    stops advancing, so the age plane read back at the end IS that sample's
    termination step -- there is no per-replay step-number parameter to read
    it from, because a captured graph replay cannot see a scalar that varies
    per replay (a ``Param`` freezes at capture time).

    ``n_active`` is a one-cell accumulator, seeded to the batch size by the
    caller and never re-zeroed: every launch subtracts exactly the number of
    samples that terminate on THIS step (``still - alive``, zero for a sample
    that stays alive or was already dead), so it decays monotonically to zero
    as the batch finishes. An ``Accum`` plane must be a float type (there is
    no integer atomic accumulate), so the caller views this float32 cell as
    uint32 to build the graph's skip guard -- exact for batches under 2**24
    samples, comfortably above anything this demo runs.

    A dead sample's own step is a mathematical no-op (``x_next = x``,
    ``alive_next = 0``, ``age_next = age``, the accumulator's contribution
    ``still - alive = 0 - 0 = 0``) -- this is WHY the unguarded graph arm
    (which never skips, so it re-runs this no-op every remaining step) still
    lands bit-for-bit on the same final state as the guarded one.
    """
    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 et_step(
        x: hawk.Vector[4],
        alive: hawk.Scalar,
        age: hawk.Scalar,
        dt: hawk.Param,
        r_esc2: hawk.Param,
        r_col2: hawk.Param,
        x_next: hawk.Mutable[hawk.Vector[4]],
        alive_next: hawk.Mutable[hawk.Scalar],
        age_next: hawk.Mutable[hawk.Scalar],
        n_active: hawk.Accum[hawk.Scalar],
    ):
        k1 = rhs(x)
        k2 = rhs(x + (0.5 * dt) * k1)
        k3 = rhs(x + (0.5 * dt) * k2)
        k4 = rhs(x + dt * k3)
        xn = x + (dt / 6.0) * (k1 + 2.0 * k2 + 2.0 * k3 + k4)

        r2 = xn[0] * xn[0] + xn[1] * xn[1]
        zero = alive * 0
        one = zero + 1
        live = alive != 0
        still = hm.select(
            live, hm.select(r2 < r_esc2, hm.select(r2 > r_col2, one, zero), zero), zero
        )

        x_next = hm.select(live, xn, x)
        alive_next = still
        age_next = age + alive
        n_active.add(still - alive, at=0)

    return et_step

By hand: seed a batch, 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(1024, 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 one step — this is the one call that actually launches the kernel, with the live-count guard evaluated eagerly:

batch = state.shape[1]
x = np.ascontiguousarray(state, dtype=np.float32)
alive = np.ones(batch, dtype=np.float32)
age = np.zeros(batch, dtype=np.float32)
x_next = np.zeros_like(x)
alive_next = np.zeros_like(alive)
age_next = np.zeros_like(age)
n_active = np.full(1, float(batch), dtype=np.float32)

plugin = next(iter(hawk.artifact.plugins(bundle).values()))
plan = eplan.plan(plugin, structure=eexec.HostTeam)
bound = plan.bind(
    x=x, alive=alive, age=age, dt=0.01,
    r_esc2=1.5 ** 2, r_col2=0.2 ** 2,  # the demo's escape/collision radii
    x_next=x_next, alive_next=alive_next, age_next=age_next, n_active=n_active)
bound.launch()
print("samples still alive after one step:", int(alive_next.sum()))
samples still alive after one step: 1024

The pieces above, assembled into one call — building, running and checking every sample terminated, for reference:

results = demo.run_demo(target="host")
target            : host (manifest exec_targets ['host'])
batch x steps     : 1024 x 400 at dt = 0.01
escape / collide  : r > 1.5 / r < 0.2
samples still active after the horizon : 0 (expected 0)
lifetime histogram (10 bins over the REALISED lifetime range [0, 284], not the full horizon): [0, 0, 0, 449, 386, 75, 53, 30, 20, 11]
replay after all-terminated: state bit-identical = True, count stays 0 = True
positive control (count restored nonzero): step runs = True
host syncs        : 200
wall per step (demo measurement, not a benchmark): 31.8 us

active_after == 0 is the correctness condition this run checks: every sample terminated within the horizon. skip_frozen/skip_count_zero prove the kernel stops touching the state once the live count hits zero. On a GPU, the same kernel’s guarded CUDA graph (one that skips a replay once the count is zero) is timed against an unguarded control (the identical graph, never skipped) and an eager host-guarded loop — see the repository README for the measured ratios (over 13x at a million samples over 4000 steps on a Quadro P2000), which this host-only run does not reproduce.