void llm_graph_input_attn_kv_msa::set_input(const llama_ubatch * ubatch) {
llm_graph_input_attn_kv::set_input(ubatch);
- mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch);
+ if (self_k_idxs_idx) {
+ mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch);
+ }
}
bool llm_graph_input_attn_kv_msa::can_reuse(const llm_graph_params & params) {
bool res = true;
- res &= self_k_idxs ->ne[0] == params.ubatch.n_tokens;
- res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens;
+ res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;
+ if (self_k_idxs_idx) {
+ res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens;
+ }
res &= can_reuse_kq_mask(self_kq_mask, this->mctx, params.ubatch, params.cparams);
return (llm_graph_input_attn_k_dsa *) res->add_input(std::move(inp));
}
-llm_graph_input_attn_kv_msa * llm_graph_context::build_attn_inp_kv_msa() const {
+llm_graph_input_attn_kv_msa * llm_graph_context::build_attn_inp_kv_msa(bool msa_enabled) const {
const auto * mctx_cur = static_cast<const llama_kv_cache_msa_context *>(mctx);
auto inp = std::make_unique<llm_graph_input_attn_kv_msa>(hparams, cparams, mctx_cur);
inp->self_k_rot = mctx_base->build_input_k_rot(ctx0);
inp->self_v_rot = mctx_base->build_input_v_rot(ctx0);
- inp->self_k_idxs_idx = mctx_idx->build_input_k_idxs(ctx0, ubatch);
+ if (msa_enabled) {
+ inp->self_k_idxs_idx = mctx_idx->build_input_k_idxs(ctx0, ubatch);
+ }
return (llm_graph_input_attn_kv_msa *) res->add_input(std::move(inp));
}
llm_graph_input_attn_k_dsa * build_attn_inp_k_dsa() const;
- llm_graph_input_attn_kv_msa * build_attn_inp_kv_msa() const;
+ llm_graph_input_attn_kv_msa * build_attn_inp_kv_msa(bool msa_enabled) const;
ggml_tensor * build_attn(
llm_graph_input_attn_k_dsa * inp,
inpL = build_inp_embd(model.tok_embd);
ggml_tensor * inp_pos = build_inp_pos();
- auto inp_attn = build_attn_inp_kv_msa();
+
+ // ==========================================
+ // TODO: avoid such kind of complexity in the model graphs
// MSA calls ggml_flash_attn_ext directly and assumes the non-transposed V layout that
// llama.cpp only provides when flash attention is enabled. Block selection is anchored
const bool streams_ok = cparams.n_seq_max == 1 || !cparams.kv_unified;
const bool msa_enabled = fa_on && streams_ok;
+ auto * inp_attn = build_attn_inp_kv_msa(msa_enabled);
+
static bool warned_no_fa = false;
if (!fa_on && !warned_no_fa) {
LLAMA_LOG_WARN("%s: flash attention disabled; MSA requires it -> running DENSE attention "
"-> running DENSE attention. Output may be degraded. Drop --kv-unified to enable MSA.\n", __func__);
warned_unified = true;
}
+ // ==========================================
// hoisted per-graph MSA state (shared by every sparse layer)
llm_graph_input_msa * msa = nullptr;