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

if x < y: / where(c, a, b)

condition false at x == y

the else branch

minimum(a, b), maximum(a, b)

a == b

the first operand

abs(x)

x == 0

+1

clip(x, lo, hi)

x == lo or x == hi

x (closed interval)

clip with lo > hi

always

hi

clip with a NaN x

always

nothing (zero)

copysign(a, b)

a == 0

a by abs’s rule; b gets zero

floor, round, sign, comparisons

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 vjp gives you.

What just happened#

  • A Python if on a per-sample value becomes a select; vjp sends 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.