diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 27a26d0948c9..80eaeb1336df 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" @@ -582,6 +585,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 +2345,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,10 +3340,14 @@ 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"); - GGML_ASSERT(!cparams.flash_attn && "training needs flash attention off: FLASH_ATTN_EXT has no backward"); + // 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"); + // 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). @@ -3404,25 +3424,82 @@ 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); 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; + } + + // 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 { + 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) { + 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) { + 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); + 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 +3510,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 +3520,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 +3554,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(); @@ -3513,7 +3586,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); @@ -3544,15 +3617,86 @@ 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_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; + }; + + // 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; + 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); + // 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 (!yield_point()) { + return; + } + 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) { + // 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; + } + if (!decode_span(c0, c1)) { + opt_stop_requested.store(true); + return; + } + } + } + } + pos = end; } } diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 64bf2e0b88f0..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"); @@ -2836,14 +2839,45 @@ 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); + // 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] + 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); + 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 + // (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..10e472fcb5ed 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()); @@ -728,17 +743,33 @@ 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). + // 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 = window; - cparams.n_ubatch = window; - cparams.n_seq_max = 1; - cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; - cparams.type_k = GGML_TYPE_F32; - cparams.type_v = GGML_TYPE_F32; + cparams.n_batch = chunk; + cparams.n_ubatch = chunk; + // 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; + // 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_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 @@ -876,6 +907,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) {