Gradients through branches and stops#
Differentiate a kernel that switches between two behaviours, and one that stops itself — and know exactly what the gradient means at the switch.
Time: ~8 min · Runs on: CPU · You need: Gradients, backward
Real models branch: a contact force that only acts on touching, a
saturation, a stop condition. hawk’s vjp handles these directly — each
branch is differentiated on its own, and the derived kernel takes the
same path the primal took.
import hawk
from hawk import Mutable, Scalar, Terminated
@hawk.kernel
def contact(x: Scalar, v: Scalar, k: Scalar, terminated: Terminated,
v_new: Mutable[Scalar]):
if x < 0.0: # touching: a stiff spring pushes back
a = -9.81 - k * x
else: # airborne: free flight
a = -9.81
v_new = v + 0.01 * a # one explicit step of 10 ms
terminated = x < -0.5 # pressed in too deep: stop this sample
contact
<Kernel contact slots=7>
An if becomes a select#
if x < 0.0: on a per-sample value is not a Python branch that runs
once; hawk reads it as a select — “take this value where the
condition holds, that value where it does not” — and evaluates the
condition per sample. That is what lets one kernel serve a whole batch
where some samples touch and others fly. It is also what makes it
differentiable: the derivative of a select is the derivative of the
branch it took, and nothing from the branch it did not.
Gradient, checked away from the switch#
import pathlib, tempfile
import numpy as np
from hawk import Kernel
from hawk.diff import vjp
contact_vjp = Kernel("contact_vjp", vjp(contact, wrt=("x", "k")))
work_dir = pathlib.Path(tempfile.mkdtemp())
_ = hawk.build([contact, contact_vjp], work_dir, targets=("host",))
def run_primal(x, v, k, stopped=None):
n = len(x)
flags = np.zeros(n, dtype=bool) if stopped is None else stopped.copy()
out = np.zeros(n)
hawk.run(hawk.load(work_dir, "contact"), x=x, v=v, k=k, terminated=flags, v_new=out)
return out, flags
def run_grad(x, v, k, stopped):
n = len(x)
bar_x, bar_k = np.zeros(n), np.zeros(n)
hawk.run(hawk.load(work_dir, "contact_vjp"), x=x, v=v, k=k, terminated=stopped,
bar_v_new=np.ones(n), bar_x=bar_x, bar_k=bar_k)
return bar_x + 0.0, bar_k + 0.0 # + 0.0 turns a signed zero into plain 0
x = np.array([0.3, 0.1, -0.1, -0.3]) # two airborne, two touching
v = np.full(4, -1.0)
k = np.full(4, 1000.0)
none_stopped = np.zeros(4, dtype=bool)
bar_x, bar_k = run_grad(x, v, k, none_stopped)
h = 1e-6
fd_x = (run_primal(x + h, v, k)[0] - run_primal(x - h, v, k)[0]) / (2 * h)
fd_k = (run_primal(x, v, k + h)[0] - run_primal(x, v, k - h)[0]) / (2 * h)
print("d v_new / d x :", bar_x)
print("finite diff :", np.round(fd_x, 6))
print("d v_new / d k :", bar_k)
print("finite diff :", np.round(fd_k, 6))
d v_new / d x : [ 0. 0. -10. -10.]
finite diff : [ 0. 0. -10. -10.]
d v_new / d k : [0. 0. 0.001 0.003]
finite diff : [0. 0. 0.001 0.003]
Airborne samples (x >= 0) have zero gradient: the spring is not in
play. Touching samples pick up -0.01 * k for x (the spring’s
stiffness times the step) and -0.01 * x for k. Each matches a central
finite difference, because none of these points sits on the switch.
What happens exactly at the switch#
Away from a switch the slope is unambiguous. At the switch there is a choice to make, and hawk makes the same one every time, so results are reproducible. The rules, each pinned by a test:
Operation |
At the tie |
Gradient goes to |
|---|---|---|
|
condition false at |
the |
|
|
the first operand |
|
|
|
|
|
|
|
always |
|
|
always |
nothing (zero) |
|
|
|
|
everywhere |
exactly zero |
Check the first row on our kernel: a sample sitting exactly at x = 0
takes the else branch (free flight), so its gradient is the airborne
one.
tie_x = np.array([0.0])
print("gradient at x == 0:", run_grad(tie_x, v[:1], k[:1], np.zeros(1, dtype=bool))[0])
gradient at x == 0: [0.]
Stops: a stopped sample adds nothing#
contact also sets terminated = x < -0.5. A sample flagged stopped is
skipped by the forward kernel, and the derived kernel is an ordinary
guarded kernel too: it skips the same samples, so a stopped sample
contributes exactly zero to the reverse pass — no mask array to
manage and no extra work, since the work is simply not done.
Let us run five samples; the last one is pressed in past -0.5:
x5 = np.array([0.3, 0.0, -0.1, -0.2, -0.8])
v5, k5 = np.full(5, -1.0), np.full(5, 1000.0)
_, stopped = run_primal(x5, v5, k5)
print("stopped flags :", stopped)
bar_x5, bar_k5 = run_grad(x5, v5, k5, stopped)
print("d v_new / d x :", bar_x5)
print("d v_new / d k :", bar_k5)
stopped flags : [False False False False True]
d v_new / d x : [ 0. 0. -10. -10. 0.]
d v_new / d k : [0. 0. 0.001 0.002 0. ]
The stopped sample (x = -0.8) shows 0 in both columns while its
touching neighbours do not. In a long rollout this is the behaviour you
want: a trajectory that has ended stops influencing the gradient the moment
it ends.
What is not differentiated#
A parameter can change where or when a system switches or stops: a stiffer or softer surface, a different threshold, a different landing time. That effect — an event-time term — is not part of the gradient hawk returns. The gradient is the piecewise one: exact at every point away from a switch, treating the switch location as fixed.
In practice this means:
Inside a branch, trust the gradient as you would any other (the finite-difference check above is the way to confirm it).
Right at a switch, you get the convention from the table; a finite difference straddling the switch measures something different, because it also sees the jump.
If your loss depends on when something happens, treat that dependence as separate from what
vjpgives you.
What just happened#
A Python
ifon a per-sample value becomes a select;vjpsends the gradient to the branch taken and zero to the other.At a tie, fixed rules decide (table above), the same on every run.
Stopped samples are skipped by the reverse kernel just as by the forward one, so they add exactly zero at no extra cost.
The piecewise gradient is exact away from a switch; the movement of the switch itself is not included.
Try this#
Change x = np.array([0.3, 0.1, -0.1, -0.3]) so one sample sits at
-1e-9 and compare the finite difference with the vjp: the finite
difference with h = 1e-6 crosses the switch and disagrees, as the last
section describes.
Next#
Writing a vocabulary — declare a reusable kernel shape once, so a whole family of models shares one vocabulary.