Skip to content

train: walk the window — decode context into the cache, train each reply chunk against it - #40

Merged
joelteply merged 7 commits into
feat/props-weight-residencyfrom
train/walk-on-head
Oct 7, 2026
Merged

joelteply merged 7 commits into
feat/props-weight-residencyfrom
train/walk-on-head

Conversation

@joelteply

Copy link
Copy Markdown

What it is: the training walk (Fable's design). A window's context is decoded into the KV cache, and each reply chunk is trained against the cached prefix, which it reads as constants. The chunks are then re-decoded, so the next chunk sees them. A training graph therefore holds one chunk, not the whole window. That's what lets Kimi's 63k-token windows train without truncating anything, which Joel's rule forbids.

Contents:

  • The walk itself, in opt_epoch_iter.
  • A snapshot of the recurrent state around each chunk (Fable), so a reply can span chunks on a hybrid.
  • Cache constants at F16 or q8_0. Training graphs build explicit attention, and the serving context keeps flash attention.
  • A yield before every decode piece and every chunk, so serving waits at most one chunk.
  • /train takes a chunk parameter (multiple of 256; interim default 512).

Measured on the 5090, rebased onto dc1a379:

Run Losses (train / eval per epoch) Windows Longest window Graph
1.5B, chunk 256 0.680/0.621, 0.422/0.455, 0.325/0.483 (eval identical before and after the rebase) 504 124 ms 692 MiB
1.5B, long examples 0.442 98 182 ms 1,249 MiB
0.8B hybrid, recompute off 0.702/0.496, 0.341/0.290 338 429 ms 718 MiB
0.8B hybrid, recompute on 0.701/0.496, 0.342/0.293 338 1,011 ms 408 MiB

Why it's a draft: the gradient is truncated. Cached prefixes are stop-gradient constants, so a chunk's loss doesn't reach the earlier context's activations. Before a gene ships from this, Fable wants an exact reverse pass with horizon H. The test plan: H = 0, 1k, 4k and all, on the 1.5B parity set, accepted when H = all matches the full-window loss (0.222). Also open: the qwen35 recompute gap (card 9355b90c). The two questions are separate.

🤖 Generated with Claude Code

https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc

joelteply and others added 7 commits October 6, 2026 10:32
…t the cached prefix (WIP)

Fable's design (continuum, 2026-10-06): walk her conversation once in one context;
context she did not write is a plain decode into the cache; her reply trains in
chunks, each attending to everything before it read from the cache as a constant;
then the chunk is decoded into the cache under the adapter as it now is. Memory
is chunk x window, not window x window, so a 63-73k lived window can train.

- build_attn (training): concat(cached prefix K/V, this chunk's K/V); V transposed
  out of the non-FA cache; the mask sliced to prefix+chunk and made contiguous.
- opt_epoch_iter: the walk (decode_span / train_chunk), nothing after the last label.
- opt_init: one ubatch per batch, not per context.
- graph_max_nodes: the prefix nodes; the training budget follows opt_ctx, not the
  per-graph flag (a decode's sched_reserve resized it under training=false).
- sched_reserve: a training context keeps its scheduler (the walk's first decode
  re-created it and left ggml-opt on a freed one: GGML_ASSERT(backend)).
- server-train: "chunk", the largest multiple of 256 dividing the window.

Measured on the 5090, Qwen2.5-Coder-1.5B Q4_K_M, 32 single-reply examples (window
1280), lr 1e-4, 3 epochs:
- forward (lr 1e-9): old 0.74658 / eval 0.82223; walk 0.74915 / 0.82235 (correct)
- eval: old 0.512 -> 0.247 -> 0.222; walk 0.628 -> 0.452 -> 0.458
- per epoch: old 15-24 s; walk 5-8 s
The stop-gradient at the cache costs about half the learning on replies that draw
on context. WIP: the WALK fprintf debug lines stay until exact-gradient work lands.
…e-decode stops the run loudly

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…able), so a reply spans chunks on a hybrid

The training forward advances a recurrent state in place, and a recurrent memory
cannot remove a partial range, so a reply longer than one chunk could not be
re-decoded on a hybrid model. Before a training chunk, the recurrent part's state
is copied to the scratch sequence 1; after it, the state is restored, the hybrid's
attention part drops the chunk's empty cells, and the chunk is decoded for real.
The training context gets n_seq_max 2, unified (one attention stream of the window).

Forward check on the 5090 (Qwen2.5-Coder-1.5B Q4_K_M, lr 1e-9, 32 single-reply
examples, window 1280): chunk=window 0.74915 / eval 0.82235; chunk=256 (replies
split across chunks, re-decoded between) 0.74668 / eval 0.82669; old engine
0.74658 / 0.82223.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
With the walk, the training context's cache holds only CONSTANTS (the context
before a chunk); the chunk's own K/V carry the gradient in the graph at F32 and are
never cached during training. So the cache is F16, cast to the chunk's type where
build_attn joins them, which halves the window's cache: about 8 GB to 4 GB at Kimi's
63k on the 27B, the margin beside her live lane.

A caller that sends no "chunk" gets 512 until the core passes the lease's S
(Fable). The window's length as one chunk is the window x window memory the walk
exists to avoid.

Forward check, 1.5B, lr 1e-9, chunk 256: 0.74845 / eval 0.82275 (old engine 0.74658
/ 0.82223). Learning, chunk 256, 3 epochs: eval 0.626 -> 0.456 -> 0.430, the same
truncated-gradient floor as at F32.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…aphs build explicit attention

A quantized V cache requires flash attention (llama_init_from_model refuses it
otherwise), and FLASH_ATTN_EXT has no backward. So the training context is created
WITH flash attention: the walk's context decodes use it, and the cache stores V
un-transposed at q8_0, the same as serving. Every training graph builds explicit
attention whatever the context's flag (build_attn_mha under cparams.training), so
opt_init no longer asserts flash attention off. At Kimi's 63k on the 27B the
cache is ~2.3 GB instead of ~4.6 (F16) or ~9 (F32).

1.5B, chunk 256: forward (lr 1e-9) 0.74694 / eval 0.82323 (old engine 0.74658 /
0.82223); learning, 3 epochs: eval 0.621 -> 0.454 -> 0.483 (the truncated floor).

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
Measured on the 5090 (1.5B, 3 examples of ~15k, recompute on): one yield per example
made a turn arriving mid-walk wait for the whole walk (max 1173-1758 ms). Yielding
before every context-decode piece and every training chunk makes the chunk the
window: max 265-345 ms, p95 137-178 ms, same loss (0.4415-0.4416). Open: epoch time
rose ~80% with no serving load (3.4 -> 6.1 s at chunk 1024), unexplained so far.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc
…ure path (returns false, like a refused graph)

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

Copy link
Copy Markdown
Author

Rebased onto 2d4d63b (#41 and #42). #42's backend-failure check now also covers the walk's chunk step. Git had merged it in with a bare return; inside a lambda that returns bool, so I fixed it to return false, the same exit a refused graph takes. Rerun on the 5090, 1.5B, chunk 256: eval 0.6211 / 0.4545 / 0.4830, identical to before.

@joelteply

Copy link
Copy Markdown
Author

Approve at 539c75d (Fable). This is the walk as designed: unlabelled runs decoded into the cache; her reply trained in chunks against the cached prefix as constants; each chunk then decoded into the cache under the adapter as it now is; nothing past last_label computed; the yield point is the chunk, not the window. The recurrent snapshot via the scratch sequence (n_seq_max = 2, unified) is right, as are flash attention for the decodes and explicit attention for the training graphs (!cparams.training in build_attn_mha), so V is cached untransposed and can be quantized. The early sched_reserve return under opt_ctx closes the freed-scheduler trap.

Two notes, both about Joel's 'align bit depth to inference' (tonight). Fine to fix on top in #47 or a follow-up:

  1. server-train.cpp hardcodes type_k = type_v = GGML_TYPE_Q8_0, overriding whatever common_context_params_to_llama(params_base) took from the serving context. Training should read the cache type serving runs at, so drop the override and inherit it. train: the exact walk, one optimizer step per window on the gradient of its whole loss (opt-in) #47's measurement says both served types are safe (f16 0.9986, q8_0 0.9978).
  2. The comment above that override still says 'F16: the cache holds only CONSTANTS…' while the code sets Q8_0. With (1) it becomes 'the serving context's cache type'.

@joelteply
joelteply marked this pull request as ready for review October 7, 2026 01:47
@joelteply
joelteply merged commit 95b41b0 into feat/props-weight-residency Oct 7, 2026
10 of 25 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant