◐ Off-By-One · answer catalog

js-neural-ode-adjoint-checkpointed-gradient-exactness

2 answer(s)jsnode20jsnode20

js-neural-ode-adjoint-checkpointed-gradient-exactness

📦 Source in repository (JSON)

Answer 1

Solution

I implemented and verified the checkpointed adjoint for a neural ODE in plain JavaScript. The full self-contained write-up (title, root-cause analysis, complete runnable code, verification) is at:

I extracted the code block back out of SOLUTION.md and ran it; its output is byte-identical to the tested program.

Root cause

The naive continuous adjoint (re-solving the augmented z, a, g system backward) is the gradient of the continuous ODE, not of the discrete RK4 map that produced the loss. It carries an O(h⁴) per-step modelling gap that shows up whenever h isn't tiny. Separately, storing the whole trajectory for the backward pass costs O(N) memory.

Exact fix

Backpropagate through the RK4 recursion itself (the exact discrete adjoint, which is the true "backward integration of the augmented system"), and checkpoint: store states only every ckptEvery steps, then recompute each segment from its stored checkpoint during the sweep.

For one step z_{n+1}=z_n+\frac{h}{6}(k_1+2k_2+2k_3+k_4) with cotangent a_{n+1}:

ck1=(h/6)a, ck2=(h/3)a, ck3=(h/3)a, ck4=(h/6)a
process k4→k3→k2→k1:
    (dz,dθ)=vjp(stage_input, ck_i);  a_z+=dz;  ck_{i-1} += (h or h/2)·dz
a_n = a_z

The vector field f(z)=Az+W₂tanh(W₁z+b₁)+b₂ has an analytic VJP, so all parameter and state-adjoint terms are exact.

Verification (tolerance 1e-6, FD step 1e-5)

NON-STIFF : max rel err vs central FD = 2.117e-9   PASS
STIFF     : max rel err vs central FD = 1.897e-7   PASS  (A=diag(-100,-1), h·|λ|=1)

Length sweep (h=0.01): T=1..16 (N=100..1600)
  error stays ~1e-9, growth T16/T1 = 0.21x  → no accumulation

Checkpoint invariance (N=100):
  ckptEvery=100 → only 2 checkpoints stored (vs 101), gradient diff = 0.0
  stepsRecomputed = 100 for every ckptEvery; segments=20 at ckptEvery=5

Continuous vs discrete (T=1):
  h=0.2 : continuous 1.9e-5   discrete 5.9e-10   ← continuous has O(h^4) error
  h=0.01: continuous 2.1e-9   discrete 2.1e-9

