From 802d4a29b56da0c05df342a244fe63d2f5cfd3fb Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Fri, 9 Oct 2026 23:56:29 -0500 Subject: [PATCH 1/2] fix(train): the exact walk rewinds the cache before retrying a refused chunk at a smaller horizon A chunk graph the device budget refuses is refused AFTER its ubatch was applied: the chunk's cells already sit in the attention cache, and on a hybrid model the recurrent cell's position has moved to the chunk's end. The adaptive horizon (#47) retried train_chunk on top of that, so the retry could not prepare its ubatch at all, and the epoch stopped instead of shrinking to a horizon that fits. Measured on the 5090 (2026-10-10 04:34Z), Kimi's first exact-walk job on Qwen3.8-27B: "the graph needs 9236.1 MiB more on CUDA0, over the 3418.0 MiB it may add" -> "the chunk at 1155 did not fit with a gradient horizon of 1155 positions: retrying at 512" -> "init_batch: failed to prepare attention ubatches" -> the job failed in 8 s. The reverse pass's rewind (pop attention to [0, c0); restore the recurrent checkpoint and decode to c0) is now one lambda, run before the first try AND before every retry, with the walk flags cleared around the retry's rewind so its decode is the same plain forward as the first. test-walk-exact case 6 (2048-token window): a device budget that lets a chunk graph grow by nothing refuses every horizon; each retry must reach the device preflight again with zero ubatch failures, and the epoch then refuses by name with the adapter untouched. On the 5090: with the fix, 3 retries (1920 -> 896 -> 384 -> 128), 4 refusals, 0 ubatch failures, OK on Qwen3.5-0.8B (hybrid) and Qwen2.5-Coder-1.5B; without it, 1 retry and 1 "failed to prepare attention ubatches", FAILED, which is the production signature. Every other case is unchanged (exact vs one graph cosine 0.998 / 0.998). Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 61 +++++++++++++++++------- tests/test-walk-exact.cpp | 99 ++++++++++++++++++++++++++++++++++++--- 2 files changed, 136 insertions(+), 24 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index c55fb2b8ad09..d12d6a7637e8 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3925,26 +3925,36 @@ void llama_context::opt_epoch_iter( } // attention: pop the cache to [0, c0); recurrent: restore the state at the chunk's // checkpoint and decode forward to its start (that decode rebuilds the attention K/V - // in between too, under the same adapter: the same values) - const size_t cp = recr != nullptr ? (size_t) j - (size_t) j % stride : (size_t) j; - const uint32_t p_pop = chunks[cp].c0; - const bool popped = recr != nullptr ? (attn == nullptr || attn->seq_rm(0, p_pop, -1)) : memory->seq_rm(0, p_pop, -1); - if (!popped) { - LLAMA_LOG_ERROR("%s: could not pop the cache to [0, %u) for the reverse pass\n", __func__, p_pop); - opt_stop_requested.store(true); - ok = false; - break; - } - if (recr != nullptr) { - if (snaps[cp].empty) { - recr->seq_rm(0, -1, -1); // the first chunk starts from a zero state - } else { - state_rows(true, snaps[cp]); + // in between too, under the same adapter: the same values). Run again before every + // retry below: a refused chunk graph is refused AFTER its ubatch was applied, so the + // chunk's cells already sit in the cache and the recurrent cell's position has moved + // to the chunk's end; retrying on top of that fails to prepare the ubatch at all + // (the 5090, 2026-10-10: "init_batch: failed to prepare attention ubatches" after + // "retrying at 512", and the epoch stopped instead of shrinking the horizon). + const auto rewind = [&]() -> bool { + const size_t cp = recr != nullptr ? (size_t) j - (size_t) j % stride : (size_t) j; + const uint32_t p_pop = chunks[cp].c0; + const bool popped = recr != nullptr ? (attn == nullptr || attn->seq_rm(0, p_pop, -1)) : memory->seq_rm(0, p_pop, -1); + if (!popped) { + LLAMA_LOG_ERROR("%s: could not pop the cache to [0, %u) for the reverse pass\n", __func__, p_pop); + opt_stop_requested.store(true); + return false; } - if (p_pop < c.c0 && !decode_span(p_pop, c.c0)) { - ok = false; - break; + if (recr != nullptr) { + if (snaps[cp].empty) { + recr->seq_rm(0, -1, -1); // the first chunk starts from a zero state + } else { + state_rows(true, snaps[cp]); + } + if (p_pop < c.c0 && !decode_span(p_pop, c.c0)) { + return false; + } } + return true; + }; + if (!rewind()) { + ok = false; + break; } cparams.walk_exact = true; cparams.walk_surrogate = c.surrogate; @@ -3973,6 +3983,21 @@ void llama_context::opt_epoch_iter( __func__, cc.c0, span, opt_walk_horizon); opt_alloc_failed.store(false); opt_stop_requested.store(false); + // the rewind's decode is a plain forward, as before the first try: no walk flags + const bool surrogate = cparams.walk_surrogate; + const bool state_surrogate = cparams.walk_state_surrogate; + cparams.walk_exact = false; + cparams.walk_grad_from = 0; + cparams.walk_surrogate = false; + cparams.walk_state_surrogate = false; + const bool rewound = rewind(); + cparams.walk_exact = true; + cparams.walk_surrogate = surrogate; + cparams.walk_state_surrogate = state_surrogate; + if (!rewound) { + ok = false; + break; + } } opt_walk_horizon_used = opt_walk_horizon; cparams.walk_exact = false; diff --git a/tests/test-walk-exact.cpp b/tests/test-walk-exact.cpp index f13fe46f2231..766c9eb2781c 100644 --- a/tests/test-walk-exact.cpp +++ b/tests/test-walk-exact.cpp @@ -30,6 +30,10 @@ static const float LR = 1e-3f; // an SGD step is linear in the gradient at any lr static const uint32_t WINDOW = 512; +// the window a step trains: WINDOW for every case but the device refusal (6), whose horizon must +// change the chunk graph's size, which a window this short does not +static const uint32_t REFUSAL_WINDOW = 2048; +static uint32_t g_window = WINDOW; static ggml_opt_optimizer_params sgd_pars(void *) { ggml_opt_optimizer_params p = ggml_opt_get_default_optimizer_params(nullptr); @@ -44,15 +48,36 @@ static ggml_opt_optimizer_params sgd_pars(void *) { static ggml_type g_cache_type = GGML_TYPE_F32; // recurrent rollback slots in the context (serving keeps some for speculative decoding) static uint32_t g_n_rs_seq = 0; +// the gradient horizon a step asks for (0 = the whole window) and the device budget its graphs +// may add (0 = the device's own free figure), and what the last step measured: its largest +// graph and the horizon it ended at +static uint32_t g_walk_horizon = 0; +static size_t g_device_budget = 0; +static size_t g_last_graph_bytes = 0; +static uint32_t g_last_horizon = 0; // counts the recurrent memory's "non-consecutive token position" warnings: a state restore that // leaves the cell's position at the window's end makes every next chunk read as non-consecutive // (the 5090 log, 2026-10-07: "position 65640 after 66345"), passing everything else through static int g_nonconsecutive = 0; +// and, for the device refusal (6): the horizon retries, the ubatches that failed to prepare, and +// the chunk graphs the device budget refused +static int g_retries = 0; +static int g_ubatch_failures = 0; +static int g_refusals = 0; static void count_log(ggml_log_level level, const char * text, void * ud) { if (text && strstr(text, "non-consecutive token position")) { ++g_nonconsecutive; } + if (text && strstr(text, "did not fit with a gradient horizon")) { + ++g_retries; + } + if (text && strstr(text, "failed to prepare attention ubatches")) { + ++g_ubatch_failures; + } + if (text && strstr(text, "it may add: not allocating it")) { + ++g_refusals; + } GGML_UNUSED(level); GGML_UNUSED(ud); fputs(text, stderr); @@ -60,7 +85,7 @@ static void count_log(ggml_log_level level, const char * text, void * ud) { static llama_context * make_ctx(const common_params & params, llama_model * model, uint32_t chunk) { auto cparams = common_context_params_to_llama(params); - cparams.n_ctx = WINDOW; + cparams.n_ctx = g_window; cparams.n_batch = chunk; cparams.n_ubatch = chunk; cparams.n_seq_max = 2; // the plain walk snapshots a recurrent state into a scratch sequence @@ -110,12 +135,13 @@ static std::vector step_delta(const common_params & params, llama_model * lopt.optimizer_type = GGML_OPT_OPTIMIZER_TYPE_SGD; lopt.adapter = adapter; lopt.walk_exact = exact; - lopt.walk_horizon = 0; + lopt.walk_horizon = g_walk_horizon; + llama_opt_set_memory_budget(ctx, g_device_budget); llama_opt_init(ctx, model, lopt); llama_opt_set_walk_host_budget(ctx, host_budget); - std::vector> seqs = { std::vector(tokens.begin(), tokens.begin() + WINDOW + 1) }; + std::vector> seqs = { std::vector(tokens.begin(), tokens.begin() + g_window + 1) }; std::vector> loss = { labelled }; - ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(WINDOW, seqs, loss, tokens[0]); + ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(g_window, seqs, loss, tokens[0]); ggml_opt_result_t result = ggml_opt_result_init(); llama_opt_epoch(ctx, dataset, result, nullptr, /*idata_split =*/ 1, nullptr, nullptr); GGML_ASSERT(!llama_opt_failed(ctx)); @@ -127,6 +153,8 @@ static std::vector step_delta(const common_params & params, llama_model * if (host_bytes) { *host_bytes = llama_opt_walk_host_bytes(ctx); } + g_last_graph_bytes = llama_opt_graph_bytes(ctx); + g_last_horizon = llama_opt_walk_horizon(ctx); ggml_opt_result_free(result); ggml_opt_dataset_free(dataset); llama_free(ctx); @@ -191,7 +219,7 @@ int main(int argc, char ** argv) { "engine", "window", "chunk", "cache", "gradient", "adapter", "verdict", "submission" }; std::string text; uint32_t lcg = 12345; - while (text.size() < 8 * WINDOW) { + while (text.size() < 8 * REFUSAL_WINDOW) { lcg = lcg * 1664525u + 1013904223u; text += words[(lcg >> 16) % 16]; text += (lcg >> 8) % 7 == 0 ? ".\n" : " "; @@ -199,7 +227,7 @@ int main(int argc, char ** argv) { llama_context * tok_ctx = make_ctx(params, model, WINDOW); std::vector tokens = common_tokenize(tok_ctx, text, true); llama_free(tok_ctx); - GGML_ASSERT(tokens.size() > WINDOW + 1); + GGML_ASSERT(tokens.size() > REFUSAL_WINDOW + 1); int failures = 0; { @@ -348,6 +376,65 @@ int main(int argc, char ** argv) { llama_adapter_lora_free(adapter); } + { + // 6. regression for the 5090 (2026-10-10 04:34Z): after the device budget refuses a chunk + // graph, the walk shrinks the gradient horizon and tries the chunk again. The refusal comes + // AFTER the chunk's ubatch was applied, so the retry used to run on top of the chunk's own + // cells and fail to prepare its ubatch at all ("init_batch: failed to prepare attention + // ubatches"), stopping the epoch instead of reaching a horizon that fits. A budget just + // under the measured graph refuses EVERY horizon, so this pins the retry itself: each retry + // must reach the device preflight again (and be refused there), never fail in init_batch, + // and the epoch then refuses by name with the adapter untouched. Needs a device: a CPU run + // measures no graph, and says so. + g_window = REFUSAL_WINDOW; + std::vector all(g_window + 1, 1); + all[0] = 0; + step_delta(params, model, init, tokens, all, WINDOW / 4, true); + const size_t graph_bytes = g_last_graph_bytes; + printf(" device graph: %.1f MiB per chunk\n", graph_bytes / 1048576.0); + if (graph_bytes == 0) { + printf(" device refusal: skipped (no device graph measured: a CPU-only run)\n"); + } else { + llama_adapter_lora * adapter = llama_adapter_lora_init(model, init.c_str()); + const std::vector before = adapter_params(adapter); + llama_context * ctx = make_ctx(params, model, WINDOW / 4); + float scale = 1.0f; + GGML_ASSERT(llama_set_adapters_lora(ctx, &adapter, 1, &scale) == 0); + llama_opt_params lopt{}; + lopt.param_filter = llama_opt_param_filter_all; + lopt.get_opt_pars = sgd_pars; + lopt.optimizer_type = GGML_OPT_OPTIMIZER_TYPE_SGD; + lopt.adapter = adapter; + lopt.walk_exact = true; + // what a graph may add is the budget less the allocator's 512 MiB margin (ggml-opt), and + // "add" is growth past the buffers the walk's own forward decodes already sized: one + // byte past the margin lets a chunk graph grow by nothing, so every horizon is refused + llama_opt_set_memory_budget(ctx, 512u*1024*1024 + 1); + llama_opt_init(ctx, model, lopt); + std::vector> seqs = { std::vector(tokens.begin(), tokens.begin() + g_window + 1) }; + std::vector> loss = { all }; + ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(g_window, seqs, loss, tokens[0]); + g_retries = 0; + g_ubatch_failures = 0; + g_refusals = 0; + llama_log_set(count_log, nullptr); + llama_opt_epoch(ctx, dataset, nullptr, nullptr, /*idata_split =*/ 1, nullptr, nullptr); + llama_log_set(nullptr, nullptr); // the default logger again + const std::string why = llama_opt_failure(ctx); + printf(" device refusal: %d retries, %d graph refusals, %d ubatch failures, failed=%d \"%s\"\n", + g_retries, g_refusals, g_ubatch_failures, (int) llama_opt_failed(ctx), why.c_str()); + if (g_retries < 1 || g_refusals < 2 || g_ubatch_failures != 0 || + !llama_opt_failed(ctx) || adapter_params(adapter) != before) { + fprintf(stderr, "FAILED: a refused chunk graph was not retried cleanly at a smaller horizon\n"); + ++failures; + } + ggml_opt_dataset_free(dataset); + llama_free(ctx); + llama_adapter_lora_free(adapter); + } + g_window = WINDOW; + } + std::remove(init.c_str()); llama_model_free(model); llama_backend_free(); From 49ee77f3c8229b51dcf182a026cbbb0a680e800c Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Fri, 9 Oct 2026 23:58:19 -0500 Subject: [PATCH 2/2] docs(train): say where a horizon retry recomputes walk_grad_from Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index d12d6a7637e8..07176adb3cd4 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3983,7 +3983,10 @@ void llama_context::opt_epoch_iter( __func__, cc.c0, span, opt_walk_horizon); opt_alloc_failed.store(false); opt_stop_requested.store(false); - // the rewind's decode is a plain forward, as before the first try: no walk flags + // the rewind's decode is a plain forward, as before the first try: no walk flags. + // walk_grad_from is not restored here because the loop's first statements recompute + // it from the shrunken horizon (cc.grad_from, then cparams.walk_grad_from) before the + // retried train_chunk runs. const bool surrogate = cparams.walk_surrogate; const bool state_surrogate = cparams.walk_state_surrogate; cparams.walk_exact = false;