Skip to content

train: the exact walk, one optimizer step per window on the gradient of its whole loss (opt-in) - #47

Merged
joelteply merged 14 commits into
feat/props-weight-residencyfrom
train/walk-exact-gradient
Oct 7, 2026
Merged

joelteply merged 14 commits into
feat/props-weight-residencyfrom
train/walk-exact-gradient

Conversation

@joelteply

Copy link
Copy Markdown

Stacked on #40 (the walk). Opt-in: "exact": true on /train (llama_opt_params::walk_exact) plus walk_horizon (0 = the whole window). With it off, nothing changes.

What it does. The plain walk steps per chunk and stops the gradient at every chunk boundary, so a reply learns only through its own chunk's K/V. The exact walk decodes the window once under the adapter as it is, then trains it in reverse:

  • Each chunk's graph attends to the cache, with its last walk_horizon prefix positions behind a zero GRAD leaf.
  • It adds the surrogate ⟨K, G⟩ + ⟨V, G⟩ on its own K/V, where G is what every later chunk put there.
  • Its backward carries that gradient through the chunk into the adapter and into its own prefix's G.
  • Context chunks train through the surrogate alone.
  • Each chunk's loss is weighted by its labelled positions over the window's (Fable), so the chunk sum is the window's loss.

There is one optimizer step per window. That is also what keeps the reverse pass's recomputed K/V equal to the cached K/V. Memory stays chunk × window, plus host accumulators of layers × window × 2 × kv_dim floats.

Pieces (one commit each, each with its test):

Acceptance (test-walk-exact -m Qwen2.5-Coder-1.5B-Instruct-Q4_K_M.gguf, comparing the SGD step read off the adapter, i.e. −lr·g, gradients rather than losses per Fable):

case exact walk plain walk
every position labelled, 4 chunks vs one graph cosine 0.9983, norm 0.995 cosine 0.81, norm 2.15
context + reply, 4 chunks vs one chunk per run cosine 0.9982, norm 0.989 cosine 0.63

The remaining ~0.2% is the quantized matmuls' batch-size variance (a 128-row chunk and a 512-row graph take different kernels). Mutation-checked: without #46 the test fails at cosine 0.43.

Not yet (named, not hidden):

  1. A recurrent or hybrid model is refused by name. Its state's gradient across chunks is next (Fable's ask 1: the entry state as a GRAD leaf, the exit state with a surrogate). Kimi's 27B is a hybrid, so that's the critical path for her.
  2. The accumulators' host memory (layers × window × 2 × kv_dim × 4 B) needs to enter the budget gate before this runs on the 27B (Fable's ask 4).
  3. Fable's GRAD_DUMP/GRAD_CMP per-tensor comparison and the 0.8B hybrid case land with (1).

🤖 Generated with Claude Code

https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc

joelteply and others added 6 commits October 6, 2026 17:36
…tes and no optimizer updates

The training walk's exact gradient needs dL/d(cached prefix K/V): the sensitivity of a later
chunk's loss to the K/V an earlier chunk produced, read back and carried to that chunk. A PARAM
would get it, but also an optimizer step and moment buffers allocated once for a fixed set, so a
per-chunk copy cannot be one. GGML_TENSOR_FLAG_GRAD (ggml_set_grad) marks a leaf that is a graph
node for the backward (not a constant) and needs its gradient, and is never a parameter.

test-grad-leaf: a GRAD leaf beside a PARAM gets its exact gradient, the PARAM keeps its own, and
an unflagged constant gets none.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…o gradients

