Commit 81ff93ea1 for llama.cpp

commit 81ff93ea1d48c0482508d30e75b814df99c25bf0
Author: Xuan-Son Nguyen <son@huggingface.co>
Date:   Wed Sep 30 18:09:08 2026 +0200

    llama: properly handle KV on training (#28520)

    * llama: properly handle KV on training

    * improve

diff --git a/src/llama-context.cpp b/src/llama-context.cpp
index 87bb3ccae..ba89ca794 100644
--- a/src/llama-context.cpp
+++ b/src/llama-context.cpp
@@ -275,6 +275,7 @@ llama_context::llama_context(

     // initialized later
     cparams.pipeline_parallel = false;
+    cparams.training = false;

     {
         const char * LLAMA_GRAPH_REUSE_DISABLE = getenv("LLAMA_GRAPH_REUSE_DISABLE");
@@ -691,7 +692,11 @@ void llama_context::sched_reserve() {
     }

     // reserve with tg (token generation) graph to get the number of splits and nodes
-    {
+    if (cparams.training) {
+        // no tg graph for training
+        n_splits_tg = n_splits_pp;
+        n_nodes_tg  = n_nodes_pp;
+    } else {
         auto * gf = graph_reserve(n_seqs, n_seqs, n_seqs, mctx.get(), model.hparams.no_alloc);
         if (!gf) {
             throw std::runtime_error("failed to allocate compute tg buffers");
@@ -2410,6 +2415,11 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const {
     if (n_sampling_outputs_max > 1) {
         res += (n_sampling_outputs_max - 1) * n_sampling_nodes_max;
     }
+
+    if (cparams.training) {
+        res *= 4;
+    }
+
     return res;
 }

@@ -3500,12 +3510,19 @@ void llama_context::opt_init(struct llama_model * model, struct llama_opt_params
     if (cparams.flash_attn) {
         LLAMA_LOG_INFO("%s: disabling flash attention, FLASH_ATTN_EXT has no backward pass\n", __func__);
         cparams.flash_attn = false;
+    }

-        // the graph changes without flash attention, need to reserve again
-        sched_need_reserve = true;
-        sched_reserve();
+    // gradients cannot flow through the KV cache, so the attention reads the K and V of the current ubatch directly
+    if (n_ubatch == cparams.n_ctx) {
+        cparams.training = true;
+    } else {
+        LLAMA_LOG_WARN("%s: n_ubatch (%u) != n_ctx (%u), the K and V projections will not receive gradients\n", __func__, n_ubatch, cparams.n_ctx);
     }

+    // the training graph is different, need to reserve again
+    sched_need_reserve = true;
+    sched_reserve();
+
     ggml_opt_params opt_params = ggml_opt_default_params(sched.get(), GGML_OPT_LOSS_TYPE_CROSS_ENTROPY);
     opt_params.opt_period      = n_batch / n_ubatch;
     opt_params.get_opt_pars    = lopt_params.get_opt_pars;
diff --git a/src/llama-cparams.h b/src/llama-cparams.h
index b592de18c..004fec5a6 100644
--- a/src/llama-cparams.h
+++ b/src/llama-cparams.h
@@ -53,6 +53,7 @@ struct llama_cparams {
     bool op_offload;
     bool kv_unified;
     bool pipeline_parallel;
+    bool training;           // set by llama_opt_init()

     std::vector<bool> embeddings_layer_inp; // [n_layer()] extract input embeddings for layer

diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp
index d398ee1d7..2df063f95 100644
--- a/src/llama-graph.cpp
+++ b/src/llama-graph.cpp
@@ -468,8 +468,13 @@ void llm_graph_input_attn_no_cache::set_input(const llama_ubatch * ubatch) {
 }

 void llm_graph_input_attn_kv::set_input(const llama_ubatch * ubatch) {
-    mctx->set_input_k_idxs(self_k_idxs, ubatch);
-    mctx->set_input_v_idxs(self_v_idxs, ubatch);
+    // the idxs are left unallocated when the KV cache is bypassed during training
+    if (self_k_idxs && self_k_idxs->buffer) {
+        mctx->set_input_k_idxs(self_k_idxs, ubatch);
+    }
+    if (self_v_idxs && self_v_idxs->buffer) {
+        mctx->set_input_v_idxs(self_v_idxs, ubatch);
+    }

     // the mask is left unallocated when the graph only stores K/V without attending
     // (e.g. DFlash's KV-injection pass)
@@ -2893,21 +2898,30 @@ ggml_tensor * llm_graph_context::build_attn(

     const auto * mctx_cur = inp->mctx;

-    // store to KV cache
-    {
-        const auto & k_idxs = inp->get_k_idxs();
-        const auto & v_idxs = inp->get_v_idxs();
+    ggml_tensor * q = q_cur;
+    ggml_tensor * k;
+    ggml_tensor * v;

-        ggml_build_forward_expand(gf, mctx_cur->cpy_k(ctx0, k_cur, k_idxs, il));
-        ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, v_cur, v_idxs, il));
+    if (cparams.training) {
+        GGML_ASSERT(mctx_cur->get_n_kv() == n_tokens);
+
+        k = k_cur;
+        v = v_cur;
+    } else {
+        {
+            const auto & k_idxs = inp->get_k_idxs();
+            const auto & v_idxs = inp->get_v_idxs();
+
+            ggml_build_forward_expand(gf, mctx_cur->cpy_k(ctx0, k_cur, k_idxs, il));
+            ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, v_cur, v_idxs, il));
+        }
+
+        k = mctx_cur->get_k(ctx0, il);
+        v = mctx_cur->get_v(ctx0, il);
     }

     ggml_tensor * kq_mask = inp->get_kq_mask();

-    ggml_tensor * q = q_cur;
-    ggml_tensor * k = mctx_cur->get_k(ctx0, il);
-    ggml_tensor * v = mctx_cur->get_v(ctx0, il);
-
     ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask, sinks, v_mla, 0, kq_scale, il);
     cb(cur, "kqv_out", il);

@@ -3144,14 +3158,22 @@ ggml_tensor * llm_graph_context::build_attn(

     const auto * mctx_cur = is_swa ? mctx_iswa->get_swa() : mctx_iswa->get_base();

+    // whole seq fits into batch in training mode
+    const bool use_kv_cur = cparams.training && k_cur && v_cur;
+    if (use_kv_cur) {
+        GGML_ASSERT(mctx_cur->get_n_kv() == n_tokens);
+    }
+
+    const bool store_kv = !use_kv_cur || hparams.n_layer_kv_from_start >= 0;
+
     // optionally store to KV cache
-    if (k_cur) {
+    if (store_kv && k_cur) {
         const auto & k_idxs = is_swa ? inp->get_k_idxs_swa() : inp->get_k_idxs();

         ggml_build_forward_expand(gf, mctx_cur->cpy_k(ctx0, k_cur, k_idxs, il));
     }

-    if (v_cur) {
+    if (store_kv && v_cur) {
         const auto & v_idxs = is_swa ? inp->get_v_idxs_swa() : inp->get_v_idxs();

         ggml_build_forward_expand(gf, mctx_cur->cpy_v(ctx0, v_cur, v_idxs, il));
@@ -3160,8 +3182,8 @@ ggml_tensor * llm_graph_context::build_attn(
     const auto & kq_mask = is_swa ? inp->get_kq_mask_swa() : inp->get_kq_mask();

     ggml_tensor * q = q_cur;
-    ggml_tensor * k = mctx_cur->get_k(ctx0, il);
-    ggml_tensor * v = mctx_cur->get_v(ctx0, il);
+    ggml_tensor * k = use_kv_cur ? k_cur : mctx_cur->get_k(ctx0, il);
+    ggml_tensor * v = use_kv_cur ? v_cur : mctx_cur->get_v(ctx0, il);

     ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask, sinks, v_mla, 0, kq_scale, il);
     cb(cur, "kqv_out", il);