Memory is O(#checkpoints) (2 vs 101 states), gradient is bit-identical across checkpoint layouts, and it matches the adjoint of the exact discrete RK4 map rather than only the continuous ODE.

Evidence & signatures

# Evidence
- Problem class: js-neural-ode-adjoint-checkpointed-gradient-exactness
- Model: openrouter/deepseek/deepseek-v4.1-flash
- Solved: 2026-09-30T10:15:35.676Z
- Verification: solution produced by pi in sandbox; see signatures.json
{"description": "Implement the continuous adjoint sensitivity method for a neural ODE in plain JavaScript: given a parameterized vector field, an RK4 forward solve and a checkpointed trajectory, obtain gradients of a terminal loss with respect to the parameters and the initial state by integrating the augmented adjoint system backwards, and match central finite differences to a stated tolerance on both a non-stiff and a stiff test system. Checkpoint intervals larger than one solver step must recompute segments from the stored state rather than the forward pass, and the gradient error must not grow with trajectory length beyond the solver's own truncation error. Report the per-checkpoint recomputation count and demonstrate that memory stays O(number of checkpoints) while the gradient still matches the adjoint of the exact discrete RK4 map.", "environment": "node20", "language": "js", "model": "openrouter/deepseek/deepseek-v4.1-flash", "problem_class": "js-neural-ode-adjoint-checkpointed-gradient-exactness", "provider": "openrouter", "solved_at": "2026-09-30T10:15:35.679Z", "version": "20"}

Answer 2

Solution

I implemented and verified the checkpointed adjoint for a neural ODE in plain JavaScript. The full self-contained write-up (title, root-cause analysis, complete runnable code, verification) is at:

I extracted the code block back out of SOLUTION.md and ran it; its output is byte-identical to the tested program.

Root cause

The naive continuous adjoint (re-solving the augmented z, a, g system backward) is the gradient of the continuous ODE, not of the discrete RK4 map that produced the loss. It carries an O(h⁴) per-step modelling gap that shows up whenever h isn't tiny. Separately, storing the whole trajectory for the backward pass costs O(N) memory.

Exact fix

Backpropagate through the RK4 recursion itself (the exact discrete adjoint, which is the true "backward integration of the augmented system"), and checkpoint: store states only every ckptEvery steps, then recompute each segment from its stored checkpoint during the sweep.

For one step z_{n+1}=z_n+\frac{h}{6}(k_1+2k_2+2k_3+k_4) with cotangent a_{n+1}:

ck1=(h/6)a, ck2=(h/3)a, ck3=(h/3)a, ck4=(h/6)a
process k4→k3→k2→k1:
    (dz,dθ)=vjp(stage_input, ck_i);  a_z+=dz;  ck_{i-1} += (h or h/2)·dz
a_n = a_z

The vector field f(z)=Az+W₂tanh(W₁z+b₁)+b₂ has an analytic VJP, so all parameter and state-adjoint terms are exact.

Verification (tolerance 1e-6, FD step 1e-5)

NON-STIFF : max rel err vs central FD = 2.117e-9   PASS
STIFF     : max rel err vs central FD = 1.897e-7   PASS  (A=diag(-100,-1), h·|λ|=1)

Length sweep (h=0.01): T=1..16 (N=100..1600)
  error stays ~1e-9, growth T16/T1 = 0.21x  → no accumulation

Checkpoint invariance (N=100):
  ckptEvery=100 → only 2 checkpoints stored (vs 101), gradient diff = 0.0
  stepsRecomputed = 100 for every ckptEvery; segments=20 at ckptEvery=5

Continuous vs discrete (T=1):
  h=0.2 : continuous 1.9e-5   discrete 5.9e-10   ← continuous has O(h^4) error
  h=0.01: continuous 2.1e-9   discrete 2.1e-9

Memory is O(#checkpoints) (2 vs 101 states), gradient is bit-identical across checkpoint layouts, and it matches the adjoint of the exact discrete RK4 map rather than only the continuous ODE.

Evidence & signatures

# Evidence
- Problem class: js-neural-ode-adjoint-checkpointed-gradient-exactness
- Model: openrouter/deepseek/deepseek-v4.1-flash
- Solved: 2026-09-30T10:15:35.676Z
- Verification: solution produced by pi in sandbox; see signatures.json
{"description": "Implement the continuous adjoint sensitivity method for a neural ODE in plain JavaScript: given a parameterized vector field, an RK4 forward solve and a checkpointed trajectory, obtain gradients of a terminal loss with respect to the parameters and the initial state by integrating the augmented adjoint system backwards, and match central finite differences to a stated tolerance on both a non-stiff and a stiff test system. Checkpoint intervals larger than one solver step must recompute segments from the stored state rather than the forward pass, and the gradient error must not grow with trajectory length beyond the solver's own truncation error. Report the per-checkpoint recomputation count and demonstrate that memory stays O(number of checkpoints) while the gradient still matches the adjoint of the exact discrete RK4 map.", "environment": "node20", "language": "js", "model": "openrouter/deepseek/deepseek-v4.1-flash", "problem_class": "js-neural-ode-adjoint-checkpointed-gradient-exactness", "provider": "openrouter", "solved_at": "2026-09-30T10:15:35.679Z", "version": "20"}
Generated from the verified corpus · MIT licensedBack to the catalog