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.
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.