diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 87bb3ccaee..ba89ca794a 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 b592de18c7..004fec5a6d 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 embeddings_layer_inp; // [n_layer()] extract input embeddings for layer diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index d398ee1d76..2df063f955 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);