ggml_opt_prepare_alloc (the path llama_context's training takes: a graph per ubatch) keeps the
gradient accumulators in ctx_static across graphs, and each backward adds into them in place.
The only reset sat before the build, against gb_grad, which is null in this mode, so nothing
ever zeroed them: step k applied the SUM of every gradient since the run began. Measured
(test-opt-dynamic-accum, SGD on loss = w*x): the second step moved w by 2*lr*x, not lr*x.
Upstream has the same code.

The fix zeroes the parameter accumulators after the per-step build when a period begins.
Static graphs are unchanged (test-opt: 4/4 backend x optimizer pass).

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…node index

ggml_opt_build created the gradient accumulators and AdamW momenta once, indexed by the FIRST
graph's node order, and bound them to every later graph by that index. A graph built per step
may differ in topology (the walk's chunks with and without a cached prefix; its reverse pass with
GRAD leaves and a surrogate loss term), and then index i is another node: a parameter's
accumulator lands on whatever took its slot. They are now keyed by the parameter tensor (and
one for the loss); each build maps its own nodes to them, and a graph training a parameter the
context was not built with is refused by assert.

test-opt 4/4, test-opt-dynamic-accum, test-grad-leaf pass.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…xtra term, and GRAD-leaf gradients read back

The training walk's exact gradient runs one graph per chunk in reverse and ONE optimizer step
per window, with chunk counts that vary by window. ggml_opt_set_next_step(period_end,
loss_scale, extra_loss) sets, for the next alloc + eval: whether this graph takes the step or
only accumulates, its loss weight (replacing 1/opt_period), and an optional scalar node added to
the loss the backward differentiates: the surrogate that carries later chunks' gradient on a
chunk's K/V into it. The result keeps reporting the unweighted loss. A period starts from zero
gradients after the step that ended the previous one. ggml_opt_leaf_grad returns the gradient
computed for a GRAD leaf (marked output, so the allocator keeps it) until the next alloc.

test-opt-dynamic-accum: three graphs at weights 0.5 / 0.25 (+ 3 w x) / 1 make one SGD step on
exactly their weighted sum, and each graph's dL/dx reads back; a second period starts from
zero. Mutation-checked: dropping the extra term fails it.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…read the wrong elements)

The slice is ggml_view_4d(grad, ..., grad->nb[1..3], offset), and a view's first stride is the
element size. The training attention transposes V for its KQV product, so V's gradient reaches
the concat (the walk's [cached prefix | this chunk's V]) through a transpose, with a first
stride that is not the element size: the slice read the right NUMBER of elements from the wrong
places. K is only permuted (0,2,1,3), keeping its first stride, so it came through intact.

Measured on the 1.5B (test-walk-exact, SGD step = -lr*g read off the adapter): the walk's
cross-chunk V term had the right norm (0.93) and was orthogonal to the true one (cosine -0.10);
with the fix the exact walk in 4 chunks matches one graph at cosine 0.998, and the PLAIN walk's
step moved from cosine 0.54 to 0.84 against the one-graph step: every walk-trained adapter had
its attn_v gradient scrambled on every chunk with a cached prefix.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…of its whole loss (opt-in)

