From 69a2f8f96cdb7b3f1880b1103b6e8b01c8433647 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 06:20:19 -0500 Subject: [PATCH 1/7] train: walk the window: decode context, train each reply chunk against 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. --- src/llama-context.cpp | 148 ++++++++++++++++++++++++++-------- src/llama-graph.cpp | 36 +++++++-- tools/server/server-train.cpp | 27 ++++++- 3 files changed, 170 insertions(+), 41 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 27a26d0948c9..669f740250b0 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -582,6 +582,13 @@ void llama_context::sched_reserve() { if (!sched_need_reserve) { return; } + // A training context's scheduler was sized for the training graph in opt_init, and the + // optimizer holds it (ggml_opt_params.backend_sched). Re-creating it here, as the first + // context decode of opt_epoch_iter's walk would, leaves the optimizer on a freed scheduler. + if (opt_ctx != nullptr) { + sched_need_reserve = false; + return; + } sched_need_reserve = false; @@ -2335,7 +2342,13 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const { if (n_sampling_outputs_max > 1) { res += (n_sampling_outputs_max - 1) * n_sampling_nodes_max; } - if (cparams.training) { + // a training context (an optimizer exists), not this graph's flag: the walk decodes context + // with cparams.training off, and a sched_reserve() run then must not shrink the scheduler + // and the graph result below what the next training chunk needs + if (cparams.training || opt_ctx != nullptr) { + // each attention layer reads the cached context in front of the chunk (build_attn: + // views of K and V, V's permute and cont, two concats, the mask's view and cont) + res += 16u * model.hparams.n_layer(); // ggml_opt duplicates the forward graph at its own size, then appends the backward // pass and the optimizer step: room for the forward three times over. res *= 3; @@ -3324,9 +3337,11 @@ void llama_context::opt_init(struct llama_model * model, struct llama_opt_params GGML_ASSERT(opt_n_ctx_train % n_batch == 0); GGML_ASSERT(n_batch % n_ubatch == 0); // A training graph cannot backprop through the KV cache (SET_ROWS writes in place and - // attention reads the cache as a leaf), so attention consumes this ubatch's K/V - // directly; that is exact only when one ubatch is the whole context. - GGML_ASSERT(n_ubatch == opt_n_ctx_train && "training needs one ubatch per context: set -ub = -b = -c"); + // attention reads the cache as a leaf), so attention consumes this chunk's K/V directly + // and the context before it from the cache as a constant (build_attn); opt_epoch_iter + // walks the window decoding context and training each labelled chunk. One chunk is one + // batch is one ubatch: a chunk's backward needs all of its own K/V in one graph. + GGML_ASSERT(n_ubatch == n_batch && "training needs one ubatch per batch: set -ub = -b"); GGML_ASSERT(!cparams.flash_attn && "training needs flash attention off: FLASH_ATTN_EXT has no backward"); cparams.training = true; // The graph result and the scheduler were sized at context creation: before training @@ -3404,25 +3419,66 @@ void llama_context::opt_epoch_iter( int64_t ndata_in_loop, int64_t t_loop_start) { GGML_ASSERT(opt_ctx); - const uint32_t n_ctx = opt_n_ctx_train; - const uint32_t n_batch = std::min(this->n_batch(), n_ctx); - const uint32_t n_ubatch = std::min(this->n_ubatch(), n_batch); + const uint32_t n_ctx = opt_n_ctx_train; + const uint32_t n_batch = std::min(this->n_batch(), n_ctx); + fprintf(stderr, "WALK start: n_ctx %u n_batch %u\n", n_ctx, n_batch); fflush(stderr); memory->clear(true); - // OUTPUTS ONLY WHERE THE LOSS IS. A position whose label is masked (< 0: the system and - // tool head, the user's turns, padding) carries no loss, so its logits are never read; - // computing them anyway cost n_vocab x window floats for the logits, the same again for - // the one-hot labels, and the same again for the logit gradients. At 151,936 x 256 that - // was the 148 MiB the memory gate named on the M5; at one of a citizen's 15k-token - // turns it was ~9 GB three times over, and "could not create the training context" on - // the 5090 (continuum card 36c3c00a). Decode already selects output rows per token - // (batch.logits); training now does the same, and the label tensor is sized to the - // rows that exist. Every window has at least one labelled position: the server refuses - // an example with none before it reaches here. - for (uint32_t pos_ctx = 0; pos_ctx < n_ctx; pos_ctx += n_batch) { - batch.n_tokens = n_batch; - for (uint32_t pos_batch = 0; pos_batch < n_batch; ++pos_batch) { + // THE WALK (Fable's design, continuum 2026-10-06): one pass over the window in one + // context. A run with no label is context she did not write: a plain decode into the + // cache, no backward. A labelled run is her reply: it trains in chunks of n_batch, each + // attending to everything before it read from the cache as a constant (build_attn, + // cparams.training), then the reply is decoded into the cache under the adapter as it + // now is, so the next reply sees it. Memory is chunk x window, never window x window, so + // her whole lived window trains. Nothing after the last label is computed. + int64_t last_label = -1; + for (int64_t i = (int64_t) n_ctx - 1; i >= 0; --i) { + if (labels_sparse[i] >= 0) { + last_label = i; + break; + } + } + if (last_label < 0) { + LLAMA_LOG_ERROR("%s: a training window with no labelled position: nothing to learn from it\n", __func__); + return; + } + + // a plain inference decode of [p0, p1) into the cache: the training graph's attention + // never writes the cache, so context and finished replies reach it this way + auto decode_span = [&](uint32_t p0, uint32_t p1) -> bool { + fprintf(stderr, "WALK decode context [%u, %u)\n", p0, p1); fflush(stderr); + const bool training = cparams.training; + cparams.training = false; + gf_res_prev->reset(); + bool ok = true; + for (uint32_t c0 = p0; c0 < p1 && ok; c0 += n_batch) { + const uint32_t c1 = std::min(c0 + n_batch, p1); + batch.n_tokens = c1 - c0; + for (uint32_t i = 0; i < c1 - c0; ++i) { + batch.token [i] = tokens[c0 + i]; + batch.pos [i] = c0 + i; + batch.n_seq_id[i] = 1; + batch.seq_id [i][0] = 0; + batch.logits [i] = i + 1 == c1 - c0; + } + const int rc = decode(batch); + fprintf(stderr, "WALK decode [%u, %u) -> %d\n", c0, c1, rc); fflush(stderr); + ok = rc == 0; + } + cparams.training = training; + gf_res_prev->reset(); + if (!ok) { + LLAMA_LOG_ERROR("%s: the context decode of [%u, %u) failed\n", __func__, p0, p1); + } + return ok; + }; + + // one training chunk [pos_ctx, pos_ctx + n_tokens): forward and backward over it alone, + // attending to the cached context before it + auto train_chunk = [&](uint32_t pos_ctx, uint32_t n_chunk) -> bool { + batch.n_tokens = n_chunk; + for (uint32_t pos_batch = 0; pos_batch < n_chunk; ++pos_batch) { batch.token [pos_batch] = tokens[pos_ctx + pos_batch]; batch.pos [pos_batch] = pos_ctx + pos_batch; batch.n_seq_id[pos_batch] = 1; @@ -3433,7 +3489,7 @@ void llama_context::opt_epoch_iter( // output_all = false: the batch's own logits flags decide which rows exist if (!balloc->init(batch, model.vocab, nullptr, model.hparams.n_embd_inp(), cparams.kv_unified ? LLAMA_MAX_SEQ : cparams.n_seq_max, false)) { LLAMA_LOG_ERROR("%s: failed to initialize batch\n", __func__); - return; + return false; } const uint32_t n_tokens_all = balloc->get_n_tokens(); @@ -3443,21 +3499,17 @@ void llama_context::opt_epoch_iter( embd_seq.clear(); const uint32_t n_outputs_all = balloc->get_n_outputs(); - // Unreachable from /train: it sets n_ctx = n_batch = n_ubatch = window (one ubatch is - // the whole context, asserted in opt_init) and refuses an example with no labelled - // position before it gets here, so every batch has at least one. Kept for any other - // caller of opt_epoch, and it RETURNS from the whole iteration (a later batch of this - // context would need the earlier one's state; with one batch per context there is - // none to lose). (Cormac and BigMama on fork #32.) + // Unreachable: the walk below trains only chunks of a labelled run. It RETURNS from the + // whole window, since a later chunk would need this one in the cache. if (n_outputs_all == 0) { LLAMA_LOG_ERROR("%s: a training batch with no labelled position (every label masked): nothing to learn from it\n", __func__); - return; + return false; } auto mctx = memory->init_batch(*balloc, cparams.n_ubatch, true); if (!mctx || mctx->get_status() != LLAMA_MEMORY_STATUS_SUCCESS) { LLAMA_LOG_ERROR("%s: could not initialize batch\n", __func__); - break; + return false; } // reserve output buffer @@ -3481,7 +3533,7 @@ void llama_context::opt_epoch_iter( if (!mctx->apply()) { LLAMA_LOG_ERROR("%s: failed to update the memory context\n", __func__); - break; + return false; } auto * res = gf_res_prev.get(); @@ -3503,6 +3555,8 @@ void llama_context::opt_epoch_iter( }; ctx_compute_opt = ggml_init(params); } + fprintf(stderr, "WALK chunk at %u (%u tokens, n_past %d): forward graph %d nodes of %d\n", + pos_ctx, ubatch.n_tokens, (int) ubatch.pos[0], ggml_graph_n_nodes(gf), (int) ggml_graph_size(gf)); fflush(stderr); ggml_opt_prepare_alloc(opt_ctx, ctx_compute_opt, gf, res->get_inp_tokens(), res->get_logits()); if (!ggml_opt_alloc(opt_ctx, train)) { const char * why = ggml_opt_refusal(opt_ctx); @@ -3513,7 +3567,7 @@ void llama_context::opt_epoch_iter( ggml_free(ctx_compute_opt); opt_alloc_failed.store(true); opt_stop_requested.store(true); - return; + return false; } res->set_inputs(&ubatch); @@ -3547,12 +3601,42 @@ void llama_context::opt_epoch_iter( return; } if (callback) { - callback(train, opt_ctx, dataset, result, idata_in_loop + (pos_ctx + pos_batch)/n_ubatch + 1, ndata_in_loop, t_loop_start); + callback(train, opt_ctx, dataset, result, idata_in_loop + (pos_ctx + pos_batch)/n_batch + 1, ndata_in_loop, t_loop_start); } ggml_free(ctx_compute_opt); pos_batch += ubatch.n_tokens; } while (mctx->next()); + return true; + }; + + uint32_t pos = 0; + while (pos <= (uint32_t) last_label && !opt_stop_requested.load(std::memory_order_relaxed)) { + const bool labelled = labels_sparse[pos] >= 0; + uint32_t end = pos; + while (end <= (uint32_t) last_label && (labels_sparse[end] >= 0) == labelled) { + ++end; + } + if (!labelled) { + if (!decode_span(pos, end)) { + return; + } + } else { + for (uint32_t c0 = pos; c0 < end; c0 += n_batch) { + const uint32_t c1 = std::min(c0 + n_batch, end); + if (!train_chunk(c0, c1 - c0)) { + return; + } + // the chunk joins the context under the adapter as it now is: the training + // graph never writes the cache, so its cells hold nothing until this decode + if (c1 <= (uint32_t) last_label) { + if (!memory->seq_rm(0, c0, -1) || !decode_span(c0, c1)) { + return; + } + } + } + } + pos = end; } } diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 64bf2e0b88f0..a0197655b96f 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2836,14 +2836,38 @@ ggml_tensor * llm_graph_context::build_attn( ggml_tensor * v; if (cparams.training) { // A backward pass cannot go through the cache: the write is an in-place SET_ROWS - // and the read is a view of a leaf, so K/V would get no gradient. Training runs one - // ubatch per context (llama_context::opt_init asserts it), so this ubatch's K/V ARE - // the whole context: attend to them directly, as the no-cache path does. + // and the read is a view of a leaf. So this ubatch's K/V are attended to directly, + // and carry the gradient. The context before it (positions [0, n_past), decoded into + // the cache by llama_context::opt_epoch_iter before this chunk trains) is read from + // the cache as a CONSTANT, concatenated in front: her reply sees the whole window, + // and memory is chunk x window instead of window x window. The prefix's own K/V get + // no gradient (a stop-gradient at the chunk boundary). + const int64_t n_tokens = k_cur->ne[2]; + const int64_t n_past = ubatch.pos[0]; k = k_cur; v = v_cur; - if (kq_mask->ne[0] > k_cur->ne[2]) { - kq_mask = ggml_view_4d(ctx0, kq_mask, k_cur->ne[2], kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3], - kq_mask->nb[1], kq_mask->nb[2], kq_mask->nb[3], 0); + if (n_past > 0) { + ggml_tensor * k_all = mctx_cur->get_k(ctx0, il); + ggml_tensor * k_prev = ggml_view_4d(ctx0, k_all, k_all->ne[0], k_all->ne[1], n_past, 1, + k_all->nb[1], k_all->nb[2], k_all->nb[3], 0); + k = ggml_concat(ctx0, k_prev, k_cur, 2); + + // a cache without flash attention stores V transposed: [n_kv, n_head_kv, n_embd_head_v] + ggml_tensor * v_all = mctx_cur->get_v(ctx0, il); + ggml_tensor * v_prev = v_all->nb[1] > v_all->nb[2] + ? ggml_cont(ctx0, ggml_permute(ctx0, + ggml_view_4d(ctx0, v_all, n_past, v_all->ne[1], v_all->ne[2], 1, + v_all->nb[1], v_all->nb[2], v_all->nb[3], 0), + 2, 1, 0, 3)) + : ggml_view_4d(ctx0, v_all, v_all->ne[0], v_all->ne[1], n_past, 1, + v_all->nb[1], v_all->nb[2], v_all->nb[3], 0); + v = ggml_concat(ctx0, v_prev, v_cur, 2); + } + // the cache's cells hold the prefix at [0, n_past) and this ubatch right after it + // (contiguous: soft_max_ext requires it, and a column slice of the padded mask is not) + if (kq_mask->ne[0] > n_past + n_tokens) { + kq_mask = ggml_cont(ctx0, ggml_view_4d(ctx0, kq_mask, n_past + n_tokens, kq_mask->ne[1], kq_mask->ne[2], kq_mask->ne[3], + kq_mask->nb[1], kq_mask->nb[2], kq_mask->nb[3], 0)); } } else { // store to KV cache diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index affdb3349bc1..930ea52bed71 100644 --- a/tools/server/server-train.cpp +++ b/tools/server/server-train.cpp @@ -314,7 +314,7 @@ json server_trainer::start(const json & body_in) { // assert in the training path would take the server down (Cormac on #14). { auto num = [&](const char * key, double def) { return body.contains(key) && body.at(key).is_number() ? body.at(key).get() : def; }; - for (const char * key : {"rank", "alpha", "window", "epochs", "lr", "val_split", "seed", "memory_budget_mib", "top_layers", "share_ppm", "max_slowdown_ppm"}) { + for (const char * key : {"rank", "alpha", "window", "epochs", "lr", "val_split", "seed", "memory_budget_mib", "top_layers", "share_ppm", "max_slowdown_ppm", "chunk"}) { if (body.contains(key) && !body.at(key).is_number()) { return json::object({{"ok", false}, {"error", std::string("\"") + key + "\" must be a number"}}); } @@ -336,6 +336,8 @@ json server_trainer::start(const json & body_in) { else if (epochs < 1 || epochs > 100 || epochs != (int64_t) epochs) why = "epochs must be an integer in [1, 100]"; else if (!(lr > 0 && lr <= 1)) why = "lr must be in (0, 1]"; else if (!(val >= 0 && val < 1)) why = "val_split must be in [0, 1)"; + else if (body.contains("chunk") && !(num("chunk", 0) >= 256 && num("chunk", 0) == (int64_t) num("chunk", 0) && (int64_t) num("chunk", 0) % 256 == 0)) + why = "chunk must be a multiple of 256, at least 256: the most tokens one training step holds in one graph"; else if (body.contains("memory_budget_mib") && !(num("memory_budget_mib", 0) >= 1)) why = "memory_budget_mib must be >= 1"; else if (body.contains("share_ppm") && !(num("share_ppm", 0) >= 1 && num("share_ppm", 0) <= 1000000 && num("share_ppm", 0) == (int64_t) num("share_ppm", 0))) @@ -698,6 +700,19 @@ json server_trainer::status() const { return s; } +// The training chunk: the largest multiple of 256 that divides the window and is at most the +// asked chunk (at least 256; the window is a multiple of 256, so 256 always divides it). The +// context must be a whole number of batches (llama_context::opt_init). +static uint32_t train_chunk_for(uint32_t window, uint32_t asked) { + const uint32_t blocks = window / 256; + for (uint32_t g = std::min(blocks, std::max(asked / 256, 1u)); g >= 1; --g) { + if (blocks % g == 0) { + return g * 256; + } + } + return 256; +} + void server_trainer::run(json req, examples_data ex) { auto fail = [&](const std::string & why) { LOG_ERR("%s: training run failed: %s\n", __func__, why.c_str()); @@ -732,9 +747,14 @@ void server_trainer::run(json req, examples_data ex) { // (the training graph attends to this ubatch's K/V directly); flash attention has no // backward; the KV cache types are F32 because OUT_PROD has no F16 path. llama_context_params cparams = common_context_params_to_llama(params_base); + // The context holds the whole window; one training chunk is one batch is one ubatch, the + // largest multiple of 256 that divides the window and is at most the caller's "chunk" (the + // memory gate's S). Context she did not write is decoded; her replies train chunk by chunk, + // each attending to everything before it from the cache (llama_context::opt_epoch_iter). + const uint32_t chunk = train_chunk_for(window, (uint32_t) req.value("chunk", (int64_t) window)); cparams.n_ctx = window; - cparams.n_batch = window; - cparams.n_ubatch = window; + cparams.n_batch = chunk; + cparams.n_ubatch = chunk; cparams.n_seq_max = 1; cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; cparams.type_k = GGML_TYPE_F32; @@ -876,6 +896,7 @@ void server_trainer::run(json req, examples_data ex) { std::lock_guard lock(mu); state["state"] = "running"; state["window"] = window; + state["chunk"] = chunk; state["tokens"] = n_tokens; state["train_tokens"] = train_tokens; if (by_example) { From 0540511f97dee927e966ccd35c56bc0c2c86a55b Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 06:22:17 -0500 Subject: [PATCH 2/7] train: drop the walk's debug prints; a reply a hybrid memory cannot re-decode stops the run loudly Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 17 +++++++++++------ 1 file changed, 11 insertions(+), 6 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 669f740250b0..82e278694781 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3422,7 +3422,6 @@ void llama_context::opt_epoch_iter( const uint32_t n_ctx = opt_n_ctx_train; const uint32_t n_batch = std::min(this->n_batch(), n_ctx); - fprintf(stderr, "WALK start: n_ctx %u n_batch %u\n", n_ctx, n_batch); fflush(stderr); memory->clear(true); // THE WALK (Fable's design, continuum 2026-10-06): one pass over the window in one @@ -3447,7 +3446,6 @@ void llama_context::opt_epoch_iter( // a plain inference decode of [p0, p1) into the cache: the training graph's attention // never writes the cache, so context and finished replies reach it this way auto decode_span = [&](uint32_t p0, uint32_t p1) -> bool { - fprintf(stderr, "WALK decode context [%u, %u)\n", p0, p1); fflush(stderr); const bool training = cparams.training; cparams.training = false; gf_res_prev->reset(); @@ -3463,7 +3461,6 @@ void llama_context::opt_epoch_iter( batch.logits [i] = i + 1 == c1 - c0; } const int rc = decode(batch); - fprintf(stderr, "WALK decode [%u, %u) -> %d\n", c0, c1, rc); fflush(stderr); ok = rc == 0; } cparams.training = training; @@ -3555,8 +3552,6 @@ void llama_context::opt_epoch_iter( }; ctx_compute_opt = ggml_init(params); } - fprintf(stderr, "WALK chunk at %u (%u tokens, n_past %d): forward graph %d nodes of %d\n", - pos_ctx, ubatch.n_tokens, (int) ubatch.pos[0], ggml_graph_n_nodes(gf), (int) ggml_graph_size(gf)); fflush(stderr); ggml_opt_prepare_alloc(opt_ctx, ctx_compute_opt, gf, res->get_inp_tokens(), res->get_logits()); if (!ggml_opt_alloc(opt_ctx, train)) { const char * why = ggml_opt_refusal(opt_ctx); @@ -3630,7 +3625,17 @@ void llama_context::opt_epoch_iter( // the chunk joins the context under the adapter as it now is: the training // graph never writes the cache, so its cells hold nothing until this decode if (c1 <= (uint32_t) last_label) { - if (!memory->seq_rm(0, c0, -1) || !decode_span(c0, c1)) { + // A recurrent or hybrid memory cannot remove a partial range, so a reply + // longer than one chunk cannot be re-decoded there: stop the run loudly, + // never train the next chunk against cells that hold nothing. + if (!memory->seq_rm(0, c0, -1)) { + LLAMA_LOG_ERROR("%s: a reply longer than one chunk (%u tokens) needs its chunks re-decoded, and this memory cannot remove a partial range (a recurrent or hybrid model): raise \"chunk\" to cover the reply\n", + __func__, n_batch); + opt_stop_requested.store(true); + return; + } + if (!decode_span(c0, c1)) { + opt_stop_requested.store(true); return; } } From 321c1fefedf1812fda30dfca937ef3b0e331bef5 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 06:24:27 -0500 Subject: [PATCH 3/7] train: a recurrent state is snapshotted around each training chunk (Fable), 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 43 ++++++++++++++++++++++++++++++----- tools/server/server-train.cpp | 6 ++++- 2 files changed, 42 insertions(+), 7 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 82e278694781..b25c62bf5536 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -7,6 +7,9 @@ #include "llama-batch.h" #include "llama-io.h" #include "llama-memory.h" +#include "llama-memory-hybrid.h" +#include "llama-memory-recurrent.h" +#include "llama-kv-cache.h" #include "llama-mmap.h" #include "llama-model.h" #include "llama-ext.h" @@ -3605,6 +3608,22 @@ void llama_context::opt_epoch_iter( return true; }; + // a recurrent part (pure recurrent or hybrid) is snapshotted around each training chunk; + // a hybrid's attention part removes the chunk's cells like a plain cache + llama_memory_recurrent * recr = nullptr; + llama_kv_cache * attn = nullptr; + if (auto * hybrid = dynamic_cast(memory.get())) { + recr = hybrid->get_mem_recr(); + attn = hybrid->get_mem_attn(); + } else { + recr = dynamic_cast(memory.get()); + } + if (recr != nullptr && cparams.n_seq_max < 2) { + LLAMA_LOG_ERROR("%s: a recurrent model trains with a scratch sequence for its state snapshot: create the training context with n_seq_max >= 2\n", __func__); + opt_stop_requested.store(true); + return; + } + uint32_t pos = 0; while (pos <= (uint32_t) last_label && !opt_stop_requested.load(std::memory_order_relaxed)) { const bool labelled = labels_sparse[pos] >= 0; @@ -3619,18 +3638,30 @@ void llama_context::opt_epoch_iter( } else { for (uint32_t c0 = pos; c0 < end; c0 += n_batch) { const uint32_t c1 = std::min(c0 + n_batch, end); + // The training forward advances a recurrent state in place, and a recurrent + // memory cannot remove a partial range: snapshot the state to the scratch + // sequence first (Fable), restore it after, then the decode below advances it + // from where the context left it. The state is per sequence and small. + const bool snapshot = recr != nullptr && c1 <= (uint32_t) last_label; + if (snapshot) { + recr->seq_cp(0, 1, -1, -1); + } if (!train_chunk(c0, c1 - c0)) { return; } + if (snapshot) { + recr->seq_rm(0, -1, -1); + recr->seq_cp(1, 0, -1, -1); + recr->seq_rm(1, -1, -1); + } // the chunk joins the context under the adapter as it now is: the training // graph never writes the cache, so its cells hold nothing until this decode if (c1 <= (uint32_t) last_label) { - // A recurrent or hybrid memory cannot remove a partial range, so a reply - // longer than one chunk cannot be re-decoded there: stop the run loudly, - // never train the next chunk against cells that hold nothing. - if (!memory->seq_rm(0, c0, -1)) { - LLAMA_LOG_ERROR("%s: a reply longer than one chunk (%u tokens) needs its chunks re-decoded, and this memory cannot remove a partial range (a recurrent or hybrid model): raise \"chunk\" to cover the reply\n", - __func__, n_batch); + // the chunk's cells hold no K/V (the training graph never writes the + // cache): drop them from the attention memory, then decode them for real + const bool removed = recr != nullptr ? (attn == nullptr || attn->seq_rm(0, c0, -1)) : memory->seq_rm(0, c0, -1); + if (!removed) { + LLAMA_LOG_ERROR("%s: could not remove the trained chunk [%u, %u) from the cache before decoding it\n", __func__, c0, c1); opt_stop_requested.store(true); return; } diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index 930ea52bed71..f5140bfa1261 100644 --- a/tools/server/server-train.cpp +++ b/tools/server/server-train.cpp @@ -755,7 +755,11 @@ void server_trainer::run(json req, examples_data ex) { cparams.n_ctx = window; cparams.n_batch = chunk; cparams.n_ubatch = chunk; - cparams.n_seq_max = 1; + // two sequences: 0 is the walk; 1 is the scratch a recurrent state is snapshotted into + // around each training chunk (llama_context::opt_epoch_iter). Unified, so the attention + // cache stays one stream of the window, not one per sequence. + cparams.n_seq_max = 2; + cparams.kv_unified = true; cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; cparams.type_k = GGML_TYPE_F32; cparams.type_v = GGML_TYPE_F32; From c424b00c3ddf9f33df99eb1a3f80e5eb2c2093dd Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 06:27:41 -0500 Subject: [PATCH 4/7] train: the cache holds constants at F16; an interim default chunk of 512 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-graph.cpp | 7 +++++++ tools/server/server-train.cpp | 16 ++++++++++------ 2 files changed, 17 insertions(+), 6 deletions(-) diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index a0197655b96f..a7b5f92c7006 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2850,6 +2850,10 @@ ggml_tensor * llm_graph_context::build_attn( ggml_tensor * k_all = mctx_cur->get_k(ctx0, il); ggml_tensor * k_prev = ggml_view_4d(ctx0, k_all, k_all->ne[0], k_all->ne[1], n_past, 1, k_all->nb[1], k_all->nb[2], k_all->nb[3], 0); + // the cache may hold the constants at a smaller type than the chunk's own K/V + if (k_prev->type != k_cur->type) { + k_prev = ggml_cast(ctx0, k_prev, k_cur->type); + } k = ggml_concat(ctx0, k_prev, k_cur, 2); // a cache without flash attention stores V transposed: [n_kv, n_head_kv, n_embd_head_v] @@ -2861,6 +2865,9 @@ ggml_tensor * llm_graph_context::build_attn( 2, 1, 0, 3)) : ggml_view_4d(ctx0, v_all, v_all->ne[0], v_all->ne[1], n_past, 1, v_all->nb[1], v_all->nb[2], v_all->nb[3], 0); + if (v_prev->type != v_cur->type) { + v_prev = ggml_cast(ctx0, v_prev, v_cur->type); + } v = ggml_concat(ctx0, v_prev, v_cur, 2); } // the cache's cells hold the prefix at [0, n_past) and this ubatch right after it diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index f5140bfa1261..d4e62d4260a8 100644 --- a/tools/server/server-train.cpp +++ b/tools/server/server-train.cpp @@ -743,15 +743,16 @@ void server_trainer::run(json req, examples_data ex) { // governed lease); without it the driver's own free figure, which on Windows is not physical const size_t budget = (size_t) (req.value("memory_budget_mib", 0.0) * 1024.0 * 1024.0); - // The training context: the SAME model, its own graph. One ubatch is the whole window - // (the training graph attends to this ubatch's K/V directly); flash attention has no - // backward; the KV cache types are F32 because OUT_PROD has no F16 path. + // The training context: the SAME model, its own graph; flash attention has no backward. llama_context_params cparams = common_context_params_to_llama(params_base); // The context holds the whole window; one training chunk is one batch is one ubatch, the // largest multiple of 256 that divides the window and is at most the caller's "chunk" (the // memory gate's S). Context she did not write is decoded; her replies train chunk by chunk, // each attending to everything before it from the cache (llama_context::opt_epoch_iter). - const uint32_t chunk = train_chunk_for(window, (uint32_t) req.value("chunk", (int64_t) window)); + // INTERIM: a caller that sends no "chunk" gets 512 until the core passes the lease's S + // (Fable); the window's whole length as one chunk is the window x window memory the walk + // exists to avoid. + const uint32_t chunk = train_chunk_for(window, (uint32_t) req.value("chunk", (int64_t) 512)); cparams.n_ctx = window; cparams.n_batch = chunk; cparams.n_ubatch = chunk; @@ -761,8 +762,11 @@ void server_trainer::run(json req, examples_data ex) { cparams.n_seq_max = 2; cparams.kv_unified = true; cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; - cparams.type_k = GGML_TYPE_F32; - cparams.type_v = GGML_TYPE_F32; + // F16: the cache holds only CONSTANTS (the context before a chunk, read by build_attn and + // cast to F32 where it joins the chunk). The chunk's own K/V, the ones with a gradient and + // the only ones OUT_PROD sees, are F32 in the graph and never cached during training. + cparams.type_k = GGML_TYPE_F16; + cparams.type_v = GGML_TYPE_F16; cparams.embeddings = false; // training reads logits for every token of the window; the serving params cap outputs per // ubatch to what sampling needs (a server-computed limit), which a training batch overruns From a8a235df5199d9be40362e23bf34fd03742e1f92 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 06:33:25 -0500 Subject: [PATCH 5/7] train: q8_0 cache constants in a flash-attention context; training graphs 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 4 +++- src/llama-graph.cpp | 5 ++++- tools/server/server-train.cpp | 9 ++++++--- 3 files changed, 13 insertions(+), 5 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index b25c62bf5536..64838d62984e 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3345,7 +3345,9 @@ void llama_context::opt_init(struct llama_model * model, struct llama_opt_params // walks the window decoding context and training each labelled chunk. One chunk is one // batch is one ubatch: a chunk's backward needs all of its own K/V in one graph. GGML_ASSERT(n_ubatch == n_batch && "training needs one ubatch per batch: set -ub = -b"); - GGML_ASSERT(!cparams.flash_attn && "training needs flash attention off: FLASH_ATTN_EXT has no backward"); + // FLASH_ATTN_EXT has no backward: training graphs build explicit attention whatever the + // context's flag (build_attn_mha checks cparams.training); flash attention, when the context + // has it, serves the walk's context decodes and lets the cache hold V un-transposed. cparams.training = true; // The graph result and the scheduler were sized at context creation: before training // (no backward pass) and before any adapter attached since (its nodes uncounted). diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index a7b5f92c7006..7fef5d164082 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2574,7 +2574,10 @@ ggml_tensor * llm_graph_context::build_attn_mha( ggml_tensor * cur; - const bool use_flash_attn = cparams.flash_attn && kq_b == nullptr; + // A training graph takes the explicit path, which has a backward (FLASH_ATTN_EXT has none), + // even in a context created with flash attention: there the walk's context decodes use + // flash attention and the cache stores V un-transposed, so it can be quantized. + const bool use_flash_attn = cparams.flash_attn && kq_b == nullptr && !cparams.training; if (use_flash_attn) { GGML_ASSERT(kq_b == nullptr && "Flash attention does not support KQ bias yet"); diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index d4e62d4260a8..10e472fcb5ed 100644 --- a/tools/server/server-train.cpp +++ b/tools/server/server-train.cpp @@ -761,12 +761,15 @@ void server_trainer::run(json req, examples_data ex) { // cache stays one stream of the window, not one per sequence. cparams.n_seq_max = 2; cparams.kv_unified = true; - cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; + // Flash attention ON for the context: the walk's context decodes use it, and the cache then + // stores V un-transposed, which a quantized V cache requires. The training graphs build + // explicit attention regardless (build_attn_mha under cparams.training). + cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED; // F16: the cache holds only CONSTANTS (the context before a chunk, read by build_attn and // cast to F32 where it joins the chunk). The chunk's own K/V, the ones with a gradient and // the only ones OUT_PROD sees, are F32 in the graph and never cached during training. - cparams.type_k = GGML_TYPE_F16; - cparams.type_v = GGML_TYPE_F16; + cparams.type_k = GGML_TYPE_Q8_0; + cparams.type_v = GGML_TYPE_Q8_0; cparams.embeddings = false; // training reads logits for every token of the window; the serving params cap outputs per // ubatch to what sampling needs (a server-computed limit), which a training batch overruns From 87e5078850429e57daa9d0f036f60f4817fe8cec Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 08:18:58 -0500 Subject: [PATCH 6/7] train: the walk yields at every chunk, not once per example window 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 64838d62984e..0c94f97852cd 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3448,6 +3448,20 @@ void llama_context::opt_epoch_iter( return; } + // THE YIELD POINT IS THE CHUNK, not the window: a turn that arrives mid-walk waits for one + // chunk, never for the whole example (measured on the 5090: one callback per example made a + // turn wait for its entire walk, 1.2-1.8 s on a 1.5B at 15k and 13.5 s on Fable's run). The + // first unit is the window's own boundary, which opt_epoch has already offered. + bool first_unit = true; + auto yield_point = [&]() -> bool { + if (first_unit) { + first_unit = false; + } else if (opt_step_callback && !opt_step_callback(train, opt_step_callback_data)) { + opt_stop_requested.store(true); + } + return !opt_stop_requested.load(std::memory_order_relaxed); + }; + // a plain inference decode of [p0, p1) into the cache: the training graph's attention // never writes the cache, so context and finished replies reach it this way auto decode_span = [&](uint32_t p0, uint32_t p1) -> bool { @@ -3456,6 +3470,11 @@ void llama_context::opt_epoch_iter( gf_res_prev->reset(); bool ok = true; for (uint32_t c0 = p0; c0 < p1 && ok; c0 += n_batch) { + if (!yield_point()) { + cparams.training = training; + gf_res_prev->reset(); + return false; // stopped or cancelled at a chunk boundary: not a decode failure + } const uint32_t c1 = std::min(c0 + n_batch, p1); batch.n_tokens = c1 - c0; for (uint32_t i = 0; i < c1 - c0; ++i) { @@ -3645,6 +3664,9 @@ void llama_context::opt_epoch_iter( // sequence first (Fable), restore it after, then the decode below advances it // from where the context left it. The state is per sequence and small. const bool snapshot = recr != nullptr && c1 <= (uint32_t) last_label; + if (!yield_point()) { + return; + } if (snapshot) { recr->seq_cp(0, 1, -1, -1); } From 539c75de5f8c0064c4d6de1426de9ac7f4f90b03 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 10:34:13 -0500 Subject: [PATCH 7/7] train: the walk's chunk step fails its run through #42's backend-failure path (returns false, like a refused graph) Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 0c94f97852cd..80eaeb1336df 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3617,7 +3617,7 @@ void llama_context::opt_epoch_iter( ggml_free(ctx_compute_opt); opt_alloc_failed.store(true); opt_stop_requested.store(true); - return; + return false; } if (callback) { callback(train, opt_ctx, dataset, result, idata_in_loop + (pos_ctx + pos_batch)/n_batch + 1, ndata_in_loop, t_loop_start);