From 44fd0548c047c0a2fa4ce270e36244d547066ece Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 17:16:28 -0500 Subject: [PATCH 01/13] ggml: GGML_TENSOR_FLAG_GRAD, a leaf whose gradient the backward computes 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- ggml/include/ggml.h | 4 +++ ggml/src/ggml.c | 9 ++++-- tests/CMakeLists.txt | 3 ++ tests/test-grad-leaf.cpp | 64 ++++++++++++++++++++++++++++++++++++++++ 4 files changed, 78 insertions(+), 2 deletions(-) create mode 100644 tests/test-grad-leaf.cpp diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index f7dedfe22e3a..619b82b88e62 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -654,6 +654,7 @@ extern "C" { GGML_TENSOR_FLAG_PARAM = 4, // ...contains trainable parameters GGML_TENSOR_FLAG_LOSS = 8, // ...defines loss for numerical optimization (multiple loss tensors add up) GGML_TENSOR_FLAG_COMPUTE = 16, // ...must be computed + GGML_TENSOR_FLAG_GRAD = 32, // ...needs its gradient computed, but no optimizer updates it (a leaf whose gradient is read back) }; enum ggml_tri_type { @@ -883,6 +884,9 @@ extern "C" { GGML_API void ggml_set_output(struct ggml_tensor * tensor); GGML_API void ggml_set_param(struct ggml_tensor * tensor); GGML_API void ggml_set_loss(struct ggml_tensor * tensor); + // the backward computes this leaf's gradient (readable after compute), and no optimizer + // step touches it: an input whose sensitivity is needed, never a trainable parameter + GGML_API void ggml_set_grad(struct ggml_tensor * tensor); // // operations on tensors with backpropagation diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 87b9e778aeaa..2298892330a6 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -7348,7 +7348,7 @@ static size_t ggml_visit_parents_graph(struct ggml_cgraph * cgraph, struct ggml_ } } - if (node->op == GGML_OP_NONE && !(node->flags & GGML_TENSOR_FLAG_PARAM)) { + if (node->op == GGML_OP_NONE && !(node->flags & (GGML_TENSOR_FLAG_PARAM | GGML_TENSOR_FLAG_GRAD))) { // reached a leaf node, not part of the gradient graph (e.g. a constant) GGML_ASSERT(cgraph->n_leafs < cgraph->size); @@ -7446,7 +7446,7 @@ void ggml_build_backward_expand( continue; } - bool node_needs_grad = (node->flags & GGML_TENSOR_FLAG_PARAM) || (node->flags & GGML_TENSOR_FLAG_LOSS); + bool node_needs_grad = (node->flags & (GGML_TENSOR_FLAG_PARAM | GGML_TENSOR_FLAG_LOSS | GGML_TENSOR_FLAG_GRAD)) != 0; bool ignore_src[GGML_MAX_SRC] = {false}; switch (node->op) { // gradients in node->src[0] for one reason or another have no effect on output gradients @@ -8077,6 +8077,11 @@ void ggml_set_param(struct ggml_tensor * tensor) { tensor->flags |= GGML_TENSOR_FLAG_PARAM; } +void ggml_set_grad(struct ggml_tensor * tensor) { + GGML_ASSERT(tensor->op == GGML_OP_NONE); + tensor->flags |= GGML_TENSOR_FLAG_GRAD; +} + void ggml_set_loss(struct ggml_tensor * tensor) { GGML_ASSERT(ggml_is_scalar(tensor)); GGML_ASSERT(tensor->type == GGML_TYPE_F32); diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index c33052830122..797532da615a 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -150,6 +150,9 @@ endif () llama_build(test-recurrent-state-rollback.cpp) +# GGML_TENSOR_FLAG_GRAD: a leaf whose gradient the backward computes, never optimized (the walk reads the cached prefix K/V gradient) +llama_build_and_test(test-grad-leaf.cpp) + if (NOT WIN32 OR NOT BUILD_SHARED_LIBS) # these tests are disabled on Windows because they use internal functions not exported with LLAMA_API (when building with shared libraries) llama_build_and_test(test-unicode.cpp) diff --git a/tests/test-grad-leaf.cpp b/tests/test-grad-leaf.cpp new file mode 100644 index 000000000000..fa34e1b66959 --- /dev/null +++ b/tests/test-grad-leaf.cpp @@ -0,0 +1,64 @@ +// what this catches: GGML_TENSOR_FLAG_GRAD, a leaf whose gradient the backward computes while +// no optimizer updates it (the training walk reads the gradient of the cached prefix K/V this +// way). The leaf must get its exact gradient beside a real parameter's, and a leaf WITHOUT the +// flag must stay a constant with no gradient. + +#include "ggml.h" +#include "ggml-cpu.h" + +#include +#include +#include + +#define CHECK(cond) do { if (!(cond)) { fprintf(stderr, "FAILED %s:%d: %s\n", __FILE__, __LINE__, #cond); exit(1); } } while (0) + +int main() { + ggml_init_params params = { 64u * 1024 * 1024, nullptr, false }; + ggml_context * ctx = ggml_init(params); + + const int n = 4; + ggml_tensor * x = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // the GRAD leaf + ggml_tensor * w = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // a real parameter + ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // a plain constant + for (int i = 0; i < n; ++i) { + ggml_set_f32_1d(x, i, 0.5f + i); + ggml_set_f32_1d(w, i, 1.0f - 0.25f * i); + ggml_set_f32_1d(c, i, 2.0f); + } + ggml_set_grad(x); + ggml_set_param(w); + + // loss = sum((w * x + c)^2): d/dx = 2 (w x + c) w, d/dw = 2 (w x + c) x + ggml_tensor * y = ggml_add(ctx, ggml_mul(ctx, w, x), c); + ggml_tensor * loss = ggml_sum(ctx, ggml_sqr(ctx, y)); + ggml_set_loss(loss); + + ggml_cgraph * gf = ggml_new_graph_custom(ctx, GGML_DEFAULT_GRAPH_SIZE, true); + ggml_build_forward_expand(gf, loss); + ggml_cgraph * gb = ggml_graph_dup(ctx, gf, true); + ggml_build_backward_expand(ctx, gb, nullptr); + ggml_graph_reset(gb); // loss grad = 1, every other grad = 0 + ggml_graph_compute_with_ctx(ctx, gb, 1); + + ggml_tensor * gx = ggml_graph_get_grad(gb, x); + ggml_tensor * gw = ggml_graph_get_grad(gb, w); + CHECK(gx != nullptr && "a GRAD leaf gets a gradient"); + CHECK(gw != nullptr && "a PARAM still gets its gradient"); + CHECK(ggml_graph_get_grad(gb, c) == nullptr && "an unflagged constant gets none"); + for (int i = 0; i < n; ++i) { + const float xi = 0.5f + i, wi = 1.0f - 0.25f * i, yi = wi * xi + 2.0f; + CHECK(std::fabs(ggml_get_f32_1d(gx, i) - 2.0f * yi * wi) < 1e-5f); + CHECK(std::fabs(ggml_get_f32_1d(gw, i) - 2.0f * yi * xi) < 1e-5f); + } + CHECK(!(x->flags & GGML_TENSOR_FLAG_PARAM) && "the GRAD leaf is not a parameter"); + // regression: GRAD first shared COMPUTE's bit, so every computed node read as a GRAD leaf and + // the backward asked for gradients of the quantized weights (GGML_ASSERT in the 1.5B walk) + for (int i = 0; i < ggml_graph_n_nodes(gb); ++i) { + ggml_tensor * node = ggml_graph_node(gb, i); + CHECK((node == x || !(node->flags & GGML_TENSOR_FLAG_GRAD)) && "only the flagged leaf is a GRAD leaf"); + } + + ggml_free(ctx); + printf("test-grad-leaf: OK\n"); + return 0; +} From f246c4efd0472ab0a7822b171341d122783edc4f Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 17:36:50 -0500 Subject: [PATCH 02/13] ggml-opt: graphs built per step start every optimizer period from zero 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- ggml/src/ggml-opt.cpp | 13 +++++ tests/CMakeLists.txt | 3 ++ tests/test-opt-dynamic-accum.cpp | 92 ++++++++++++++++++++++++++++++++ 3 files changed, 108 insertions(+) create mode 100644 tests/test-opt-dynamic-accum.cpp diff --git a/ggml/src/ggml-opt.cpp b/ggml/src/ggml-opt.cpp index 5220ec5f65b4..8814325e9720 100644 --- a/ggml/src/ggml-opt.cpp +++ b/ggml/src/ggml-opt.cpp @@ -906,6 +906,19 @@ bool ggml_opt_alloc(ggml_opt_context_t opt_ctx, bool backward) { if (!opt_ctx->static_graphs) { ggml_opt_build(opt_ctx); + + // Graphs built per step keep their gradient accumulators in ctx_static across graphs, + // and the reset above found no graph to reset (gb_grad is rebuilt every step). Without + // this, a period's step applied the SUM of every gradient since the run began, since + // each backward adds into the accumulator in place (test-opt-dynamic-accum). + if (backward && opt_ctx->opt_i == 0) { + for (size_t i = 0; i < opt_ctx->grad_accs.size(); ++i) { + ggml_tensor * acc = opt_ctx->grad_accs[i]; + if (acc && (opt_ctx->gf->nodes[i]->flags & GGML_TENSOR_FLAG_PARAM)) { + ggml_set_zero(acc); + } + } + } } struct ggml_cgraph * graph = nullptr; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 797532da615a..cbb9f6481699 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -153,6 +153,9 @@ llama_build(test-recurrent-state-rollback.cpp) # GGML_TENSOR_FLAG_GRAD: a leaf whose gradient the backward computes, never optimized (the walk reads the cached prefix K/V gradient) llama_build_and_test(test-grad-leaf.cpp) +# ggml-opt with per-step graphs (the llama training path) starts each optimizer period from zero accumulators +llama_build_and_test(test-opt-dynamic-accum.cpp) + if (NOT WIN32 OR NOT BUILD_SHARED_LIBS) # these tests are disabled on Windows because they use internal functions not exported with LLAMA_API (when building with shared libraries) llama_build_and_test(test-unicode.cpp) diff --git a/tests/test-opt-dynamic-accum.cpp b/tests/test-opt-dynamic-accum.cpp new file mode 100644 index 000000000000..f282b3d7c51f --- /dev/null +++ b/tests/test-opt-dynamic-accum.cpp @@ -0,0 +1,92 @@ +// what this catches: ggml-opt with graphs built per step (ggml_opt_prepare_alloc, the path +// llama_context's training takes) must start every optimizer period from ZERO gradient +// accumulators. Each step here is SGD on loss = sum(w * x), so its gradient is x and each +// step moves w by exactly -lr * x. If the accumulators carried the previous period's gradient, +// the second step would move w by -2 lr x, the third by -3 lr x. + +#include "ggml.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml-cpu.h" +#include "ggml-opt.h" + +#include +#include +#include + +#define CHECK(cond) do { if (!(cond)) { fprintf(stderr, "FAILED %s:%d: %s\n", __FILE__, __LINE__, #cond); exit(1); } } while (0) + +static const float LR = 0.125f; +static const float X = 2.0f; + +static ggml_opt_optimizer_params sgd_pars(void *) { + ggml_opt_optimizer_params p = ggml_opt_get_default_optimizer_params(nullptr); + p.sgd.alpha = LR; + p.sgd.wd = 0.0f; + return p; +} + +static void run(int32_t opt_period) { + ggml_backend_t cpu = ggml_backend_cpu_init(); + ggml_backend_t backends[] = { cpu }; + ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 1, GGML_DEFAULT_GRAPH_SIZE, false, true); + + ggml_init_params sp = { 8 * ggml_tensor_overhead(), nullptr, true }; + ggml_context * ctx_static = ggml_init(sp); + ggml_tensor * w = ggml_new_tensor_1d(ctx_static, GGML_TYPE_F32, 1); + ggml_set_param(w); + ggml_backend_buffer_t buf = ggml_backend_alloc_ctx_tensors(ctx_static, cpu); + float w_now = 1.0f; + ggml_backend_tensor_set(w, &w_now, 0, sizeof(float)); + + ggml_opt_params params = ggml_opt_default_params(sched, GGML_OPT_LOSS_TYPE_SUM); + params.optimizer = GGML_OPT_OPTIMIZER_TYPE_SGD; + params.get_opt_pars = sgd_pars; + params.opt_period = opt_period; + ggml_opt_context_t opt_ctx = ggml_opt_init(params); // no ctx_compute: graphs are built per step + ggml_opt_result_t result = ggml_opt_result_init(); + + const int n_evals = 3 * opt_period; + for (int i = 0; i < n_evals; ++i) { + ggml_init_params cp = { GGML_DEFAULT_GRAPH_SIZE * ggml_tensor_overhead() + 4 * ggml_graph_overhead_custom(GGML_DEFAULT_GRAPH_SIZE, true), nullptr, true }; + ggml_context * ctx_compute = ggml_init(cp); + ggml_tensor * x = ggml_new_tensor_1d(ctx_compute, GGML_TYPE_F32, 1); + ggml_tensor * out = ggml_mul(ctx_compute, w, x); + ggml_cgraph * gf = ggml_new_graph_custom(ctx_compute, GGML_DEFAULT_GRAPH_SIZE, true); + ggml_build_forward_expand(gf, out); + ggml_opt_prepare_alloc(opt_ctx, ctx_compute, gf, x, out); + CHECK(ggml_opt_alloc(opt_ctx, true)); + ggml_backend_tensor_set(x, &X, 0, sizeof(float)); + ggml_opt_eval(opt_ctx, result); + ggml_free(ctx_compute); + + float w_after; + ggml_backend_tensor_get(w, &w_after, 0, sizeof(float)); + if ((i + 1) % opt_period == 0) { + // one period = opt_period evals of gradient X each, scaled by nothing (LOSS_TYPE_SUM) + const float expected = w_now - LR * X * opt_period; + if (std::fabs(w_after - expected) > 1e-6f) { + fprintf(stderr, "opt_period %d, step %d: w moved to %f, expected %f (from %f)\n", + opt_period, (i + 1) / opt_period, w_after, expected, w_now); + exit(1); + } + w_now = w_after; + } else { + CHECK(std::fabs(w_after - w_now) < 1e-7f && "no update inside a period"); + } + } + + ggml_opt_result_free(result); + ggml_opt_free(opt_ctx); + ggml_backend_buffer_free(buf); + ggml_free(ctx_static); + ggml_backend_sched_free(sched); + ggml_backend_free(cpu); +} + +int main() { + run(1); + run(2); + printf("test-opt-dynamic-accum: OK\n"); + return 0; +} From 0863877da31f4819785b36fb6c4e64a6f6464735 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 17:51:33 -0500 Subject: [PATCH 03/13] ggml-opt: accumulators and momenta belong to a parameter, never to a 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- ggml/src/ggml-opt.cpp | 72 ++++++++++++++++++++++++++----------------- 1 file changed, 43 insertions(+), 29 deletions(-) diff --git a/ggml/src/ggml-opt.cpp b/ggml/src/ggml-opt.cpp index 8814325e9720..49c6aa783785 100644 --- a/ggml/src/ggml-opt.cpp +++ b/ggml/src/ggml-opt.cpp @@ -67,9 +67,17 @@ struct ggml_opt_context { size_t alloc_budget = 0; // bytes a graph may add on a non-CPU device; 0 = the device's own free figure size_t peak_graph_bytes = 0; // largest graph measured by the preflight on a non-CPU device bool eval_ready = false; + // Accumulators and momenta belong to a PARAMETER (or to the loss), never to a node index: a + // graph built per step may differ in topology from the first one (the training walk's chunks + // with and without a prefix, its reverse pass with GRAD leaves and a surrogate term), and an + // index would bind one tensor's accumulator to whatever node took that index. grad_accs is + // the current graph's view of them, rebuilt by every ggml_opt_build. + std::unordered_map param_grad_acc; + std::unordered_map param_m; + std::unordered_map param_v; + struct ggml_tensor * loss_grad_acc = nullptr; + bool accumulators_created = false; std::vector grad_accs; - std::vector grad_m; - std::vector grad_v; int64_t iter = 1; int32_t opt_period = 1; @@ -606,32 +614,41 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) { return; } - if (opt_ctx->grad_accs.empty()) { + if (!opt_ctx->accumulators_created) { GGML_ASSERT(opt_ctx->build_type_alloc >= GGML_OPT_BUILD_TYPE_GRAD); + opt_ctx->accumulators_created = true; - const int n_nodes = opt_ctx->gf->n_nodes; - opt_ctx->grad_accs.resize(n_nodes); - for (int i = 0; i < n_nodes; ++i) { + // created once, in ctx_static, for the parameters of the first graph: every later graph + // trains the same parameters (asserted below), whatever else its topology holds + for (int i = 0; i < opt_ctx->gf->n_nodes; ++i) { ggml_tensor * node = opt_ctx->gf->nodes[i]; - if ((accumulate && (node->flags & GGML_TENSOR_FLAG_PARAM)) || (node->flags & GGML_TENSOR_FLAG_LOSS)) { - opt_ctx->grad_accs[i] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne); - } else { - opt_ctx->grad_accs[i] = nullptr; + if (node->flags & GGML_TENSOR_FLAG_LOSS) { + opt_ctx->loss_grad_acc = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne); + } + if (!(node->flags & GGML_TENSOR_FLAG_PARAM)) { + continue; + } + if (accumulate) { + opt_ctx->param_grad_acc[node] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne); + } + if (need_momenta && opt_ctx->build_type_alloc >= GGML_OPT_BUILD_TYPE_OPT) { + opt_ctx->param_m[node] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne); + opt_ctx->param_v[node] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne); } } + } - if (need_momenta && opt_ctx->build_type_alloc >= GGML_OPT_BUILD_TYPE_OPT) { - opt_ctx->grad_m.resize(n_nodes); - opt_ctx->grad_v.resize(n_nodes); - for (int i = 0; i < n_nodes; ++i) { - ggml_tensor * node = opt_ctx->gf->nodes[i]; - if (node->flags & GGML_TENSOR_FLAG_PARAM) { - opt_ctx->grad_m[i] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne); - opt_ctx->grad_v[i] = ggml_new_tensor(opt_ctx->ctx_static, GGML_TYPE_F32, GGML_MAX_DIMS, node->ne); - } else { - opt_ctx->grad_m[i] = nullptr; - opt_ctx->grad_v[i] = nullptr; - } + // this graph's view of the accumulators, by tensor + opt_ctx->grad_accs.assign(opt_ctx->gf->n_nodes, nullptr); + for (int i = 0; i < opt_ctx->gf->n_nodes; ++i) { + ggml_tensor * node = opt_ctx->gf->nodes[i]; + if (node->flags & GGML_TENSOR_FLAG_LOSS) { + opt_ctx->grad_accs[i] = opt_ctx->loss_grad_acc; + } else if (node->flags & GGML_TENSOR_FLAG_PARAM) { + GGML_ASSERT((!accumulate || opt_ctx->param_grad_acc.count(node)) && + "a graph trains a parameter the optimizer context was not built with"); + if (accumulate) { + opt_ctx->grad_accs[i] = opt_ctx->param_grad_acc.at(node); } } } @@ -670,8 +687,8 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) { struct ggml_tensor * m = nullptr; struct ggml_tensor * v = nullptr; if (need_momenta) { - m = opt_ctx->grad_m[i]; - v = opt_ctx->grad_v[i]; + m = opt_ctx->param_m.at(node); + v = opt_ctx->param_v.at(node); ggml_format_name(m, "AdamW m for %s", node->name); ggml_format_name(v, "AdamW v for %s", node->name); } @@ -912,11 +929,8 @@ bool ggml_opt_alloc(ggml_opt_context_t opt_ctx, bool backward) { // this, a period's step applied the SUM of every gradient since the run began, since // each backward adds into the accumulator in place (test-opt-dynamic-accum). if (backward && opt_ctx->opt_i == 0) { - for (size_t i = 0; i < opt_ctx->grad_accs.size(); ++i) { - ggml_tensor * acc = opt_ctx->grad_accs[i]; - if (acc && (opt_ctx->gf->nodes[i]->flags & GGML_TENSOR_FLAG_PARAM)) { - ggml_set_zero(acc); - } + for (auto & [param, acc] : opt_ctx->param_grad_acc) { + ggml_set_zero(acc); } } } From 3075a3571feb3f0e3705dc0b4be081b3a9008738 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 18:02:10 -0500 Subject: [PATCH 04/13] ggml-opt: a caller-driven optimizer period, a weighted loss with an extra 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- ggml/include/ggml-opt.h | 19 +++++++ ggml/src/ggml-opt.cpp | 72 +++++++++++++++++++++++-- tests/test-opt-dynamic-accum.cpp | 93 ++++++++++++++++++++++++++++++++ 3 files changed, 179 insertions(+), 5 deletions(-) diff --git a/ggml/include/ggml-opt.h b/ggml/include/ggml-opt.h index 9dc27b7bf9be..c04ffba47f67 100644 --- a/ggml/include/ggml-opt.h +++ b/ggml/include/ggml-opt.h @@ -211,6 +211,25 @@ extern "C" { // do forward pass, increment result if not NULL, do backward pass if allocated GGML_API void ggml_opt_eval(ggml_opt_context_t opt_ctx, ggml_opt_result_t result); + // A caller-driven optimizer period, for graphs built per step (the training walk's reverse + // pass: one graph per chunk, ONE optimizer step per window). Applies to the next + // ggml_opt_alloc + ggml_opt_eval only: + // period_end: true = this graph runs the optimizer step, false = it only accumulates + // loss_scale: the weight of this graph's loss in the period's total (replaces 1/opt_period) + // extra_loss: NULL, or a scalar F32 node of the forward graph added to the loss the backward + // differentiates (a surrogate carrying a later graph's gradient into this one); + // the loss the result reports stays the unweighted loss of the outputs + // A period starts from zero gradients after the step that ended the previous one. + GGML_API void ggml_opt_set_next_step( + ggml_opt_context_t opt_ctx, + bool period_end, + float loss_scale, + struct ggml_tensor * extra_loss); + + // After ggml_opt_eval and until the next ggml_opt_alloc: the gradient the backward computed + // for a GRAD leaf (ggml_set_grad) of the evaluated graph, or NULL when the graph has none. + GGML_API struct ggml_tensor * ggml_opt_leaf_grad(ggml_opt_context_t opt_ctx, struct ggml_tensor * leaf); + // ############################################################################ // ## The high-level functions start here. They do not depend on any private ## // ## functions or structs and can be copied to and adapted for user code. ## diff --git a/ggml/src/ggml-opt.cpp b/ggml/src/ggml-opt.cpp index 49c6aa783785..6379f996f4f3 100644 --- a/ggml/src/ggml-opt.cpp +++ b/ggml/src/ggml-opt.cpp @@ -77,6 +77,15 @@ struct ggml_opt_context { std::unordered_map param_v; struct ggml_tensor * loss_grad_acc = nullptr; bool accumulators_created = false; + + // ggml_opt_set_next_step: a caller-driven period, for the next graph only + bool next_manual = false; + bool next_period_end = false; + float next_loss_scale = 1.0f; + struct ggml_tensor * next_extra_loss = nullptr; + bool period_fresh = true; // the next backward starts a period (zeroed gradients) + // the gradients of the evaluated graph's GRAD leaves, readable until the next alloc + std::unordered_map leaf_grads; std::vector grad_accs; int64_t iter = 1; @@ -564,7 +573,9 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) { ggml_set_name(opt_ctx->labels, "labels"); opt_ctx->loss = ggml_cross_entropy_loss(ctx_results, opt_ctx->outputs, opt_ctx->labels); ggml_set_name(opt_ctx->loss, "loss_cross_entropy"); - if (opt_ctx->opt_period > 1) { + if (opt_ctx->next_manual) { + // weighted by the caller below, beside its extra term + } else if (opt_ctx->opt_period > 1) { opt_ctx->loss = ggml_scale(ctx_results, opt_ctx->loss, 1.0f / opt_ctx->opt_period); ggml_set_name(opt_ctx->loss, "loss_cross_entropy_scaled"); } @@ -589,8 +600,24 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) { } } ggml_set_output(opt_ctx->loss); - ggml_set_loss(opt_ctx->loss); - ggml_build_forward_expand(opt_ctx->gf, opt_ctx->loss); + if (opt_ctx->next_manual) { + GGML_ASSERT(!opt_ctx->static_graphs && "a caller-driven period needs graphs built per step"); + GGML_ASSERT(opt_ctx->loss_type == GGML_OPT_LOSS_TYPE_CROSS_ENTROPY); + // the loss the backward differentiates: the caller's weight of this graph's loss, plus + // its extra term; opt_ctx->loss stays the unweighted loss the result reports + struct ggml_tensor * total = ggml_scale(ctx_results, opt_ctx->loss, opt_ctx->next_loss_scale); + if (opt_ctx->next_extra_loss) { + GGML_ASSERT(ggml_is_scalar(opt_ctx->next_extra_loss) && opt_ctx->next_extra_loss->type == GGML_TYPE_F32); + total = ggml_add(ctx_results, total, opt_ctx->next_extra_loss); + } + ggml_set_name(total, "loss_total"); + ggml_set_loss(total); + ggml_build_forward_expand(opt_ctx->gf, opt_ctx->loss); + ggml_build_forward_expand(opt_ctx->gf, total); + } else { + ggml_set_loss(opt_ctx->loss); + ggml_build_forward_expand(opt_ctx->gf, opt_ctx->loss); + } if (opt_ctx->loss_type == GGML_OPT_LOSS_TYPE_CROSS_ENTROPY) { opt_ctx->pred = ggml_argmax(ctx_results, opt_ctx->outputs); @@ -656,6 +683,17 @@ static void ggml_opt_build(ggml_opt_context_t opt_ctx) { // gb_grad == graph backward gradients, forward pass, then backward pass to calculate gradients. opt_ctx->gb_grad = ggml_graph_dup(opt_ctx->ctx_compute, opt_ctx->gf, /*force_grads =*/ true); ggml_build_backward_expand(opt_ctx->ctx_compute, opt_ctx->gb_grad, opt_ctx->grad_accs.data()); + opt_ctx->leaf_grads.clear(); + for (int i = 0; i < opt_ctx->gf->n_nodes; ++i) { + ggml_tensor * node = opt_ctx->gf->nodes[i]; + if (node->flags & GGML_TENSOR_FLAG_GRAD) { + ggml_tensor * grad = ggml_graph_get_grad(opt_ctx->gb_grad, node); + if (grad) { + ggml_set_output(grad); // read back after the step: the allocator must not reuse it + opt_ctx->leaf_grads[node] = grad; + } + } + } if (!opt_ctx->checkpoint_prefix.empty()) { ggml_opt_checkpoint(opt_ctx->ctx_compute, opt_ctx->gb_grad, opt_ctx->gf->n_nodes, opt_ctx->checkpoint_prefix.c_str()); } @@ -765,6 +803,19 @@ void ggml_opt_free(ggml_opt_context_t opt_ctx) { delete opt_ctx; } +void ggml_opt_set_next_step(ggml_opt_context_t opt_ctx, bool period_end, float loss_scale, struct ggml_tensor * extra_loss) { + GGML_ASSERT(!opt_ctx->eval_ready && "set the next step before ggml_opt_alloc"); + opt_ctx->next_manual = true; + opt_ctx->next_period_end = period_end; + opt_ctx->next_loss_scale = loss_scale; + opt_ctx->next_extra_loss = extra_loss; +} + +struct ggml_tensor * ggml_opt_leaf_grad(ggml_opt_context_t opt_ctx, struct ggml_tensor * leaf) { + const auto it = opt_ctx->leaf_grads.find(leaf); + return it == opt_ctx->leaf_grads.end() ? nullptr : it->second; +} + void ggml_opt_reset(ggml_opt_context_t opt_ctx, bool optimizer) { if (optimizer) { ggml_graph_reset(opt_ctx->gb_opt); @@ -914,7 +965,9 @@ bool ggml_opt_alloc(ggml_opt_context_t opt_ctx, bool backward) { if (opt_ctx->build_type == GGML_OPT_BUILD_TYPE_OPT && opt_ctx->opt_period > 1 && opt_ctx->opt_i == 0) { ggml_graph_reset(opt_ctx->gb_grad); } - if (backward) { + if (backward && opt_ctx->next_manual) { + opt_ctx->build_type = opt_ctx->next_period_end ? GGML_OPT_BUILD_TYPE_OPT : GGML_OPT_BUILD_TYPE_GRAD; + } else if (backward) { const int32_t opt_i_next = (opt_ctx->opt_i + 1) % opt_ctx->opt_period; opt_ctx->build_type = opt_i_next == 0 ? GGML_OPT_BUILD_TYPE_OPT : GGML_OPT_BUILD_TYPE_GRAD; } else { @@ -928,7 +981,7 @@ bool ggml_opt_alloc(ggml_opt_context_t opt_ctx, bool backward) { // and the reset above found no graph to reset (gb_grad is rebuilt every step). Without // this, a period's step applied the SUM of every gradient since the run began, since // each backward adds into the accumulator in place (test-opt-dynamic-accum). - if (backward && opt_ctx->opt_i == 0) { + if (backward && (opt_ctx->next_manual ? opt_ctx->period_fresh : opt_ctx->opt_i == 0)) { for (auto & [param, acc] : opt_ctx->param_grad_acc) { ggml_set_zero(acc); } @@ -1127,6 +1180,15 @@ void ggml_opt_eval(ggml_opt_context_t opt_ctx, ggml_opt_result_t result) { } opt_ctx->iter += opt_ctx->allocated_graph == opt_ctx->gb_opt; opt_ctx->opt_i = (opt_ctx->opt_i + 1) % opt_ctx->opt_period; + if (opt_ctx->allocated_graph == opt_ctx->gb_opt) { + opt_ctx->period_fresh = true; + } else if (opt_ctx->allocated_graph == opt_ctx->gb_grad) { + opt_ctx->period_fresh = false; + } + opt_ctx->next_manual = false; + opt_ctx->next_period_end = false; + opt_ctx->next_loss_scale = 1.0f; + opt_ctx->next_extra_loss = nullptr; if (!opt_ctx->static_graphs) { opt_ctx->gf = nullptr; diff --git a/tests/test-opt-dynamic-accum.cpp b/tests/test-opt-dynamic-accum.cpp index f282b3d7c51f..191c9e78f6ce 100644 --- a/tests/test-opt-dynamic-accum.cpp +++ b/tests/test-opt-dynamic-accum.cpp @@ -84,9 +84,102 @@ static void run(int32_t opt_period) { ggml_backend_free(cpu); } +// what this catches: ggml_opt_set_next_step, the training walk's one-step-per-window period. +// Three graphs on loss = sum(w * x): the first two only accumulate, at weights 0.5 and 0.25, +// the second with an extra term 3 * w * x; the third (weight 1) takes the step. So the step +// is SGD on (0.5 + 0.25 + 3 + 1) * x. Each graph's x is a GRAD leaf, whose gradient (the +// weight times w, plus 3 w in the second) must read back after its eval; a following period +// starts from zero. +static void run_manual() { + ggml_backend_t cpu = ggml_backend_cpu_init(); + ggml_backend_t backends[] = { cpu }; + ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 1, GGML_DEFAULT_GRAPH_SIZE, false, true); + + ggml_init_params sp = { 8 * ggml_tensor_overhead(), nullptr, true }; + ggml_context * ctx_static = ggml_init(sp); + ggml_tensor * w = ggml_new_tensor_1d(ctx_static, GGML_TYPE_F32, 1); + ggml_set_param(w); + ggml_backend_buffer_t buf = ggml_backend_alloc_ctx_tensors(ctx_static, cpu); + float w_now = 1.0f; + ggml_backend_tensor_set(w, &w_now, 0, sizeof(float)); + + ggml_opt_params params = ggml_opt_default_params(sched, GGML_OPT_LOSS_TYPE_CROSS_ENTROPY); + params.optimizer = GGML_OPT_OPTIMIZER_TYPE_SGD; + params.get_opt_pars = sgd_pars; + ggml_opt_context_t opt_ctx = ggml_opt_init(params); + + // cross-entropy over 2 classes of logits [w x, 0] with label class 0: + // dCE/dlogit0 = softmax0 - 1, so d/dw = (s0 - 1) x and d/dx = (s0 - 1) w + const float scales[3] = { 0.5f, 0.25f, 1.0f }; + for (int period = 0; period < 2; ++period) { + float expected_grad_w = 0.0f; + for (int k = 0; k < 3; ++k) { + ggml_init_params cp = { GGML_DEFAULT_GRAPH_SIZE * ggml_tensor_overhead() + 4 * ggml_graph_overhead_custom(GGML_DEFAULT_GRAPH_SIZE, true), nullptr, true }; + ggml_context * ctx_compute = ggml_init(cp); + ggml_tensor * x = ggml_new_tensor_1d(ctx_compute, GGML_TYPE_F32, 1); + ggml_set_input(x); + ggml_set_grad(x); + ggml_tensor * z = ggml_new_tensor_2d(ctx_compute, GGML_TYPE_F32, 1, 1); + ggml_set_input(z); + ggml_tensor * wx = ggml_reshape_2d(ctx_compute, ggml_mul(ctx_compute, w, x), 1, 1); + ggml_tensor * logits = ggml_concat(ctx_compute, wx, z, 0); // [2, 1] + ggml_tensor * extra = k == 1 ? ggml_scale(ctx_compute, ggml_sum(ctx_compute, ggml_mul(ctx_compute, w, x)), 3.0f) : nullptr; + ggml_cgraph * gf = ggml_new_graph_custom(ctx_compute, GGML_DEFAULT_GRAPH_SIZE, true); + ggml_build_forward_expand(gf, logits); + if (extra) { + ggml_build_forward_expand(gf, extra); + } + ggml_opt_prepare_alloc(opt_ctx, ctx_compute, gf, x, logits); + ggml_opt_set_next_step(opt_ctx, /*period_end =*/ k == 2, scales[k], extra); + CHECK(ggml_opt_alloc(opt_ctx, true)); + ggml_backend_tensor_set(x, &X, 0, sizeof(float)); + const float zero = 0.0f, one = 1.0f; + ggml_backend_tensor_set(z, &zero, 0, sizeof(float)); + ggml_tensor * labels = ggml_opt_labels(opt_ctx); + ggml_backend_tensor_set(labels, &one, 0, sizeof(float)); + ggml_backend_tensor_set(labels, &zero, sizeof(float), sizeof(float)); + ggml_opt_eval(opt_ctx, nullptr); + + const float s0 = 1.0f / (1.0f + std::exp(-w_now * X)); + const float dlogit = scales[k] * (s0 - 1.0f); + const float dx = dlogit * w_now + (k == 1 ? 3.0f * w_now : 0.0f); + expected_grad_w += dlogit * X + (k == 1 ? 3.0f * X : 0.0f); + ggml_tensor * gx = ggml_opt_leaf_grad(opt_ctx, x); + CHECK(gx != nullptr && "a GRAD leaf's gradient reads back after the eval"); + float gx_val; + ggml_backend_tensor_get(gx, &gx_val, 0, sizeof(float)); + if (std::fabs(gx_val - dx) > 1e-5f) { + fprintf(stderr, "period %d graph %d: dL/dx %f, expected %f\n", period, k, gx_val, dx); + exit(1); + } + ggml_free(ctx_compute); + + float w_after; + ggml_backend_tensor_get(w, &w_after, 0, sizeof(float)); + if (k < 2) { + CHECK(std::fabs(w_after - w_now) < 1e-7f && "no step before the period ends"); + } else { + const float expected = w_now - LR * expected_grad_w; + if (std::fabs(w_after - expected) > 1e-5f) { + fprintf(stderr, "period %d: w moved to %f, expected %f (from %f)\n", period, w_after, expected, w_now); + exit(1); + } + w_now = w_after; + } + } + } + + ggml_opt_free(opt_ctx); + ggml_backend_buffer_free(buf); + ggml_free(ctx_static); + ggml_backend_sched_free(sched); + ggml_backend_free(cpu); +} + int main() { run(1); run(2); + run_manual(); printf("test-opt-dynamic-accum: OK\n"); return 0; } From 8f0512fd87baf768ccda259185f18b912d066093 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 18:55:04 -0500 Subject: [PATCH 05/13] ggml: CONCAT backward slices a contiguous gradient (a transposed one 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- ggml/src/ggml.c | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 2298892330a6..8f64ff6399bb 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -7115,18 +7115,23 @@ static void ggml_compute_backward( // layer concatenates the cached conv state (a leaf) with this ubatch's qkv, so the // second slice is the one that carries a gradient back into the model. const int dim = ggml_get_op_params_i32(tensor, 0); + // A view's first stride is the element size (ggml_view_4d takes nb[1..3] only), so the + // gradient is sliced from a CONTIGUOUS copy: one that arrives through a transpose (the + // training attention transposes V for its KQV product) has a first stride that is not, + // and slicing it in place read the right number of elements from the wrong places. + struct ggml_tensor * g = ggml_is_contiguous(grad) ? grad : ggml_cont(ctx, grad); size_t offset = 0; for (int j = 0; j < 2; ++j) { struct ggml_tensor * src = tensor->src[j]; const size_t isrc = j == 0 ? isrc0 : isrc1; const bool needs = j == 0 ? src0_needs_grads : src1_needs_grads; if (needs) { - struct ggml_tensor * slice = ggml_view_4d(ctx, grad, + struct ggml_tensor * slice = ggml_view_4d(ctx, g, src->ne[0], src->ne[1], src->ne[2], src->ne[3], - grad->nb[1], grad->nb[2], grad->nb[3], offset); + g->nb[1], g->nb[2], g->nb[3], offset); ggml_add_or_set(ctx, cgraph, isrc, ggml_cont(ctx, slice)); } - offset += src->ne[dim] * grad->nb[dim]; + offset += src->ne[dim] * g->nb[dim]; } } break; case GGML_OP_SSM_CONV: { From 366d44cdba8d3e5869e42e15bdeee35d8ee4ecb6 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 18:55:05 -0500 Subject: [PATCH 06/13] train: the exact walk, one optimizer step per window on the gradient 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 + 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- include/llama.h | 8 ++ src/llama-context.cpp | 166 +++++++++++++++++++++++++- src/llama-context.h | 2 + src/llama-cparams.h | 9 ++ src/llama-graph.cpp | 93 ++++++++++++--- src/llama-graph.h | 14 +++ tests/CMakeLists.txt | 3 + tests/test-walk-exact.cpp | 213 ++++++++++++++++++++++++++++++++++ tools/server/server-train.cpp | 10 +- 9 files changed, 494 insertions(+), 24 deletions(-) create mode 100644 tests/test-walk-exact.cpp diff --git a/include/llama.h b/include/llama.h index 231d00a6b5e2..1a2679a090a7 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1714,6 +1714,14 @@ extern "C" { // per-layer residual, "l_out") instead of keeping every layer's alive: about one extra // forward, and one layer's attention scores in memory instead of all of them. bool recompute; + + // THE EXACT WALK (pure attention models): one optimizer step per training window, with + // each chunk's loss carried backward through the cached K/V of every chunk before it, so + // the step is the gradient of the whole window's loss, not of each chunk alone. false keeps + // the walk's stop-gradient at each chunk boundary. walk_horizon: how many cached positions + // before a chunk receive its gradient (0 = the whole window). + bool walk_exact; + uint32_t walk_horizon; }; LLAMA_API void llama_opt_init(struct llama_context * lctx, struct llama_model * model, struct llama_opt_params lopt_params); diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 80eaeb1336df..cd7f98a65758 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3335,6 +3335,8 @@ void llama_context::opt_init(struct llama_model * model, struct llama_opt_params // The training window lives on this context, not the shared model (a serving context may // be running on the same weights). opt_n_ctx_train = lopt_params.n_ctx_train > 0 ? lopt_params.n_ctx_train : n_ctx(); + opt_walk_exact = lopt_params.walk_exact; + opt_walk_horizon = lopt_params.walk_horizon; const uint32_t n_batch = std::min(this->n_batch(), opt_n_ctx_train); const uint32_t n_ubatch = std::min(this->n_ubatch(), n_batch); GGML_ASSERT(opt_n_ctx_train % n_batch == 0); @@ -3495,6 +3497,22 @@ void llama_context::opt_epoch_iter( return ok; }; + // THE EXACT WALK's state for the chunk being trained (null in the plain walk): its place in + // the window's one optimizer period, and the gradient accumulated on the window's cached K/V + struct exact_chunk { + uint32_t c0; + uint32_t n; + uint32_t n_labels; + uint32_t grad_from; // its GRAD leaves cover cached positions [grad_from, c0) + bool surrogate; // a later chunk's gradient sits on its own K/V + bool needed; + bool period_end; // the window's optimizer step + float loss_scale; // its labelled positions over the window's + }; + const exact_chunk * xc = nullptr; + // per layer, per position: dL/d(cached K) and dL/d(cached V), [position][n_embd_*_gqa] + std::vector> walk_gk, walk_gv; + // 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 { @@ -3504,7 +3522,10 @@ void llama_context::opt_epoch_iter( batch.pos [pos_batch] = pos_ctx + pos_batch; batch.n_seq_id[pos_batch] = 1; batch.seq_id [pos_batch][0] = 0; - batch.logits [pos_batch] = labels_sparse[pos_ctx + pos_batch] >= 0; + // an exact walk's context chunk has no label but still trains (its K/V carry later + // chunks' gradient): one output row, weighted zero, so the graph has its loss + batch.logits [pos_batch] = labels_sparse[pos_ctx + pos_batch] >= 0 + || (xc != nullptr && xc->n_labels == 0 && pos_batch + 1 == n_chunk); } // output_all = false: the batch's own logits flags decide which rows exist @@ -3577,6 +3598,9 @@ void llama_context::opt_epoch_iter( ctx_compute_opt = ggml_init(params); } ggml_opt_prepare_alloc(opt_ctx, ctx_compute_opt, gf, res->get_inp_tokens(), res->get_logits()); + if (xc != nullptr) { + ggml_opt_set_next_step(opt_ctx, xc->period_end, xc->loss_scale, res->t_walk_surrogate); + } if (!ggml_opt_alloc(opt_ctx, train)) { const char * why = ggml_opt_refusal(opt_ctx); opt_failure = why; @@ -3603,13 +3627,32 @@ void llama_context::opt_epoch_iter( continue; } const uint32_t ilabel = pos_ctx + pos_batch + pos_ubatch; - GGML_ASSERT(labels_sparse[ilabel] >= 0 && labels_sparse[ilabel] < labels->ne[0]); - ggml_backend_tensor_set(labels, &onef, (row*labels->ne[0] + labels_sparse[ilabel])*sizeof(float), sizeof(float)); + // an exact walk's context chunk: its one row has no label (weighted zero) + GGML_ASSERT((xc != nullptr && labels_sparse[ilabel] < 0) || (labels_sparse[ilabel] >= 0 && labels_sparse[ilabel] < labels->ne[0])); + if (labels_sparse[ilabel] >= 0) { + ggml_backend_tensor_set(labels, &onef, (row*labels->ne[0] + labels_sparse[ilabel])*sizeof(float), sizeof(float)); + } ++row; } GGML_ASSERT(row == n_outputs); } - ggml_opt_eval(opt_ctx, result); + if (xc != nullptr) { + // the GRAD leaves start at zero; the surrogate inputs are what later chunks + // accumulated on this chunk's own K/V + for (const auto & io : res->t_walk) { + if (io.dk) { + ggml_backend_tensor_memset(io.dk, 0, 0, ggml_nbytes(io.dk)); + ggml_backend_tensor_memset(io.dv, 0, 0, ggml_nbytes(io.dv)); + } + if (io.gk) { + const size_t kd = io.gk->ne[0]*io.gk->ne[1]; + const size_t vd = io.gv->ne[0]*io.gv->ne[1]; + ggml_backend_tensor_set(io.gk, walk_gk[io.il].data() + (size_t) pos_ctx*kd, 0, ggml_nbytes(io.gk)); + ggml_backend_tensor_set(io.gv, walk_gv[io.il].data() + (size_t) pos_ctx*vd, 0, ggml_nbytes(io.gv)); + } + } + } + ggml_opt_eval(opt_ctx, xc != nullptr && xc->n_labels == 0 ? nullptr : result); if (const char * why = ggml_opt_refusal(opt_ctx); why[0] != '\0') { // the backend failed the graph (ggml_opt_eval): the run fails, as a refused graph does opt_failure = why; @@ -3619,7 +3662,27 @@ void llama_context::opt_epoch_iter( opt_stop_requested.store(true); return false; } - if (callback) { + if (xc != nullptr) { + // dL/d(cached K/V) of this chunk's prefix joins what later chunks put there + std::vector g; + for (const auto & io : res->t_walk) { + if (!io.dk) { + continue; + } + for (int kv = 0; kv < 2; ++kv) { + ggml_tensor * leaf = kv == 0 ? io.dk : io.dv; + ggml_tensor * grad = ggml_opt_leaf_grad(opt_ctx, leaf); + GGML_ASSERT(grad != nullptr && grad->type == GGML_TYPE_F32 && ggml_is_contiguous(grad)); + g.resize(ggml_nelements(grad)); + ggml_backend_tensor_get(grad, g.data(), 0, ggml_nbytes(grad)); + float * acc = (kv == 0 ? walk_gk : walk_gv)[io.il].data() + (size_t) io.grad_from*leaf->ne[0]*leaf->ne[1]; + for (size_t e = 0; e < g.size(); ++e) { + acc[e] += g[e]; + } + } + } + } + if (callback && !(xc != nullptr && xc->n_labels == 0)) { 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); @@ -3645,6 +3708,99 @@ void llama_context::opt_epoch_iter( return; } + if (train && opt_walk_exact) { + // THE EXACT WALK. The plain walk below steps per chunk and stops the gradient at every + // chunk boundary: a reply learns only through its own chunk's K/V, never through what + // an earlier chunk computed for it. Here the window is decoded once under the adapter + // as it is, then trained in REVERSE: each chunk's graph attends to the cache, with its + // last walk_horizon prefix positions behind a GRAD leaf, and adds the surrogate + // + on its own K/V, where G is what every later chunk put there. Its + // backward then carries that gradient through itself into the adapter and into its own + // prefix's G. ONE optimizer step per window, on the gradient of the window's loss + // (exact within the horizon). Memory stays chunk x window; the cost is one decode plus + // one forward and backward per chunk, context chunks included. + if (recr != nullptr) { + opt_failure = "the exact walk trains pure attention models: a recurrent state's gradient across chunks is not carried yet"; + LLAMA_LOG_ERROR("%s: %s\n", __func__, opt_failure.c_str()); + opt_alloc_failed.store(true); + opt_stop_requested.store(true); + return; + } + std::vector chunks; + uint32_t n_labels_window = 0; + for (uint32_t p = 0; p <= (uint32_t) last_label; ) { + const bool labelled = labels_sparse[p] >= 0; + uint32_t end = p; + while (end <= (uint32_t) last_label && (labels_sparse[end] >= 0) == labelled) { + ++end; + } + for (uint32_t c0 = p; c0 < end; c0 += n_batch) { + const uint32_t c1 = std::min(c0 + n_batch, end); + chunks.push_back({ c0, c1 - c0, labelled ? c1 - c0 : 0, 0, false, false, false, 0.0f }); + n_labels_window += labelled ? c1 - c0 : 0; + } + p = end; + } + // which chunks train, and how far back each one's gradient reaches + uint32_t reach = UINT32_MAX; // the lowest position a later trained chunk's gradient reaches + int64_t first = -1; // the last chunk the reverse pass trains: the step + for (int64_t j = (int64_t) chunks.size() - 1; j >= 0; --j) { + exact_chunk & c = chunks[j]; + c.surrogate = reach < c.c0 + c.n; + c.needed = c.n_labels > 0 || c.surrogate; + if (!c.needed) { + continue; + } + c.grad_from = opt_walk_horizon == 0 || c.c0 < opt_walk_horizon ? 0 : c.c0 - opt_walk_horizon; + c.loss_scale = (float) c.n_labels / (float) n_labels_window; + reach = std::min(reach, c.grad_from); + first = j; + } + GGML_ASSERT(first >= 0); + chunks[first].period_end = true; + + walk_gk.assign(model.hparams.n_layer(), {}); + walk_gv.assign(model.hparams.n_layer(), {}); + for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { + walk_gk[il].assign((size_t) (last_label + 1)*model.hparams.n_embd_k_gqa(il), 0.0f); + walk_gv[il].assign((size_t) (last_label + 1)*model.hparams.n_embd_v_gqa(il), 0.0f); + } + + // the window as the adapter now reads it, then the reverse pass, popping the cache + if (!decode_span(0, (uint32_t) last_label + 1)) { + return; + } + bool ok = true; + for (int64_t j = (int64_t) chunks.size() - 1; j >= first && ok; --j) { + const exact_chunk & c = chunks[j]; + if (!c.needed) { + continue; + } + if (!yield_point()) { + ok = false; + break; + } + if (!memory->seq_rm(0, c.c0, -1)) { + LLAMA_LOG_ERROR("%s: could not pop the cache to [0, %u) for the reverse pass\n", __func__, c.c0); + opt_stop_requested.store(true); + ok = false; + break; + } + cparams.walk_exact = true; + cparams.walk_grad_from = c.grad_from; + cparams.walk_surrogate = c.surrogate; + xc = &c; + ok = train_chunk(c.c0, c.n); + xc = nullptr; + cparams.walk_exact = false; + cparams.walk_grad_from = 0; + cparams.walk_surrogate = false; + } + walk_gk.clear(); + walk_gv.clear(); + 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; diff --git a/src/llama-context.h b/src/llama-context.h index 405a74275898..661af77c6028 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -369,6 +369,8 @@ struct llama_context { // model's hparams, because several contexts share one model (a serving context and a // training context on the same resident weights) uint32_t opt_n_ctx_train = 0; + bool opt_walk_exact = false; // llama_opt_params::walk_exact + uint32_t opt_walk_horizon = 0; // llama_opt_params::walk_horizon ggml_threadpool_t threadpool = nullptr; ggml_threadpool_t threadpool_batch = nullptr; diff --git a/src/llama-cparams.h b/src/llama-cparams.h index 85fbde9c3c9a..3199dfb3fff1 100644 --- a/src/llama-cparams.h +++ b/src/llama-cparams.h @@ -12,6 +12,15 @@ struct llama_cparams { // set by llama_opt_init: the graph is built for a backward pass, so attention reads K/V // from this ubatch directly rather than through the KV cache (see build_attn) bool training = false; + + // THE EXACT WALK (llama_context::opt_epoch_iter, walk_exact): a chunk's training graph in + // the reverse pass. walk_grad_from: the first cached prefix position whose K/V carry a + // gradient (positions [walk_grad_from, n_past) attend through a GRAD leaf; [0, + // walk_grad_from) stay constants). walk_surrogate: later chunks accumulated a gradient on + // this chunk's own K/V, so its graph adds the surrogate term that carries it in. + bool walk_exact = false; + uint32_t walk_grad_from = 0; + bool walk_surrogate = false; uint32_t n_ctx_seq; // context for a single sequence uint32_t n_batch; uint32_t n_ubatch; diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 7fef5d164082..bce1f65b0f95 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -1340,6 +1340,9 @@ void llm_graph_result::reset() { t_layer_inp.resize(LLAMA_MAX_LAYERS + 1); std::fill(t_layer_inp.begin(), t_layer_inp.end(), nullptr); + t_walk.clear(); + t_walk_surrogate = nullptr; + t_sampled.clear(); t_sampled_probs.clear(); t_sampled_logits.clear(); @@ -2849,29 +2852,83 @@ ggml_tensor * llm_graph_context::build_attn( const int64_t n_past = ubatch.pos[0]; k = k_cur; v = v_cur; - 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 cached prefix positions [p0, p1) as this chunk's K / V type + ggml_tensor * k_all = n_past > 0 ? mctx_cur->get_k(ctx0, il) : nullptr; + ggml_tensor * v_all = n_past > 0 ? mctx_cur->get_v(ctx0, il) : nullptr; + auto k_span = [&](int64_t p0, int64_t p1) { + ggml_tensor * t = ggml_view_4d(ctx0, k_all, k_all->ne[0], k_all->ne[1], p1 - p0, 1, + k_all->nb[1], k_all->nb[2], k_all->nb[3], p0*k_all->nb[2]); // 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); - + return t->type != k_cur->type ? ggml_cast(ctx0, t, k_cur->type) : t; + }; + auto v_span = [&](int64_t p0, int64_t p1) { // 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_tensor * t = 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), + ggml_view_4d(ctx0, v_all, p1 - p0, v_all->ne[1], v_all->ne[2], 1, + v_all->nb[1], v_all->nb[2], v_all->nb[3], p0*v_all->nb[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); + : ggml_view_4d(ctx0, v_all, v_all->ne[0], v_all->ne[1], p1 - p0, 1, + v_all->nb[1], v_all->nb[2], v_all->nb[3], p0*v_all->nb[2]); + return t->type != v_cur->type ? ggml_cast(ctx0, t, v_cur->type) : t; + }; + if (cparams.walk_exact) { + // THE EXACT WALK's chunk graph. The prefix positions [grad_from, n_past) are the + // cached values PLUS a zero GRAD leaf, so the backward yields dL/d(cached K/V) for + // them without copying the cache; earlier positions stay constants. The chunk's own + // K/V meet the gradient later chunks put on them through the surrogate + // + : its gradient on k_cur / v_cur IS gk / gv, which the + // backward carries through this chunk into the adapter and into its own prefix. + const int64_t g0 = std::min(cparams.walk_grad_from, n_past); + llm_graph_result::walk_layer io = { il, (uint32_t) g0, (uint32_t) n_past }; + // [const prefix][gradient prefix][this chunk]: the order the cells hold them in + std::vector k_parts, v_parts; + if (g0 > 0) { + k_parts.push_back(k_span(0, g0)); + v_parts.push_back(v_span(0, g0)); + } + if (n_past > g0) { + ggml_tensor * kc = ggml_cont(ctx0, ggml_cast(ctx0, k_span(g0, n_past), GGML_TYPE_F32)); + ggml_tensor * vc = ggml_cont(ctx0, ggml_cast(ctx0, v_span(g0, n_past), GGML_TYPE_F32)); + io.dk = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, kc->ne[0], kc->ne[1], kc->ne[2], 1); + io.dv = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, vc->ne[0], vc->ne[1], vc->ne[2], 1); + ggml_format_name(io.dk, "walk_dk-%d", il); + ggml_format_name(io.dv, "walk_dv-%d", il); + ggml_set_input(io.dk); + ggml_set_input(io.dv); + ggml_set_grad(io.dk); + ggml_set_grad(io.dv); + ggml_tensor * kg = ggml_add(ctx0, kc, io.dk); + ggml_tensor * vg = ggml_add(ctx0, vc, io.dv); + k_parts.push_back(kg->type != k_cur->type ? ggml_cast(ctx0, kg, k_cur->type) : kg); + v_parts.push_back(vg->type != v_cur->type ? ggml_cast(ctx0, vg, v_cur->type) : vg); + } + k_parts.push_back(k_cur); + v_parts.push_back(v_cur); + k = k_parts[0]; + v = v_parts[0]; + for (size_t i = 1; i < k_parts.size(); ++i) { + k = ggml_concat(ctx0, k, k_parts[i], 2); + v = ggml_concat(ctx0, v, v_parts[i], 2); + } + if (cparams.walk_surrogate) { + io.gk = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, k_cur->ne[0], k_cur->ne[1], k_cur->ne[2]); + io.gv = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, v_cur->ne[0], v_cur->ne[1], v_cur->ne[2]); + ggml_format_name(io.gk, "walk_gk-%d", il); + ggml_format_name(io.gv, "walk_gv-%d", il); + ggml_set_input(io.gk); + ggml_set_input(io.gv); + ggml_tensor * k32 = k_cur->type == GGML_TYPE_F32 ? k_cur : ggml_cast(ctx0, k_cur, GGML_TYPE_F32); + ggml_tensor * v32 = v_cur->type == GGML_TYPE_F32 ? v_cur : ggml_cast(ctx0, v_cur, GGML_TYPE_F32); + ggml_tensor * term = ggml_add(ctx0, + ggml_sum(ctx0, ggml_mul(ctx0, k32, io.gk)), + ggml_sum(ctx0, ggml_mul(ctx0, v32, io.gv))); + res->t_walk_surrogate = res->t_walk_surrogate ? ggml_add(ctx0, res->t_walk_surrogate, term) : term; } - v = ggml_concat(ctx0, v_prev, v_cur, 2); + res->t_walk.push_back(io); + } else if (n_past > 0) { + k = ggml_concat(ctx0, k_span(0, n_past), k_cur, 2); + v = ggml_concat(ctx0, v_span(0, n_past), 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) diff --git a/src/llama-graph.h b/src/llama-graph.h index b388e028cb53..d5dc8bd701c0 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -938,6 +938,20 @@ class llm_graph_result { std::vector t_layer_inp; + // THE EXACT WALK: per attention layer of a reverse-pass chunk graph, the tensors the walk + // fills before the step and reads after it (llama_context::opt_epoch_iter) + struct walk_layer { + int il; + uint32_t grad_from; // the GRAD leaves cover cached positions [grad_from, n_past) + uint32_t n_past; + ggml_tensor * dk = nullptr; // GRAD leaves, zero-filled: their gradient is dL/d(cached K/V) + ggml_tensor * dv = nullptr; // [n_embd_head, n_head_kv, n_past - grad_from], V untransposed + ggml_tensor * gk = nullptr; // inputs: the gradient later chunks accumulated on this chunk's + ggml_tensor * gv = nullptr; // own K/V, shaped like k_cur / v_cur (null without a surrogate) + }; + std::vector t_walk; + ggml_tensor * t_walk_surrogate = nullptr; // sum over layers of + + std::vector t_sampled; std::vector t_sampled_probs; std::vector t_sampled_logits; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index cbb9f6481699..96484141528c 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -150,6 +150,9 @@ endif () llama_build(test-recurrent-state-rollback.cpp) +# the exact walk trains on the window's whole gradient whatever the chunking (needs -m ) +llama_build(test-walk-exact.cpp) + # GGML_TENSOR_FLAG_GRAD: a leaf whose gradient the backward computes, never optimized (the walk reads the cached prefix K/V gradient) llama_build_and_test(test-grad-leaf.cpp) diff --git a/tests/test-walk-exact.cpp b/tests/test-walk-exact.cpp new file mode 100644 index 000000000000..7620dfcc4ecb --- /dev/null +++ b/tests/test-walk-exact.cpp @@ -0,0 +1,213 @@ +// what this catches: the exact walk (llama_opt_params::walk_exact) training on anything other +// than the gradient of the window's whole loss. The plain walk stops the gradient at every +// chunk boundary; the exact walk carries each chunk's gradient back through the cached K/V of +// every chunk before it, so how the window is CHUNKED must not change the step it takes. +// +// One SGD step from the same fresh adapter, read off the adapter itself: SGD moves every A and B +// by exactly -lr * gradient, so the steps compare gradients directly (Fable: compare gradients, +// not losses; a loss-change proxy was measured to be nonlinear at practical step sizes): +// 1. every position labelled: a chunk of the whole window is ONE graph, the true gradient; the +// exact walk in 4 chunks must match it, and the plain walk in 4 chunks must not +// 2. context, then a reply (a masked window, the walk's real case): the exact walk in 4 chunks +// must take the step it takes with one chunk per run, and the plain walk must not +// +// Needs a pure-attention model: test-walk-exact -m Qwen2.5-Coder-1.5B-Instruct-Q4_K_M.gguf + +#include "arg.h" +#include "common.h" +#include "llama.h" +#include "../src/llama-adapter.h" + +#include +#include +#include +#include +#include +#include + +static const float LR = 1e-3f; // an SGD step is linear in the gradient at any lr +static const uint32_t WINDOW = 512; + +static ggml_opt_optimizer_params sgd_pars(void *) { + ggml_opt_optimizer_params p = ggml_opt_get_default_optimizer_params(nullptr); + p.sgd.alpha = LR; + p.sgd.wd = 0.0f; + return p; +} + +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_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; // the walk's cached constants equal the chunk's own K/V + cparams.type_v = GGML_TYPE_F32; + return llama_init_from_model(model, cparams); +} + +// every A and B tensor of the adapter, in name order: SGD moves them by exactly -lr * gradient +static std::vector adapter_params(const llama_adapter_lora * adapter) { + std::map sorted; + for (const auto & [name, w] : adapter->ab_map) { + sorted[name] = &w; + } + std::vector out; + for (const auto & [name, w] : sorted) { + for (ggml_tensor * t : { w->a, w->b }) { + GGML_ASSERT(t->type == GGML_TYPE_F32); + const size_t n = out.size(); + out.resize(n + ggml_nelements(t)); + ggml_backend_tensor_get(t, out.data() + n, 0, ggml_nbytes(t)); + } + } + return out; +} + +// one SGD step from the fresh adapter: the step itself, -lr * the gradient it took +static std::vector step_delta(const common_params & params, llama_model * model, const std::string & init, + const std::vector & tokens, const std::vector & labelled, + uint32_t chunk, bool exact) { + llama_adapter_lora * adapter = llama_adapter_lora_init(model, init.c_str()); + GGML_ASSERT(adapter != nullptr); + const std::vector before = adapter_params(adapter); + + llama_context * ctx = make_ctx(params, model, chunk); + float scale = 1.0f; + GGML_ASSERT(llama_set_adapters_lora(ctx, &adapter, 1, &scale) == 0); + llama_opt_params lopt{}; + lopt.n_ctx_train = 0; + 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 = exact; + lopt.walk_horizon = 0; + llama_opt_init(ctx, model, lopt); + std::vector> seqs = { std::vector(tokens.begin(), tokens.begin() + WINDOW + 1) }; + std::vector> loss = { labelled }; + ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(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)); + { + double l = 0.0, unc = 0.0; + ggml_opt_result_loss(result, &l, &unc); + printf(" chunk %3u %-5s: training loss %.6f (the forward the step saw)\n", chunk, exact ? "exact" : "plain", l); + } + ggml_opt_result_free(result); + ggml_opt_dataset_free(dataset); + llama_free(ctx); + + const std::vector after = adapter_params(adapter); + llama_adapter_lora_free(adapter); + std::vector d(before.size()); + for (size_t i = 0; i < d.size(); ++i) { + d[i] = after[i] - before[i]; + } + return d; +} + +struct agreement { + double cosine; + double norm_ratio; +}; + +// the angle and the size of a step against the reference's: a wrong direction (a misrouted +// gradient) shows in the cosine, a wrong scale (a misweighted loss) in the norm ratio +static agreement compare(const char * what, const std::vector & a, const std::vector & ref) { + double dot = 0.0, na = 0.0, nr = 0.0; + for (size_t i = 0; i < a.size(); ++i) { + dot += (double) a[i] * ref[i]; + na += (double) a[i] * a[i]; + nr += (double) ref[i] * ref[i]; + } + const agreement r = { dot / std::sqrt(na * nr), std::sqrt(na / nr) }; + printf(" %s: cosine %.4f, norm ratio %.4f\n", what, r.cosine, r.norm_ratio); + return r; +} + +// Measured on the 1.5B Q4_K_M (5090): the exact walk agrees with one graph at cosine 0.998 and +// norm 0.995 (what is left is the quantized matmuls' batch-size variance: a 128-row chunk and a +// 512-row graph take different kernels); the plain walk sits at cosine 0.81, norm 2.1. +static bool same_step(const agreement & r) { + return r.cosine > 0.995 && std::fabs(r.norm_ratio - 1.0) < 0.02; +} + +int main(int argc, char ** argv) { + common_params params; + if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_COMMON)) { + return 1; + } + params.no_extra_bufts = true; // OUT_PROD dequantizes the standard layout only + llama_backend_init(); + auto mparams = common_model_params_to_llama(params); + llama_model * model = llama_model_load_from_file(params.model.path.c_str(), mparams); + GGML_ASSERT(model != nullptr); + + const std::string init = "test-walk-exact.init.gguf"; + GGML_ASSERT(common_lora_write_fresh(model, init, 8, 16, { "attn_q", "attn_v" }, 42)); + + // words in a fixed pseudo-random order: a window the model cannot already predict, so its + // loss, and the gradient under test, is well above the kernels' batch-size noise + static const char * words[] = { "room", "card", "claim", "review", "patch", "lease", "board", "turn", + "engine", "window", "chunk", "cache", "gradient", "adapter", "verdict", "submission" }; + std::string text; + uint32_t lcg = 12345; + while (text.size() < 8 * WINDOW) { + lcg = lcg * 1664525u + 1013904223u; + text += words[(lcg >> 16) % 16]; + text += (lcg >> 8) % 7 == 0 ? ".\n" : " "; + } + 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); + + int failures = 0; + { + // 1. every position labelled: chunk = window is one graph, the true gradient + std::vector all(WINDOW + 1, 1); + all[0] = 0; + const auto ref = step_delta(params, model, init, tokens, all, WINDOW, false); + const auto exact = step_delta(params, model, init, tokens, all, WINDOW / 4, true); + const auto plain = step_delta(params, model, init, tokens, all, WINDOW / 4, false); + if (!same_step(compare("all labelled: exact walk in 4 chunks vs one graph", exact, ref))) { + fprintf(stderr, "FAILED: the exact walk's step is not the window's gradient\n"); + ++failures; + } + if (same_step(compare("all labelled: plain walk in 4 chunks vs one graph", plain, ref))) { + fprintf(stderr, "FAILED: the plain walk matches one graph too: this window cannot tell the two apart\n"); + ++failures; + } + } + { + // 2. context, then a reply (the walk's real case): the exact walk is invariant to how the + // window is chunked, context chunks included (they train through the surrogate alone) + std::vector reply(WINDOW + 1, 0); + for (uint32_t i = WINDOW / 2 + 37; i <= WINDOW; ++i) { + reply[i] = 1; + } + const auto runs = step_delta(params, model, init, tokens, reply, WINDOW, true); // one chunk per run + const auto exact4 = step_delta(params, model, init, tokens, reply, WINDOW / 4, true); + const auto plain4 = step_delta(params, model, init, tokens, reply, WINDOW / 4, false); + if (!same_step(compare("context + reply: exact walk in 4 chunks vs one chunk per run", exact4, runs))) { + fprintf(stderr, "FAILED: the exact walk's step depends on the chunking\n"); + ++failures; + } + if (same_step(compare("context + reply: plain walk in 4 chunks vs the exact walk", plain4, runs))) { + fprintf(stderr, "FAILED: the plain walk matches the exact one: this window cannot tell the two apart\n"); + ++failures; + } + } + + std::remove(init.c_str()); + llama_model_free(model); + llama_backend_free(); + if (failures) { + return 1; + } + printf("test-walk-exact: OK\n"); + return 0; +} diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index 10e472fcb5ed..b6753855d22e 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", "chunk"}) { + for (const char * key : {"rank", "alpha", "window", "epochs", "lr", "val_split", "seed", "memory_budget_mib", "top_layers", "share_ppm", "max_slowdown_ppm", "chunk", "walk_horizon"}) { if (body.contains(key) && !body.at(key).is_number()) { return json::object({{"ok", false}, {"error", std::string("\"") + key + "\" must be a number"}}); } @@ -338,6 +338,10 @@ json server_trainer::start(const json & body_in) { 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("exact") && !body.at("exact").is_boolean()) + why = "exact must be true or false: the exact walk (one step per window on its whole loss) or the per-chunk walk"; + else if (body.contains("walk_horizon") && !(num("walk_horizon", 0) >= 0 && num("walk_horizon", 0) == (int64_t) num("walk_horizon", 0) && num("walk_horizon", 0) <= 4294967295.0)) + why = "walk_horizon must be an integer >= 0: how many cached positions before a chunk receive its gradient (0 = the whole window)"; 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))) @@ -861,6 +865,10 @@ void server_trainer::run(json req, examples_data ex) { // with it on is measured beside #39's receipt; Cormac): one layer's // attention scores alive in the backward pass instead of every layer's /*recompute =*/ req.value("recompute", false), + // the exact walk (OFF unless the request says "exact": true, until it is measured against + // the single-context gradient): one step per window on the window's whole loss + /*walk_exact =*/ req.value("exact", false), + /*walk_horizon =*/ (uint32_t) req.value("walk_horizon", (int64_t) 0), }; llama_opt_set_memory_budget(ctx, budget); llama_opt_init(ctx, model, lopt); From ed292df41f810249a19f7d44ed2d5a9329669949 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 19:05:03 -0500 Subject: [PATCH 07/13] train: the exact walk carries a recurrent state's gradient across chunks (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 , 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 104 +++++++++++++++++++++++++++++----- src/llama-cparams.h | 2 + src/llama-graph.cpp | 28 +++++++++ src/llama-graph.h | 21 +++++++ src/models/delta-net-base.cpp | 3 + tests/test-walk-exact.cpp | 7 ++- 6 files changed, 150 insertions(+), 15 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index cd7f98a65758..37d7b1622090 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3512,6 +3512,9 @@ void llama_context::opt_epoch_iter( const exact_chunk * xc = nullptr; // per layer, per position: dL/d(cached K) and dL/d(cached V), [position][n_embd_*_gqa] std::vector> walk_gk, walk_gv; + // per recurrent state tensor: dL/d(the state the chunk just before leaves), from the chunk + // trained last (each chunk's exit state feeds exactly one chunk, the next) + std::unordered_map> walk_gstate; // one training chunk [pos_ctx, pos_ctx + n_tokens): forward and backward over it alone, // attending to the cached context before it @@ -3652,6 +3655,18 @@ void llama_context::opt_epoch_iter( } } } + if (xc != nullptr) { + for (const auto & ws : res->t_walk_state) { + if (ws.ds) { + ggml_backend_tensor_memset(ws.ds, 0, 0, ggml_nbytes(ws.ds)); + } + if (ws.gs) { + const auto it = walk_gstate.find(ws.cache); + GGML_ASSERT(it != walk_gstate.end() && it->second.size() == (size_t) ggml_nelements(ws.gs)); + ggml_backend_tensor_set(ws.gs, it->second.data(), 0, ggml_nbytes(ws.gs)); + } + } + } ggml_opt_eval(opt_ctx, xc != nullptr && xc->n_labels == 0 ? nullptr : result); if (const char * why = ggml_opt_refusal(opt_ctx); why[0] != '\0') { // the backend failed the graph (ggml_opt_eval): the run fails, as a refused graph does @@ -3682,6 +3697,19 @@ void llama_context::opt_epoch_iter( } } } + if (xc != nullptr) { + // dL/d(entry state): what the previous chunk's exit state takes in its surrogate + for (const auto & ws : res->t_walk_state) { + if (!ws.ds) { + continue; + } + ggml_tensor * grad = ggml_opt_leaf_grad(opt_ctx, ws.ds); + GGML_ASSERT(grad != nullptr && grad->type == GGML_TYPE_F32 && ggml_is_contiguous(grad)); + auto & g = walk_gstate[ws.cache]; + g.resize(ggml_nelements(grad)); + ggml_backend_tensor_get(grad, g.data(), 0, ggml_nbytes(grad)); + } + } if (callback && !(xc != nullptr && xc->n_labels == 0)) { callback(train, opt_ctx, dataset, result, idata_in_loop + (pos_ctx + pos_batch)/n_batch + 1, ndata_in_loop, t_loop_start); } @@ -3702,7 +3730,8 @@ void llama_context::opt_epoch_iter( } else { recr = dynamic_cast(memory.get()); } - if (recr != nullptr && cparams.n_seq_max < 2) { + // the plain walk snapshots a recurrent state into a scratch sequence; the exact walk keeps host snapshots + if (recr != nullptr && cparams.n_seq_max < 2 && !(train && opt_walk_exact)) { 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; @@ -3719,13 +3748,10 @@ void llama_context::opt_epoch_iter( // prefix's G. ONE optimizer step per window, on the gradient of the window's loss // (exact within the horizon). Memory stays chunk x window; the cost is one decode plus // one forward and backward per chunk, context chunks included. - if (recr != nullptr) { - opt_failure = "the exact walk trains pure attention models: a recurrent state's gradient across chunks is not carried yet"; - LLAMA_LOG_ERROR("%s: %s\n", __func__, opt_failure.c_str()); - opt_alloc_failed.store(true); - opt_stop_requested.store(true); - return; - } + // A recurrent model's state is carried the same way, one chunk at a time: the state + // entering a chunk is a GRAD leaf, the state it leaves meets the next chunk's gradient + // (build_rs, build_walk_state_exit). The reverse pass restores each chunk's entry state + // from a host snapshot taken at its boundary in the forward pass. std::vector chunks; uint32_t n_labels_window = 0; for (uint32_t p = 0; p <= (uint32_t) last_label; ) { @@ -3747,7 +3773,8 @@ void llama_context::opt_epoch_iter( for (int64_t j = (int64_t) chunks.size() - 1; j >= 0; --j) { exact_chunk & c = chunks[j]; c.surrogate = reach < c.c0 + c.n; - c.needed = c.n_labels > 0 || c.surrogate; + // a recurrent state reaches every later chunk: a chunk before a trained one trains + c.needed = c.n_labels > 0 || c.surrogate || (recr != nullptr && first >= 0); if (!c.needed) { continue; } @@ -3766,9 +3793,48 @@ void llama_context::opt_epoch_iter( walk_gv[il].assign((size_t) (last_label + 1)*model.hparams.n_embd_v_gqa(il), 0.0f); } - // the window as the adapter now reads it, then the reverse pass, popping the cache - if (!decode_span(0, (uint32_t) last_label + 1)) { - return; + // the window as the adapter now reads it, chunk by chunk: a recurrent model's state is + // snapshotted at each chunk's start (empty before the first: it starts from zero) + struct state_snapshot { + bool empty = true; + std::vector> r, s; // per layer + }; + std::vector snaps(recr != nullptr ? chunks.size() : 0); + auto state_rows = [&](bool write, state_snapshot & snap) { + const int32_t cell = recr->cells[0].tail; // the cell holding sequence 0's state + if (!write) { + snap.empty = cell < 0; + snap.r.assign(recr->r_l.size(), {}); + snap.s.assign(recr->s_l.size(), {}); + } + if (cell < 0) { + return; + } + for (size_t il = 0; il < recr->r_l.size(); ++il) { + for (int kind = 0; kind < 2; ++kind) { + ggml_tensor * t = kind == 0 ? recr->r_l[il] : recr->s_l[il]; + if (t == nullptr) { + continue; + } + GGML_ASSERT(t->type == GGML_TYPE_F32); + auto & row = kind == 0 ? snap.r[il] : snap.s[il]; + const size_t bytes = t->nb[1]; + if (write) { + ggml_backend_tensor_set(t, row.data(), (size_t) cell*bytes, bytes); + } else { + row.resize(t->ne[0]); + ggml_backend_tensor_get(t, row.data(), (size_t) cell*bytes, bytes); + } + } + } + }; + for (size_t j = 0; j < chunks.size(); ++j) { + if (recr != nullptr) { + state_rows(false, snaps[j]); + } + if (!decode_span(chunks[j].c0, chunks[j].c0 + chunks[j].n)) { + return; + } } bool ok = true; for (int64_t j = (int64_t) chunks.size() - 1; j >= first && ok; --j) { @@ -3780,24 +3846,36 @@ void llama_context::opt_epoch_iter( ok = false; break; } - if (!memory->seq_rm(0, c.c0, -1)) { + // attention: pop the cache to [0, c0); recurrent: restore the chunk's entry state + const bool popped = recr != nullptr ? (attn == nullptr || attn->seq_rm(0, c.c0, -1)) : memory->seq_rm(0, c.c0, -1); + if (!popped) { LLAMA_LOG_ERROR("%s: could not pop the cache to [0, %u) for the reverse pass\n", __func__, c.c0); opt_stop_requested.store(true); ok = false; break; } + if (recr != nullptr) { + if (snaps[j].empty) { + recr->seq_rm(0, -1, -1); // the first chunk starts from a zero state + } else { + state_rows(true, snaps[j]); + } + } cparams.walk_exact = true; cparams.walk_grad_from = c.grad_from; cparams.walk_surrogate = c.surrogate; + cparams.walk_state_surrogate = recr != nullptr && j < (int64_t) chunks.size() - 1 && !walk_gstate.empty(); xc = &c; ok = train_chunk(c.c0, c.n); xc = nullptr; cparams.walk_exact = false; cparams.walk_grad_from = 0; cparams.walk_surrogate = false; + cparams.walk_state_surrogate = false; } walk_gk.clear(); walk_gv.clear(); + walk_gstate.clear(); return; } diff --git a/src/llama-cparams.h b/src/llama-cparams.h index 3199dfb3fff1..5df493dec5da 100644 --- a/src/llama-cparams.h +++ b/src/llama-cparams.h @@ -21,6 +21,8 @@ struct llama_cparams { bool walk_exact = false; uint32_t walk_grad_from = 0; bool walk_surrogate = false; + // a later chunk's gradient sits on the recurrent state this chunk leaves + bool walk_state_surrogate = false; uint32_t n_ctx_seq; // context for a single sequence uint32_t n_batch; uint32_t n_ubatch; diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index bce1f65b0f95..f87e4bc54e5a 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -1341,6 +1341,7 @@ void llm_graph_result::reset() { std::fill(t_layer_inp.begin(), t_layer_inp.end(), nullptr); t_walk.clear(); + t_walk_state.clear(); t_walk_surrogate = nullptr; t_sampled.clear(); @@ -3551,6 +3552,17 @@ ggml_tensor * llm_graph_context::build_rs( // NOTE: assuming the copy destinations are ALL contained between rs_head and rs_head + n_rs // {state_size, rs_size} -> {state_size, n_seqs} ggml_tensor * output_states = get_state_rows(ctx0, states, state_copy_main); + if (cparams.training && cparams.walk_exact) { + // THE EXACT WALK: the state entering the chunk is the cached one PLUS a zero GRAD leaf, + // so the backward yields dL/d(entry state): the gradient the previous chunk's exit takes + GGML_ASSERT(output_states->type == GGML_TYPE_F32 && "the exact walk carries an F32 recurrent state"); + ggml_tensor * ds = ggml_new_tensor(ctx0, GGML_TYPE_F32, GGML_MAX_DIMS, output_states->ne); + ggml_format_name(ds, "walk_ds-%s", s->name); + ggml_set_input(ds); + ggml_set_grad(ds); + output_states = ggml_add(ctx0, output_states, ds); + res->walk_state_of(s).ds = ds; + } ggml_build_forward_expand(gf, output_states); // copy extra states which won't be changed further (between n_seqs and n_rs) @@ -3606,6 +3618,22 @@ ggml_tensor * llm_graph_context::build_rs( get_state_rows); } +void llm_graph_context::build_walk_state_exit(ggml_tensor * cache, ggml_tensor * exit) const { + if (!(cparams.training && cparams.walk_exact && cparams.walk_state_surrogate)) { + return; + } + ggml_tensor * flat = ggml_reshape_1d(ctx0, ggml_cont(ctx0, exit), ggml_nelements(exit)); + if (flat->type != GGML_TYPE_F32) { + flat = ggml_cast(ctx0, flat, GGML_TYPE_F32); + } + ggml_tensor * gs = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, ggml_nelements(exit)); + ggml_format_name(gs, "walk_gs-%s", cache->name); + ggml_set_input(gs); + ggml_tensor * term = ggml_sum(ctx0, ggml_mul(ctx0, flat, gs)); + res->t_walk_surrogate = res->t_walk_surrogate ? ggml_add(ctx0, res->t_walk_surrogate, term) : term; + res->walk_state_of(cache).gs = gs; +} + ggml_tensor * llm_graph_context::build_rwkv_token_shift_load( llm_graph_input_rs * inp, const llama_ubatch & ubatch, diff --git a/src/llama-graph.h b/src/llama-graph.h index d5dc8bd701c0..5f924306f7f6 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -950,6 +950,23 @@ class llm_graph_result { ggml_tensor * gv = nullptr; // own K/V, shaped like k_cur / v_cur (null without a surrogate) }; std::vector t_walk; + // the recurrent states of a reverse-pass chunk graph, one per state tensor of the memory + // (a layer's conv state r_l and its recurrent state s_l): the chain runs through them too + struct walk_state { + ggml_tensor * cache = nullptr; // the memory's state tensor this entry belongs to + ggml_tensor * ds = nullptr; // GRAD leaf, zero-filled, on the state ENTERING the chunk + ggml_tensor * gs = nullptr; // input: the gradient later chunks put on the state it LEAVES + }; + std::vector t_walk_state; + walk_state & walk_state_of(ggml_tensor * cache) { + for (auto & w : t_walk_state) { + if (w.cache == cache) { + return w; + } + } + t_walk_state.push_back({ cache, nullptr, nullptr }); + return t_walk_state.back(); + } ggml_tensor * t_walk_surrogate = nullptr; // sum over layers of + std::vector t_sampled; @@ -1340,6 +1357,10 @@ struct llm_graph_context { int32_t n_seqs, const llm_graph_get_rows_fn & get_state_rows = ggml_get_rows) const; + // THE EXACT WALK: the state a chunk leaves (`exit`, what is written back to `cache`) meets the + // gradient later chunks put on it, through the surrogate . A no-op outside the walk. + void build_walk_state_exit(ggml_tensor * cache, ggml_tensor * exit) const; + ggml_tensor * build_rwkv_token_shift_load( llm_graph_input_rs * inp, const llama_ubatch & ubatch, diff --git a/src/models/delta-net-base.cpp b/src/models/delta-net-base.cpp index ad6612647736..e0a6368bd2ed 100644 --- a/src/models/delta-net-base.cpp +++ b/src/models/delta-net-base.cpp @@ -494,7 +494,9 @@ ggml_tensor * llm_build_delta_net_base::build_conv_state( cb(conv_state_update, "conv_state_update", il); ggml_build_forward_expand(gf, ggml_cpy(ctx0, conv_state_last, conv_state_update)); + build_walk_state_exit(conv_states_all, conv_state_last); } else { + GGML_ASSERT(!cparams.walk_exact && "the exact walk keeps one recurrent state per sequence (n_rs_seq = 0)"); // [TAG_RECURRENT_ROLLBACK_SPLITS] // this logic assumes that the last (n_rs_seq + 1) tokens of a sequence in a batch are inside // the same ubatch, which `split_equal()` guarantees via its n_keep_tail argument @@ -556,6 +558,7 @@ ggml_tensor * llm_build_delta_net_base::build_recurrent_attn( ggml_cpy(ctx0, new_state, ggml_view_2d(ctx0, ssm_states_all, hparams.n_embd_s(), n_seqs, ssm_states_all->nb[1], kv_head * hparams.n_embd_s() * ggml_element_size(ssm_states_all)))); + build_walk_state_exit(ssm_states_all, new_state); return output; } diff --git a/tests/test-walk-exact.cpp b/tests/test-walk-exact.cpp index 7620dfcc4ecb..27d63d74b1b1 100644 --- a/tests/test-walk-exact.cpp +++ b/tests/test-walk-exact.cpp @@ -11,7 +11,9 @@ // 2. context, then a reply (a masked window, the walk's real case): the exact walk in 4 chunks // must take the step it takes with one chunk per run, and the plain walk must not // -// Needs a pure-attention model: test-walk-exact -m Qwen2.5-Coder-1.5B-Instruct-Q4_K_M.gguf +// Run on a pure-attention model and on a hybrid (attention + gated delta-net): +// test-walk-exact -m Qwen2.5-Coder-1.5B-Instruct-Q4_K_M.gguf +// test-walk-exact -m Qwen3.5-0.8B-Q8_0.gguf #include "arg.h" #include "common.h" @@ -40,7 +42,8 @@ static llama_context * make_ctx(const common_params & params, llama_model * mode cparams.n_ctx = WINDOW; cparams.n_batch = chunk; cparams.n_ubatch = chunk; - cparams.n_seq_max = 1; + cparams.n_seq_max = 2; // the plain walk snapshots a recurrent state into a scratch sequence + cparams.kv_unified = true; // ...and the window keeps every cell cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; cparams.type_k = GGML_TYPE_F32; // the walk's cached constants equal the chunk's own K/V cparams.type_v = GGML_TYPE_F32; From 8b5f095c5e5f9f8fdabb466f5666e3c4e1b7791e Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 19:09:34 -0500 Subject: [PATCH 08/13] train: the exact walk's host memory is counted before a window runs and 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- include/llama.h | 10 ++++++++ src/llama-context.cpp | 43 +++++++++++++++++++++++++++++++++++ src/llama-context.h | 3 +++ tests/test-walk-exact.cpp | 31 +++++++++++++++++++++++++ tools/server/server-train.cpp | 9 +++++++- 5 files changed, 95 insertions(+), 1 deletion(-) diff --git a/include/llama.h b/include/llama.h index 1a2679a090a7..d7972a233f12 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1755,6 +1755,16 @@ extern "C" { // batch): the footprint a run of this shape needs, from the run's own allocation preflight. LLAMA_API size_t llama_opt_graph_bytes(struct llama_context * lctx); + // Cap, in bytes, the HOST memory the exact walk keeps per training window (0 = no cap): the + // gradient accumulated on every cached K/V position of every attention layer, and a recurrent + // model's state snapshot at every chunk boundary. A window over the cap refuses the run by + // name before anything is decoded (llama_opt_failure). Set before llama_opt_epoch. + LLAMA_API void llama_opt_set_walk_host_budget(struct llama_context * lctx, size_t bytes); + + // The host bytes the exact walk's largest window so far kept (0 before one ran, or when the + // walk is not exact): what a run of this shape needs, from the run's own arithmetic. + LLAMA_API size_t llama_opt_walk_host_bytes(struct llama_context * lctx); + LLAMA_API void llama_opt_epoch( struct llama_context * lctx, ggml_opt_dataset_t dataset, diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 37d7b1622090..124b0f0a0522 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3786,9 +3786,44 @@ void llama_context::opt_epoch_iter( GGML_ASSERT(first >= 0); chunks[first].period_end = true; + // THE HOST BUDGET (Fable): the accumulators are layers x positions x (k + v) floats, the + // snapshots chunks x one sequence's state rows; both are known before anything runs. + { + size_t bytes = 0; + for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { + if (!model.hparams.is_recr(il)) { + bytes += (size_t) (last_label + 1)*(model.hparams.n_embd_k_gqa(il) + model.hparams.n_embd_v_gqa(il))*sizeof(float); + } + } + if (recr != nullptr) { + size_t state = 0; + for (size_t il = 0; il < recr->r_l.size(); ++il) { + state += recr->r_l[il] ? recr->r_l[il]->nb[1] : 0; + state += recr->s_l[il] ? recr->s_l[il]->nb[1] : 0; + } + bytes += chunks.size()*state; + } + opt_walk_host_bytes = std::max(opt_walk_host_bytes, bytes); + if (opt_walk_host_budget > 0 && bytes > opt_walk_host_budget) { + char why[512]; + snprintf(why, sizeof(why), + "the exact walk needs %.1f MiB of host memory for this window (%lld positions, %zu chunks), over its budget of %.1f MiB: " + "the K/V gradient grows with the window and a recurrent state's snapshots with the chunk count (a larger chunk has fewer); " + "the horizon does not change either", + bytes/1048576.0, (long long) (last_label + 1), chunks.size(), opt_walk_host_budget/1048576.0); + opt_failure = why; + LLAMA_LOG_ERROR("%s: %s\n", __func__, opt_failure.c_str()); + opt_alloc_failed.store(true); + opt_stop_requested.store(true); + return; + } + } walk_gk.assign(model.hparams.n_layer(), {}); walk_gv.assign(model.hparams.n_layer(), {}); for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { + if (model.hparams.is_recr(il)) { + continue; // a recurrent layer has no K/V: its state carries the chain + } walk_gk[il].assign((size_t) (last_label + 1)*model.hparams.n_embd_k_gqa(il), 0.0f); walk_gv[il].assign((size_t) (last_label + 1)*model.hparams.n_embd_v_gqa(il), 0.0f); } @@ -4698,6 +4733,14 @@ void llama_opt_set_step_callback(struct llama_context * ctx, llama_opt_step_call ctx->opt_step_callback_data = user_data; } +void llama_opt_set_walk_host_budget(struct llama_context * ctx, size_t bytes) { + ctx->opt_walk_host_budget = bytes; +} + +size_t llama_opt_walk_host_bytes(struct llama_context * ctx) { + return ctx->opt_walk_host_bytes; +} + void llama_opt_set_memory_budget(struct llama_context * ctx, size_t bytes) { ctx->opt_memory_budget = bytes; } diff --git a/src/llama-context.h b/src/llama-context.h index 661af77c6028..b5b9f7bdc939 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -208,6 +208,9 @@ struct llama_context { std::string opt_failure; // what training may add per GPU device (llama_opt_set_memory_budget); 0 = no cap size_t opt_memory_budget = 0; + // the exact walk's host memory per window (llama_opt_set_walk_host_budget); 0 = no cap + size_t opt_walk_host_budget = 0; + size_t opt_walk_host_bytes = 0; // the largest window's, measured by its own arithmetic // the largest training graph the allocation preflight measured on a GPU device (bytes) size_t opt_graph_bytes() const; diff --git a/tests/test-walk-exact.cpp b/tests/test-walk-exact.cpp index 27d63d74b1b1..4b8b0efe9ab2 100644 --- a/tests/test-walk-exact.cpp +++ b/tests/test-walk-exact.cpp @@ -205,6 +205,37 @@ int main(int argc, char ** argv) { } } + { + // 3. the host budget: a window whose accumulators and snapshots exceed it refuses by + // name before anything runs, and the adapter is untouched + 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; + llama_opt_init(ctx, model, lopt); + llama_opt_set_walk_host_budget(ctx, 1); + std::vector> seqs = { std::vector(tokens.begin(), tokens.begin() + WINDOW + 1) }; + std::vector> loss = { std::vector(WINDOW + 1, 1) }; + ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(WINDOW, seqs, loss, tokens[0]); + llama_opt_epoch(ctx, dataset, nullptr, nullptr, /*idata_split =*/ 1, nullptr, nullptr); + const std::string why = llama_opt_failure(ctx); + printf(" host budget of 1 byte: failed=%d, \"%s\", needs %zu MiB\n", (int) llama_opt_failed(ctx), why.c_str(), llama_opt_walk_host_bytes(ctx) >> 20); + if (!llama_opt_failed(ctx) || why.find("host memory") == std::string::npos || adapter_params(adapter) != before) { + fprintf(stderr, "FAILED: a window over the host budget was not refused by name before it ran\n"); + ++failures; + } + ggml_opt_dataset_free(dataset); + llama_free(ctx); + llama_adapter_lora_free(adapter); + } + std::remove(init.c_str()); llama_model_free(model); llama_backend_free(); diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index b6753855d22e..c9f91586af78 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", "chunk", "walk_horizon"}) { + for (const char * key : {"rank", "alpha", "window", "epochs", "lr", "val_split", "seed", "memory_budget_mib", "top_layers", "share_ppm", "max_slowdown_ppm", "chunk", "walk_horizon", "walk_host_budget_mib"}) { if (body.contains(key) && !body.at(key).is_number()) { return json::object({{"ok", false}, {"error", std::string("\"") + key + "\" must be a number"}}); } @@ -340,6 +340,8 @@ json server_trainer::start(const json & body_in) { 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("exact") && !body.at("exact").is_boolean()) why = "exact must be true or false: the exact walk (one step per window on its whole loss) or the per-chunk walk"; + else if (body.contains("walk_host_budget_mib") && !(num("walk_host_budget_mib", 0) >= 1 && num("walk_host_budget_mib", 0) == (int64_t) num("walk_host_budget_mib", 0))) + why = "walk_host_budget_mib must be an integer >= 1: the host memory the exact walk may keep per window"; else if (body.contains("walk_horizon") && !(num("walk_horizon", 0) >= 0 && num("walk_horizon", 0) == (int64_t) num("walk_horizon", 0) && num("walk_horizon", 0) <= 4294967295.0)) why = "walk_horizon must be an integer >= 0: how many cached positions before a chunk receive its gradient (0 = the whole window)"; else if (body.contains("memory_budget_mib") && !(num("memory_budget_mib", 0) >= 1)) @@ -871,6 +873,8 @@ void server_trainer::run(json req, examples_data ex) { /*walk_horizon =*/ (uint32_t) req.value("walk_horizon", (int64_t) 0), }; llama_opt_set_memory_budget(ctx, budget); + // the exact walk's host memory per window: refused by name over the caller's cap + llama_opt_set_walk_host_budget(ctx, (size_t) req.value("walk_host_budget_mib", (int64_t) 0) << 20); llama_opt_init(ctx, model, lopt); llama_opt_set_step_callback(ctx, &server_trainer::before_window, this); { @@ -974,6 +978,7 @@ void server_trainer::run(json req, examples_data ex) { // copied before llama_free: the refusal names a node the device cannot run, when that was it const std::string refusal = llama_opt_failure(ctx); const double graph_mib = llama_opt_graph_bytes(ctx) / 1048576.0; + const double walk_host_mib = llama_opt_walk_host_bytes(ctx) / 1048576.0; const int32_t saved = (cancelled || no_fit) ? 0 : llama_adapter_lora_save(adapter, out.c_str()); llama_adapter_lora_free(adapter); llama_free(ctx); @@ -982,6 +987,8 @@ void server_trainer::run(json req, examples_data ex) { // what the training graph needed on the GPU, measured by its own preflight (also on a refusal: // it is the number to size the next attempt, or a caller's admission, by) state["graph_mib"] = graph_mib; + // what the exact walk kept on the host per window (its K/V gradient and state snapshots) + state["walk_host_mib"] = walk_host_mib; if (no_fit) { state["state"] = "error"; state["error"] = !refusal.empty() ? refusal + " (window " + std::to_string(window) + ")" : From 5a06f8bcd96d0216ddad929615578bb24318c9fc Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 19:14:13 -0500 Subject: [PATCH 09/13] train: the exact walk checkpoints a recurrent state every few chunks 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 54 +++++++++++++++++++++++++++------------ tests/test-walk-exact.cpp | 20 +++++++++++++-- 2 files changed, 56 insertions(+), 18 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 124b0f0a0522..1f4011260277 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3786,37 +3786,51 @@ void llama_context::opt_epoch_iter( GGML_ASSERT(first >= 0); chunks[first].period_end = true; - // THE HOST BUDGET (Fable): the accumulators are layers x positions x (k + v) floats, the - // snapshots chunks x one sequence's state rows; both are known before anything runs. + // THE HOST BUDGET (Fable): the accumulators are layers x positions x (k + v) floats, and a + // recurrent state is checkpointed every `stride` chunks (one sequence's state rows each; + // the first chunk's entry is the zero state and costs nothing). Both are known before + // anything runs. Under a budget the stride is the smallest that fits: the reverse pass + // rebuilds a chunk's entry state by decoding forward from its checkpoint, so a longer + // stride trades host memory for decode, never for exactness. + size_t stride = 1; { - size_t bytes = 0; + size_t kv_bytes = 0; for (uint32_t il = 0; il < model.hparams.n_layer(); ++il) { if (!model.hparams.is_recr(il)) { - bytes += (size_t) (last_label + 1)*(model.hparams.n_embd_k_gqa(il) + model.hparams.n_embd_v_gqa(il))*sizeof(float); + kv_bytes += (size_t) (last_label + 1)*(model.hparams.n_embd_k_gqa(il) + model.hparams.n_embd_v_gqa(il))*sizeof(float); } } + size_t state = 0; if (recr != nullptr) { - size_t state = 0; for (size_t il = 0; il < recr->r_l.size(); ++il) { state += recr->r_l[il] ? recr->r_l[il]->nb[1] : 0; state += recr->s_l[il] ? recr->s_l[il]->nb[1] : 0; } - bytes += chunks.size()*state; } + auto bytes_at = [&](size_t s) { return kv_bytes + ((chunks.size() + s - 1)/s - 1)*state; }; + if (opt_walk_host_budget > 0) { + while (stride < chunks.size() && bytes_at(stride) > opt_walk_host_budget) { + ++stride; + } + } + const size_t bytes = bytes_at(stride); opt_walk_host_bytes = std::max(opt_walk_host_bytes, bytes); if (opt_walk_host_budget > 0 && bytes > opt_walk_host_budget) { char why[512]; snprintf(why, sizeof(why), - "the exact walk needs %.1f MiB of host memory for this window (%lld positions, %zu chunks), over its budget of %.1f MiB: " - "the K/V gradient grows with the window and a recurrent state's snapshots with the chunk count (a larger chunk has fewer); " - "the horizon does not change either", - bytes/1048576.0, (long long) (last_label + 1), chunks.size(), opt_walk_host_budget/1048576.0); + "the exact walk needs %.1f MiB of host memory for this window (%lld positions, the K/V gradient alone), over its budget of %.1f MiB: " + "the K/V gradient grows with the window; neither the horizon nor the chunk size changes it", + bytes/1048576.0, (long long) (last_label + 1), opt_walk_host_budget/1048576.0); opt_failure = why; LLAMA_LOG_ERROR("%s: %s\n", __func__, opt_failure.c_str()); opt_alloc_failed.store(true); opt_stop_requested.store(true); return; } + if (stride > 1) { + LLAMA_LOG_INFO("%s: the exact walk checkpoints the recurrent state every %zu chunks to fit %.1f MiB of host memory (%.1f MiB)\n", + __func__, stride, opt_walk_host_budget/1048576.0, bytes/1048576.0); + } } walk_gk.assign(model.hparams.n_layer(), {}); walk_gv.assign(model.hparams.n_layer(), {}); @@ -3864,7 +3878,7 @@ void llama_context::opt_epoch_iter( } }; for (size_t j = 0; j < chunks.size(); ++j) { - if (recr != nullptr) { + if (recr != nullptr && j % stride == 0) { state_rows(false, snaps[j]); } if (!decode_span(chunks[j].c0, chunks[j].c0 + chunks[j].n)) { @@ -3881,19 +3895,27 @@ void llama_context::opt_epoch_iter( ok = false; break; } - // attention: pop the cache to [0, c0); recurrent: restore the chunk's entry state - const bool popped = recr != nullptr ? (attn == nullptr || attn->seq_rm(0, c.c0, -1)) : memory->seq_rm(0, c.c0, -1); + // 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__, c.c0); + 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[j].empty) { + if (snaps[cp].empty) { recr->seq_rm(0, -1, -1); // the first chunk starts from a zero state } else { - state_rows(true, snaps[j]); + state_rows(true, snaps[cp]); + } + if (p_pop < c.c0 && !decode_span(p_pop, c.c0)) { + ok = false; + break; } } cparams.walk_exact = true; diff --git a/tests/test-walk-exact.cpp b/tests/test-walk-exact.cpp index 4b8b0efe9ab2..5068c9b84db0 100644 --- a/tests/test-walk-exact.cpp +++ b/tests/test-walk-exact.cpp @@ -71,7 +71,7 @@ static std::vector adapter_params(const llama_adapter_lora * adapter) { // one SGD step from the fresh adapter: the step itself, -lr * the gradient it took static std::vector step_delta(const common_params & params, llama_model * model, const std::string & init, const std::vector & tokens, const std::vector & labelled, - uint32_t chunk, bool exact) { + uint32_t chunk, bool exact, size_t host_budget = 0, size_t * host_bytes = nullptr) { llama_adapter_lora * adapter = llama_adapter_lora_init(model, init.c_str()); GGML_ASSERT(adapter != nullptr); const std::vector before = adapter_params(adapter); @@ -88,6 +88,7 @@ static std::vector step_delta(const common_params & params, llama_model * lopt.walk_exact = exact; lopt.walk_horizon = 0; 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> loss = { labelled }; ggml_opt_dataset_t dataset = common_opt_dataset_init_masked(WINDOW, seqs, loss, tokens[0]); @@ -99,6 +100,9 @@ static std::vector step_delta(const common_params & params, llama_model * ggml_opt_result_loss(result, &l, &unc); printf(" chunk %3u %-5s: training loss %.6f (the forward the step saw)\n", chunk, exact ? "exact" : "plain", l); } + if (host_bytes) { + *host_bytes = llama_opt_walk_host_bytes(ctx); + } ggml_opt_result_free(result); ggml_opt_dataset_free(dataset); llama_free(ctx); @@ -174,7 +178,8 @@ int main(int argc, char ** argv) { std::vector all(WINDOW + 1, 1); all[0] = 0; const auto ref = step_delta(params, model, init, tokens, all, WINDOW, false); - const auto exact = step_delta(params, model, init, tokens, all, WINDOW / 4, true); + size_t full_bytes = 0; + const auto exact = step_delta(params, model, init, tokens, all, WINDOW / 4, true, 0, &full_bytes); const auto plain = step_delta(params, model, init, tokens, all, WINDOW / 4, false); if (!same_step(compare("all labelled: exact walk in 4 chunks vs one graph", exact, ref))) { fprintf(stderr, "FAILED: the exact walk's step is not the window's gradient\n"); @@ -184,6 +189,17 @@ int main(int argc, char ** argv) { fprintf(stderr, "FAILED: the plain walk matches one graph too: this window cannot tell the two apart\n"); ++failures; } + if (llama_model_is_hybrid(model)) { + // a host budget below every chunk's state snapshot: the walk checkpoints every few + // chunks and decodes forward from the checkpoint, and the step must not change + size_t strided_bytes = 0; + const auto strided = step_delta(params, model, init, tokens, all, WINDOW / 4, true, full_bytes * 3 / 4, &strided_bytes); + printf(" host memory: %.1f MiB with a snapshot per chunk, %.1f MiB checkpointed\n", full_bytes / 1048576.0, strided_bytes / 1048576.0); + if (!(strided_bytes < full_bytes) || !same_step(compare("all labelled: checkpointed state vs a snapshot per chunk", strided, exact))) { + fprintf(stderr, "FAILED: checkpointing the recurrent state changed the step (or saved nothing)\n"); + ++failures; + } + } } { // 2. context, then a reply (the walk's real case): the exact walk is invariant to how the From 1ef1a29f7b8b826d27d2b4da30246b398185bf8e Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 19:18:31 -0500 Subject: [PATCH 10/13] train: the exact walk's gradient horizon shrinks to fit the device, and 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- include/llama.h | 4 ++++ src/llama-context.cpp | 34 ++++++++++++++++++++++++++++++---- src/llama-context.h | 3 ++- tools/server/server-train.cpp | 6 ++++++ 4 files changed, 42 insertions(+), 5 deletions(-) diff --git a/include/llama.h b/include/llama.h index d7972a233f12..f128e2b6ba7e 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1765,6 +1765,10 @@ extern "C" { // walk is not exact): what a run of this shape needs, from the run's own arithmetic. LLAMA_API size_t llama_opt_walk_host_bytes(struct llama_context * lctx); + // The gradient horizon the exact walk's last window trained at, in positions (0 = the whole + // window): the requested one, or smaller where a chunk's graph did not fit the device. + LLAMA_API uint32_t llama_opt_walk_horizon(struct llama_context * lctx); + LLAMA_API void llama_opt_epoch( struct llama_context * lctx, ggml_opt_dataset_t dataset, diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 1f4011260277..f8dba12068db 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3919,12 +3919,34 @@ void llama_context::opt_epoch_iter( } } cparams.walk_exact = true; - cparams.walk_grad_from = c.grad_from; cparams.walk_surrogate = c.surrogate; cparams.walk_state_surrogate = recr != nullptr && j < (int64_t) chunks.size() - 1 && !walk_gstate.empty(); - xc = &c; - ok = train_chunk(c.c0, c.n); - xc = nullptr; + // THE HORIZON FITS THE DEVICE. Its GRAD leaves are the chunk graph's largest growth + // with the window (layers x horizon x (k + v) floats, twice with their gradient); the + // graph preflight refuses a chunk that does not fit before allocating anything, so a + // memory refusal halves the horizon (whole chunks) and the chunk is tried again; every + // later chunk keeps the smaller one. The forward still attends to the whole window and + // the 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. + exact_chunk cc = c; + for (;;) { + cc.grad_from = opt_walk_horizon == 0 || cc.c0 < opt_walk_horizon ? 0 : cc.c0 - opt_walk_horizon; + cparams.walk_grad_from = cc.grad_from; + xc = &cc; + ok = train_chunk(cc.c0, cc.n); + xc = nullptr; + const uint32_t span = cc.c0 - cc.grad_from; + if (ok || !opt_failure.empty() || !opt_alloc_failed.load() || span <= n_batch) { + break; + } + opt_walk_horizon = std::max(n_batch, (span/2)/n_batch*n_batch); + LLAMA_LOG_WARN("%s: the chunk at %u did not fit with a gradient horizon of %u positions: retrying at %u\n", + __func__, cc.c0, span, opt_walk_horizon); + opt_alloc_failed.store(false); + opt_stop_requested.store(false); + } + opt_walk_horizon_used = opt_walk_horizon; cparams.walk_exact = false; cparams.walk_grad_from = 0; cparams.walk_surrogate = false; @@ -4763,6 +4785,10 @@ size_t llama_opt_walk_host_bytes(struct llama_context * ctx) { return ctx->opt_walk_host_bytes; } +uint32_t llama_opt_walk_horizon(struct llama_context * ctx) { + return ctx->opt_walk_horizon_used; +} + void llama_opt_set_memory_budget(struct llama_context * ctx, size_t bytes) { ctx->opt_memory_budget = bytes; } diff --git a/src/llama-context.h b/src/llama-context.h index b5b9f7bdc939..95449131eebc 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -211,6 +211,7 @@ struct llama_context { // the exact walk's host memory per window (llama_opt_set_walk_host_budget); 0 = no cap size_t opt_walk_host_budget = 0; size_t opt_walk_host_bytes = 0; // the largest window's, measured by its own arithmetic + uint32_t opt_walk_horizon_used = 0; // the gradient horizon the last exact window trained at (0 = all) // the largest training graph the allocation preflight measured on a GPU device (bytes) size_t opt_graph_bytes() const; @@ -373,7 +374,7 @@ struct llama_context { // training context on the same resident weights) uint32_t opt_n_ctx_train = 0; bool opt_walk_exact = false; // llama_opt_params::walk_exact - uint32_t opt_walk_horizon = 0; // llama_opt_params::walk_horizon + uint32_t opt_walk_horizon = 0; // llama_opt_params::walk_horizon; shrinks to fit the device ggml_threadpool_t threadpool = nullptr; ggml_threadpool_t threadpool_batch = nullptr; diff --git a/tools/server/server-train.cpp b/tools/server/server-train.cpp index c9f91586af78..9328d433ff17 100644 --- a/tools/server/server-train.cpp +++ b/tools/server/server-train.cpp @@ -979,6 +979,7 @@ void server_trainer::run(json req, examples_data ex) { const std::string refusal = llama_opt_failure(ctx); const double graph_mib = llama_opt_graph_bytes(ctx) / 1048576.0; const double walk_host_mib = llama_opt_walk_host_bytes(ctx) / 1048576.0; + const int64_t walk_horizon = llama_opt_walk_horizon(ctx); const int32_t saved = (cancelled || no_fit) ? 0 : llama_adapter_lora_save(adapter, out.c_str()); llama_adapter_lora_free(adapter); llama_free(ctx); @@ -989,6 +990,11 @@ void server_trainer::run(json req, examples_data ex) { state["graph_mib"] = graph_mib; // what the exact walk kept on the host per window (its K/V gradient and state snapshots) state["walk_host_mib"] = walk_host_mib; + // the exact walk's gradient horizon as trained (0 = the whole window), smaller where the device + // could not hold the requested one + if (req.value("exact", false)) { + state["walk_horizon"] = walk_horizon; + } if (no_fit) { state["state"] = "error"; state["error"] = !refusal.empty() ? refusal + " (window " + std::to_string(window) + ")" : From 9f7e57511bab1177074d109513cd2d2775488bea Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 19:35:30 -0500 Subject: [PATCH 11/13] tests: test-grad-leaf and test-opt-dynamic-accum take the CPU backend by type (CI's backend-DL build) Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- tests/test-grad-leaf.cpp | 36 +++++++++++++++++++++++--------- tests/test-opt-dynamic-accum.cpp | 12 ++++++++--- 2 files changed, 35 insertions(+), 13 deletions(-) diff --git a/tests/test-grad-leaf.cpp b/tests/test-grad-leaf.cpp index fa34e1b66959..82d09e576e03 100644 --- a/tests/test-grad-leaf.cpp +++ b/tests/test-grad-leaf.cpp @@ -4,7 +4,8 @@ // flag must stay a constant with no gradient. #include "ggml.h" -#include "ggml-cpu.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" #include #include @@ -13,18 +14,19 @@ #define CHECK(cond) do { if (!(cond)) { fprintf(stderr, "FAILED %s:%d: %s\n", __FILE__, __LINE__, #cond); exit(1); } } while (0) int main() { - ggml_init_params params = { 64u * 1024 * 1024, nullptr, false }; + // by type, not a CPU-backend symbol: in a dynamically loaded backend build (CI) the CPU + // backend is its own library + ggml_backend_load_all(); + ggml_backend_t cpu = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + CHECK(cpu != nullptr); + + ggml_init_params params = { 64 * ggml_tensor_overhead() + ggml_graph_overhead_custom(GGML_DEFAULT_GRAPH_SIZE, true) * 2, nullptr, true }; ggml_context * ctx = ggml_init(params); const int n = 4; ggml_tensor * x = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // the GRAD leaf ggml_tensor * w = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // a real parameter ggml_tensor * c = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, n); // a plain constant - for (int i = 0; i < n; ++i) { - ggml_set_f32_1d(x, i, 0.5f + i); - ggml_set_f32_1d(w, i, 1.0f - 0.25f * i); - ggml_set_f32_1d(c, i, 2.0f); - } ggml_set_grad(x); ggml_set_param(w); @@ -37,8 +39,17 @@ int main() { ggml_build_forward_expand(gf, loss); ggml_cgraph * gb = ggml_graph_dup(ctx, gf, true); ggml_build_backward_expand(ctx, gb, nullptr); + + ggml_backend_buffer_t buf = ggml_backend_alloc_ctx_tensors(ctx, cpu); + CHECK(buf != nullptr); + for (int i = 0; i < n; ++i) { + const float xi = 0.5f + i, wi = 1.0f - 0.25f * i, ci = 2.0f; + ggml_backend_tensor_set(x, &xi, i * sizeof(float), sizeof(float)); + ggml_backend_tensor_set(w, &wi, i * sizeof(float), sizeof(float)); + ggml_backend_tensor_set(c, &ci, i * sizeof(float), sizeof(float)); + } ggml_graph_reset(gb); // loss grad = 1, every other grad = 0 - ggml_graph_compute_with_ctx(ctx, gb, 1); + CHECK(ggml_backend_graph_compute(cpu, gb) == GGML_STATUS_SUCCESS); ggml_tensor * gx = ggml_graph_get_grad(gb, x); ggml_tensor * gw = ggml_graph_get_grad(gb, w); @@ -47,8 +58,11 @@ int main() { CHECK(ggml_graph_get_grad(gb, c) == nullptr && "an unflagged constant gets none"); for (int i = 0; i < n; ++i) { const float xi = 0.5f + i, wi = 1.0f - 0.25f * i, yi = wi * xi + 2.0f; - CHECK(std::fabs(ggml_get_f32_1d(gx, i) - 2.0f * yi * wi) < 1e-5f); - CHECK(std::fabs(ggml_get_f32_1d(gw, i) - 2.0f * yi * xi) < 1e-5f); + float gxi, gwi; + ggml_backend_tensor_get(gx, &gxi, i * sizeof(float), sizeof(float)); + ggml_backend_tensor_get(gw, &gwi, i * sizeof(float), sizeof(float)); + CHECK(std::fabs(gxi - 2.0f * yi * wi) < 1e-5f); + CHECK(std::fabs(gwi - 2.0f * yi * xi) < 1e-5f); } CHECK(!(x->flags & GGML_TENSOR_FLAG_PARAM) && "the GRAD leaf is not a parameter"); // regression: GRAD first shared COMPUTE's bit, so every computed node read as a GRAD leaf and @@ -58,7 +72,9 @@ int main() { CHECK((node == x || !(node->flags & GGML_TENSOR_FLAG_GRAD)) && "only the flagged leaf is a GRAD leaf"); } + ggml_backend_buffer_free(buf); ggml_free(ctx); + ggml_backend_free(cpu); printf("test-grad-leaf: OK\n"); return 0; } diff --git a/tests/test-opt-dynamic-accum.cpp b/tests/test-opt-dynamic-accum.cpp index 191c9e78f6ce..ff5445fda605 100644 --- a/tests/test-opt-dynamic-accum.cpp +++ b/tests/test-opt-dynamic-accum.cpp @@ -7,7 +7,6 @@ #include "ggml.h" #include "ggml-alloc.h" #include "ggml-backend.h" -#include "ggml-cpu.h" #include "ggml-opt.h" #include @@ -27,7 +26,10 @@ static ggml_opt_optimizer_params sgd_pars(void *) { } static void run(int32_t opt_period) { - ggml_backend_t cpu = ggml_backend_cpu_init(); + // by type, not ggml_backend_cpu_init: in a dynamically loaded backend build (CI) the CPU + // backend is its own library + ggml_backend_t cpu = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + GGML_ASSERT(cpu != nullptr); ggml_backend_t backends[] = { cpu }; ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 1, GGML_DEFAULT_GRAPH_SIZE, false, true); @@ -91,7 +93,10 @@ static void run(int32_t opt_period) { // weight times w, plus 3 w in the second) must read back after its eval; a following period // starts from zero. static void run_manual() { - ggml_backend_t cpu = ggml_backend_cpu_init(); + // by type, not ggml_backend_cpu_init: in a dynamically loaded backend build (CI) the CPU + // backend is its own library + ggml_backend_t cpu = ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr); + GGML_ASSERT(cpu != nullptr); ggml_backend_t backends[] = { cpu }; ggml_backend_sched_t sched = ggml_backend_sched_new(backends, nullptr, 1, GGML_DEFAULT_GRAPH_SIZE, false, true); @@ -177,6 +182,7 @@ static void run_manual() { } int main() { + ggml_backend_load_all(); run(1); run(2); run_manual(); From 477e7cdcf4a4412df97db3a91956700d2d1dbf25 Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 20:02:34 -0500 Subject: [PATCH 12/13] examples: the finetune examples name the exact walk's fields (off) in 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- examples/training/finetune-lora.cpp | 2 ++ examples/training/finetune.cpp | 2 ++ 2 files changed, 4 insertions(+) diff --git a/examples/training/finetune-lora.cpp b/examples/training/finetune-lora.cpp index 5244b45f6da6..1f97b2bcde03 100644 --- a/examples/training/finetune-lora.cpp +++ b/examples/training/finetune-lora.cpp @@ -134,6 +134,8 @@ int main(int argc, char ** argv) { /*optimizer_type =*/params.optimizer, /*adapter =*/adapter, /*recompute =*/false, // keep every layer's intermediates, as before + /*walk_exact =*/false, // the per-chunk walk, as before + /*walk_horizon =*/0, }; llama_opt_init(ctx, model, lopt_params); diff --git a/examples/training/finetune.cpp b/examples/training/finetune.cpp index 55d3c9dc9558..c19e5f82dd19 100644 --- a/examples/training/finetune.cpp +++ b/examples/training/finetune.cpp @@ -75,6 +75,8 @@ int main(int argc, char ** argv) { /*optimizer_type =*/params.optimizer, /*adapter =*/nullptr, // full-model training: every tensor the filter allows /*recompute =*/false, // keep every layer's intermediates, as before + /*walk_exact =*/false, // the per-chunk walk, as before + /*walk_horizon =*/0, }; llama_opt_init(ctx, model, lopt_params); From 7b90bd163b15e8fbe17802f8ed1b351c13565e1b Mon Sep 17 00:00:00 2001 From: Joel Teply Date: Tue, 6 Oct 2026 20:28:39 -0500 Subject: [PATCH 13/13] train: the exact walk resets its horizon to the request each window; 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 Claude-Session: https://claude.ai/code/session_01Q4NU4VNiELPQfBpCacDZGc --- src/llama-context.cpp | 12 +++++++++++- src/llama-context.h | 3 ++- src/llama-graph.cpp | 5 +++++ tests/test-walk-exact.cpp | 36 +++++++++++++++++++++++++++++++++--- 4 files changed, 51 insertions(+), 5 deletions(-) diff --git a/src/llama-context.cpp b/src/llama-context.cpp index f8dba12068db..2ffe62e3d52b 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -3336,7 +3336,8 @@ void llama_context::opt_init(struct llama_model * model, struct llama_opt_params // be running on the same weights). opt_n_ctx_train = lopt_params.n_ctx_train > 0 ? lopt_params.n_ctx_train : n_ctx(); opt_walk_exact = lopt_params.walk_exact; - opt_walk_horizon = lopt_params.walk_horizon; + opt_walk_horizon_req = lopt_params.walk_horizon; + opt_walk_horizon = lopt_params.walk_horizon; const uint32_t n_batch = std::min(this->n_batch(), opt_n_ctx_train); const uint32_t n_ubatch = std::min(this->n_ubatch(), n_batch); GGML_ASSERT(opt_n_ctx_train % n_batch == 0); @@ -3748,6 +3749,15 @@ void llama_context::opt_epoch_iter( // prefix's G. ONE optimizer step per window, on the gradient of the window's loss // (exact within the horizon). Memory stays chunk x window; the cost is one decode plus // one forward and backward per chunk, context chunks included. + // THE INVARIANT THAT MAKES THE REVERSE PASS VALID (Fable): the adapter does not change + // between the forward decode and the last chunk's backward. One optimizer step per window, + // taken by the chunk the reverse pass trains last, is what keeps each chunk's recomputed + // K/V (and recurrent state) equal to the cached values its successors attended to; a step + // inside the window would make every surrogate a gradient of a different function. + // + // Each window starts from the REQUESTED horizon: a halving that fit one window's device + // budget must not shrink every later window and epoch, which may be shorter (Fable). + opt_walk_horizon = opt_walk_horizon_req; // A recurrent model's state is carried the same way, one chunk at a time: the state // entering a chunk is a GRAD leaf, the state it leaves meets the next chunk's gradient // (build_rs, build_walk_state_exit). The reverse pass restores each chunk's entry state diff --git a/src/llama-context.h b/src/llama-context.h index 95449131eebc..1cd3e75d2cee 100644 --- a/src/llama-context.h +++ b/src/llama-context.h @@ -374,7 +374,8 @@ struct llama_context { // training context on the same resident weights) uint32_t opt_n_ctx_train = 0; bool opt_walk_exact = false; // llama_opt_params::walk_exact - uint32_t opt_walk_horizon = 0; // llama_opt_params::walk_horizon; shrinks to fit the device + uint32_t opt_walk_horizon_req = 0; // llama_opt_params::walk_horizon, as requested + uint32_t opt_walk_horizon = 0; // this window's: starts at the request, shrinks to fit the device ggml_threadpool_t threadpool = nullptr; ggml_threadpool_t threadpool_batch = nullptr; diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index f87e4bc54e5a..698ec24bdaa7 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2880,6 +2880,11 @@ ggml_tensor * llm_graph_context::build_attn( // K/V meet the gradient later chunks put on them through the surrogate // + : its gradient on k_cur / v_cur IS gk / gv, which the // backward carries through this chunk into the adapter and into its own prefix. + // At a quantized or f16 cache the prefix's gradient is taken at the cache's values and + // applied to the chunk's own full-precision K/V: straight-through. Measured + // (test-walk-exact, q8_0 and f16 caches): cosine 0.998 against one graph, the same as + // at an F32 cache, so training reads the numbers serving runs on (Joel: align the bit + // depth to inference). const int64_t g0 = std::min(cparams.walk_grad_from, n_past); llm_graph_result::walk_layer io = { il, (uint32_t) g0, (uint32_t) n_past }; // [const prefix][gradient prefix][this chunk]: the order the cells hold them in diff --git a/tests/test-walk-exact.cpp b/tests/test-walk-exact.cpp index 5068c9b84db0..958c14ae0063 100644 --- a/tests/test-walk-exact.cpp +++ b/tests/test-walk-exact.cpp @@ -37,6 +37,11 @@ static ggml_opt_optimizer_params sgd_pars(void *) { return p; } +// the cache the walk reads its prefix from: F32 for the exactness checks; the served types (f16, +// q8_0, in a flash-attention context that stores V untransposed) measure the straight-through +// approximation the walk makes there (Fable on #47: the gradient is taken at the cache's values) +static ggml_type g_cache_type = GGML_TYPE_F32; + 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; @@ -44,9 +49,11 @@ static llama_context * make_ctx(const common_params & params, llama_model * mode cparams.n_ubatch = chunk; cparams.n_seq_max = 2; // the plain walk snapshots a recurrent state into a scratch sequence cparams.kv_unified = true; // ...and the window keeps every cell - cparams.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_DISABLED; - cparams.type_k = GGML_TYPE_F32; // the walk's cached constants equal the chunk's own K/V - cparams.type_v = GGML_TYPE_F32; + // a quantized V cache needs flash attention (it stores V untransposed); training graphs take + // the explicit path either way + cparams.flash_attn_type = g_cache_type == GGML_TYPE_F32 ? LLAMA_FLASH_ATTN_TYPE_DISABLED : LLAMA_FLASH_ATTN_TYPE_ENABLED; + cparams.type_k = g_cache_type; + cparams.type_v = g_cache_type; return llama_init_from_model(model, cparams); } @@ -138,6 +145,11 @@ static agreement compare(const char * what, const std::vector & a, const // Measured on the 1.5B Q4_K_M (5090): the exact walk agrees with one graph at cosine 0.998 and // norm 0.995 (what is left is the quantized matmuls' batch-size variance: a 128-row chunk and a // 512-row graph take different kernels); the plain walk sits at cosine 0.81, norm 2.1. +// the served cache types' bar (see 1b). Measured on the 5090: f16 0.9986 (1.5B) / 0.9981 (hybrid), +// q8_0 0.9978 / 0.9979, against F32's 0.998: the straight-through approximation costs nothing +// measurable, so the bar sits just under F32's own agreement. +static const double STRAIGHT_THROUGH_COSINE = 0.99; + static bool same_step(const agreement & r) { return r.cosine > 0.995 && std::fabs(r.norm_ratio - 1.0) < 0.02; } @@ -201,6 +213,24 @@ int main(int argc, char ** argv) { } } } + { + // 1b. the served cache types: the walk reads the prefix at the cache's precision and its + // gradient is taken there, then applied to the chunk's own full-precision K/V (straight- + // through). Measured against the same one-graph reference, which never reads the cache. + std::vector all(WINDOW + 1, 1); + all[0] = 0; + const auto ref = step_delta(params, model, init, tokens, all, WINDOW, false); + for (ggml_type t : { GGML_TYPE_F16, GGML_TYPE_Q8_0 }) { + g_cache_type = t; + const auto exact = step_delta(params, model, init, tokens, all, WINDOW / 4, true); + g_cache_type = GGML_TYPE_F32; + const agreement r = compare((std::string("all labelled, ") + ggml_type_name(t) + " cache: exact walk in 4 chunks vs one graph").c_str(), exact, ref); + if (!(r.cosine > STRAIGHT_THROUGH_COSINE)) { + fprintf(stderr, "FAILED: at a %s cache the exact walk's step strays from the window's gradient (cosine %.4f)\n", ggml_type_name(t), r.cosine); + ++failures; + } + } + } { // 2. context, then a reply (the walk's real case): the exact walk is invariant to how the // window is chunked, context chunks included (they train through the surrogate alone)