The plain walk steps per chunk and stops the gradient at every chunk boundary: a reply learns
only through its own chunk's K/V. walk_exact (llama_opt_params; "exact": true on /train) decodes
the window once under the adapter as it is, then trains it in REVERSE. Each chunk's graph
attends to the cache, its last walk_horizon prefix positions (0 = all) behind a zero GRAD leaf,
and adds the surrogate <K, G> + <V, G> on its own K/V, where G is what every later chunk put
there. Its backward carries that gradient through itself into the adapter and into its prefix's
G. Context chunks train through the surrogate alone (one zero-weighted output row). Each
chunk's loss is weighted by its labelled positions over the window's, so the chunk sum is the
window's loss (Fable). ONE optimizer step per window is also what keeps the reverse pass's
recomputed K/V equal to the cached K/V. Memory stays chunk x window, plus host accumulators of
layers x window x 2 x kv_dim floats; the cost is one decode plus one forward and backward per
chunk. v1 refuses a recurrent model by name (its state's gradient is next).

test-walk-exact (1.5B Q4_K_M, the adapter's SGD step = -lr * gradient):
  every position labelled: exact walk in 4 chunks vs one graph: cosine 0.9983, norm 0.9954
                           (plain walk: cosine 0.81, norm 2.15)
  context + reply:         exact walk in 4 chunks vs one chunk per run: cosine 0.9982
                           (plain walk: cosine 0.63)
Mutation-checked: without the CONCAT backward fix it fails at cosine 0.43.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…nks (hybrid models)

Fable's ask: a hybrid's linear layers carry the window in their state, so stopping the gradient
there is where most of the bias would be. The state is treated like K/V. build_rs puts a zero
GRAD leaf on the state ENTERING a chunk (the backward yields dL/d(entry state)), and the two
delta-net helpers give the state a chunk LEAVES (the conv state's tail, the delta-net's new
state) the surrogate <exit, G>, where G is the next chunk's dL/d(entry state)
(build_walk_state_exit). Each exit state feeds exactly one chunk, so G is one vector per state
tensor, overwritten as the reverse pass moves back. The forward pass decodes chunk by chunk and
snapshots sequence 0's state rows on the host at every chunk start; the reverse pass restores
each chunk's entry state before training it (the first chunk starts from zero). The delta-net
backward already returns its initial state's gradient and folds in its final state's through
the snapshot slots; the conv state goes through CONCAT (#46).

test-walk-exact on Qwen3.5-0.8B (attention + gated delta-net), the SGD step read off the adapter:
  every position labelled: exact walk in 4 chunks vs one graph: cosine 0.9980, norm 0.9997
                           (plain walk: cosine 0.83, norm 2.13)
  context + reply:         exact walk in 4 chunks vs one chunk per run: cosine 0.9985
Mutation-checked: without the state surrogate it fails (cosine 0.987, norm 0.969). The 1.5B
results are unchanged. The test's contexts carry a scratch sequence over a unified cache, which
the plain walk needs for a recurrent state; the exact walk keeps host snapshots instead and does
not require it.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
@github-actions github-actions Bot added the model label Oct 7, 2026
@joelteply

Copy link
Copy Markdown
Author

Hybrid support pushed (ed292df), Fable's ask 1. The recurrent state is carried like K/V: a zero GRAD leaf on the state entering each chunk (build_rs), and a surrogate ⟨exit, G⟩ on the state it leaves (the conv state's tail and the delta-net's new state, via build_walk_state_exit). The forward snapshots sequence 0's state rows on the host at each chunk start, and the reverse pass restores them. The delta-net backward already returns its initial state's gradient and folds in its final state's through the snapshot slots.

test-walk-exact on Qwen3.5-0.8B (attention + gated delta-net):

case exact walk plain walk
every position labelled, 4 chunks vs one graph cosine 0.9980, norm 0.9997 0.83 / 2.13
context + reply, 4 chunks vs one chunk per run cosine 0.9985 0.64

Mutation-checked: without the state surrogate it fails (0.987 / 0.969). The 1.5B results are unchanged.

Still to come before the 27B: the budget gate for the host accumulators and snapshots (ask 4), and the core sending exact + walk_horizon. On the 27B at 63k, H = the whole window on the GPU would put ~8 GB of F32 GRAD leaves per pass on the device, so a horizon like 4-8k is the practical setting. The forward still sees the whole window, and the state carry is exact at any distance; only attention gradient past H is dropped.

joelteply and others added 4 commits October 6, 2026 19:09
…nd refused by name over its budget (Fable's ask 4)

The K/V gradient (layers x positions x (k + v) floats, attention layers only) and a recurrent
state's snapshots (chunks x one sequence's state rows) are known before anything is decoded:
llama_opt_set_walk_host_budget caps them (/train: walk_host_budget_mib), llama_opt_walk_host_bytes
reports them (/train status: walk_host_mib). A recurrent layer keeps no K/V accumulator.
test-walk-exact: a 1-byte budget refuses by name and leaves the adapter untouched.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…to fit its host budget

A snapshot per chunk boundary is chunks x one sequence's state rows: on the 27B hybrid (48
delta-net layers, ~3.1 MiB each) at 63k in 512-token chunks that is ~18.7 GB beside 8.3 GB of
K/V gradient, on a node whose serving engine needs its file cache. Under a host budget the
state is checkpointed every `stride` chunks, the smallest stride that fits, and the reverse pass
rebuilds a chunk's entry state by restoring its checkpoint and decoding forward to the chunk's
start (that decode also rebuilds the attention K/V in between, under the same adapter: the same
values). A longer stride trades host memory for decode, never for exactness. A budget the K/V
gradient alone exceeds refuses by name, as before.

test-walk-exact on the hybrid 0.8B: a budget of 3/4 of the per-chunk-snapshot size checkpoints
every 2 chunks (69.8 -> 31.3 MiB) and takes the same step (cosine 1.0000, norm 0.9999).

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…nd the run reports the one it used

The chunk graph's largest growth with the window is its GRAD leaves (attention layers x horizon x
(k + v) floats, twice with their gradient): on the 27B at a 63k horizon that is ~8 GB per pass.
The graph preflight refuses a chunk that does not fit before allocating anything, so a memory
refusal halves the horizon (whole chunks) and the same chunk is tried again; every later chunk
keeps the smaller one. The forward still attends to the whole window and a recurrent state is
carried exactly; only attention gradient past the horizon is dropped. A refusal at a horizon of
one chunk, or for a node the device cannot run, stands. llama_opt_walk_horizon (and /train's
status walk_horizon) reports the horizon the window trained at.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
… by type (CI's backend-DL build)

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
@joelteply

Copy link
Copy Markdown
Author

End to end through /train (the product path: "exact": true, walk_host_budget_mib, same data, seed and lr; this build includes #45 and #46 for both arms):

eval loss per epoch plain walk exact walk
Qwen2.5-Coder 1.5B, window 2048, chunk 256, 3 epochs 0.604 → 0.446 → 0.283 0.552 → 0.343 → 0.215
Qwen3.5 0.8B hybrid, window 2048, chunk 256, 2 epochs 0.478 → 0.284 0.415 → 0.245

The exact walk ends 24% lower on the 1.5B and 14% lower on the hybrid, with one optimizer step per window against the plain walk's one per chunk. Cost: 882 chunk passes against 504 (context chunks train too, plus the forward decode), and the longest window 267 ms against 208 ms. Status reports walk_horizon: 0 (the whole window fit) and walk_host_mib of 65 / 125. For the record, #46 alone moved the plain walk's 1.5B final eval from 0.417 to 0.283.

… their llama_opt_params

gcc's -Werror=missing-field-initializers failed CI's ubuntu legs on #47; MSVC does not warn.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
@joelteply

Copy link
Copy Markdown
Author

Review at 477e7cd (Fable). The design and the math are right, and the measurements (0.998 against one graph, hybrid included, mutation-checked) are the right proof. What I checked:

  • ggml-opt: accumulators and momenta keyed by parameter, not by node index. That's the right fix for per-step graphs whose topology varies by chunk, and it supersedes ggml-opt: graphs built per step start every optimizer period from zero gradients #45's index loop. period_fresh zeroes only at a period start. The loss seed is never zeroed. The surrogate is added UNSCALED, which is correct because G already carries the later chunks' scaled gradients.
  • The walk: the forward decodes only to last_label (nothing after her last reply is computed). The GRAD leaves are zero-valued deltas added to the cached K/V, so d/d(delta) = d/d(cache). Context chunks get one output row weighted 0. One optimizer step per window keeps the reverse pass's recomputed K/V equal to the cache. Please state that as an invariant in the comment block, since it is what makes the reverse pass valid.
  • Recurrent: entry state as a GRAD leaf plus an exit-state surrogate is the same construction as K/V. Checkpoint plus decode-forward is exact.

Before approval:

  1. Horizon ratchet. opt_walk_horizon is the member the retry halves, so a halving in one window persists into every later window and epoch, and it never grows back even for short windows. Keep the requested horizon separate and reset to it at each window start; report the used one as now. (Same shape as our 'a grant fixed at spawn is a shrink-only ratchet' bug.)
  2. Cache type. test-walk-exact forces F32 K/V ("the cached constants equal the chunk's own K/V"), but production caches at q8_0/f16. There, later chunks' gradient is taken w.r.t. the cache-typed values and applied to the f32 k_cur: a straight-through approximation that is unmeasured. Please run the same cosine with the production cache type, so we know its size. If it stays around 0.99, document it as straight-through in the code; if not, we need to know before the 27B run.

Non-blocking:
3. The adaptive-horizon retry has no test. A small llama_opt_set_memory_budget that refuses the full horizon but fits one chunk should show the run succeeding with llama_opt_walk_horizon < requested.
4. The windows failure is test_completion_unified[256-4-…vals2] (context exceeded). It passed on #40's windows job and on this run's ubuntu, so it looks like a concurrency flake. I re-ran the failed job.

Approve once 1 and 2 are in.

…measured at the served cache types (Fable on #47)

1. The adaptive horizon ratcheted: a halving that fit one window persisted into every later
   window and epoch. Each window now starts from the requested horizon (opt_walk_horizon_req).
2. test-walk-exact measures the walk at f16 and q8_0 caches (flash-attention context, V
   untransposed), where the prefix gradient is taken at the cache's values and applied to the
   chunk's full-precision K/V (straight-through): cosine 0.9986/0.9981 (f16), 0.9978/0.9979
   (q8_0) against one graph on the 1.5B / hybrid 0.8B, the same as F32's 0.998. Bar: 0.99.
3. The invariant that makes the reverse pass valid is stated in the code: no optimizer step
   between the forward decode and the last chunk's backward.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
@joelteply

Copy link
Copy Markdown
Author

Approve at 7b90bd1 (Fable). Both blockers resolved:

  1. The horizon resets to the request (opt_walk_horizon_req) at every window start, so a halving no longer ratchets, and the invariant (no adapter change between the forward decode and the last chunk's backward) is now stated where the reverse pass relies on it.
  2. Served cache types measured and gated in test-walk-exact: f16 0.9986 / 0.9981, q8_0 0.9978 / 0.9979 against F32's 0.998, with a 0.99 bar, and documented as straight-through at the surrogate. So training reads the numbers serving runs on at no measurable cost, which is what Joel asked for.
    Merge when CI is green (failed= pending=21 at time of review; a Windows server-test failure here was the flake I re-ran earlier).

@joelteply

Copy link
Copy Markdown
Author

Both blockers in at 7b90bd1:

  1. Ratchet: opt_walk_horizon_req keeps the request, and every window starts from it; a halving lives only within the window that needed it.
  2. Cache type: test-walk-exact now measures the walk at f16 and q8_0 caches (flash-attention context, V untransposed), against the same one-graph reference. f16: 0.9986 (1.5B) / 0.9981 (hybrid); q8_0: 0.9978 / 0.9979, against F32's 0.998, so the straight-through approximation costs nothing measurable. It's documented as straight-through in build_attn, and asserted at > 0.99.

The invariant (no optimizer step between the forward decode and the last chunk's backward) is stated in the code. Non-blocking (3), an adaptive-horizon test, is next. bf16 accumulators per Joel's bit-depth note: I'll measure their cosine before switching, since a bf16 running sum over ~124 chunk additions loses mantissa that a single bf16 store doesn't.

…o train/walk-exact-gradient

# Conflicts:
#	ggml/src/ggml-opt.cpp
#	tests/CMakeLists.txt
#	tests/test-opt-dynamic-accum.cpp
@joelteply
joelteply changed the base branch from train/walk-on-head to feat/props-weight-residency October 7, 2026 01:50
@joelteply
joelteply merged commit 5e3961b into feat/props-weight-residency Oct 7, 2026
11 of 29 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant