add Q rot when cache is quantized

This commit is contained in:
Aman Gupta
2026-06-04 18:54:12 +08:00
parent 7b87cd3598
commit 27461cd888
3 changed files with 46 additions and 5 deletions
+29
View File
@@ -623,6 +623,19 @@ bool llm_graph_input_attn_kv_iswa::can_reuse(const llm_graph_params & params) {
void llm_graph_input_attn_src_kv_iswa::set_input(const llama_ubatch * ubatch) {
src_mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);
src_mctx->get_swa() ->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);
if (self_k_rot) {
src_mctx->get_base()->set_input_k_rot(self_k_rot);
}
if (self_v_rot) {
src_mctx->get_base()->set_input_v_rot(self_v_rot);
}
if (self_k_rot_swa) {
src_mctx->get_swa()->set_input_k_rot(self_k_rot_swa);
}
if (self_v_rot_swa) {
src_mctx->get_swa()->set_input_v_rot(self_v_rot_swa);
}
}
bool llm_graph_input_attn_src_kv_iswa::can_reuse(const llm_graph_params & params) {
@@ -2635,6 +2648,11 @@ llm_graph_input_attn_src_kv_iswa * llm_graph_context::build_attn_inp_src_kv_iswa
inp->self_kq_mask_swa = build_attn_inp_kq_mask(ctx0, src_iswa->get_swa(), ubatch, cparams);
inp->self_kq_mask_swa_cnv = cparams.flash_attn ? ggml_cast(ctx0, inp->self_kq_mask_swa, GGML_TYPE_F16) : inp->self_kq_mask_swa;
inp->self_k_rot = src_iswa->get_base()->build_input_k_rot(ctx0);
inp->self_v_rot = src_iswa->get_base()->build_input_v_rot(ctx0);
inp->self_k_rot_swa = src_iswa->get_swa()->build_input_k_rot(ctx0);
inp->self_v_rot_swa = src_iswa->get_swa()->build_input_v_rot(ctx0);
return (llm_graph_input_attn_src_kv_iswa *) res->add_input(std::move(inp));
}
@@ -2657,6 +2675,13 @@ ggml_tensor * llm_graph_context::build_attn(
const auto & kq_mask = is_swa ? inp->get_kq_mask_swa() : inp->get_kq_mask();
auto * k_rot = is_swa ? inp->self_k_rot_swa : inp->self_k_rot;
auto * v_rot = is_swa ? inp->self_v_rot_swa : inp->self_v_rot;
if (k_rot) {
q_cur = ggml_mul_mat_aux(ctx0, q_cur, k_rot);
}
ggml_build_forward_expand(gf, q_cur);
ggml_tensor * q = q_cur;
@@ -2689,6 +2714,10 @@ ggml_tensor * llm_graph_context::build_attn(
ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask, sinks, v_mla, kq_scale, il_assist);
cb(cur, "kqv_out", il_assist);
if (v_rot) {
cur = ggml_mul_mat_aux(ctx0, cur, v_rot);
}
if (wo) {
cur = build_lora_mm(wo, cur, wo_s);
}
+5
View File
@@ -484,6 +484,11 @@ public:
ggml_tensor * self_kq_mask_swa = nullptr;
ggml_tensor * self_kq_mask_swa_cnv = nullptr;
ggml_tensor * self_k_rot = nullptr;
ggml_tensor * self_v_rot = nullptr;
ggml_tensor * self_k_rot_swa = nullptr;
ggml_tensor * self_v_rot_swa = nullptr;
const llama_hparams hparams;
const llama_cparams cparams;
+12 -5
View File
@@ -913,13 +913,23 @@ private:
SRV_INF("loading draft model '%s'\n", params_spec.mparams.path.c_str());
const bool spec_mtp = std::find(params_base.speculative.types.begin(),
params_base.speculative.types.end(),
COMMON_SPECULATIVE_TYPE_DRAFT_MTP) != params_base.speculative.types.end();
auto params_dft = params_base;
params_dft.devices = params_spec.devices;
params_dft.model = params_spec.mparams;
params_dft.n_gpu_layers = params_spec.n_gpu_layers;
params_dft.cache_type_k = params_spec.cache_type_k;
params_dft.cache_type_v = params_spec.cache_type_v;
// TODO: find a better way to expose that the cache is shared
if (spec_mtp) {
params_dft.cache_type_k = params_base.cache_type_k;
params_dft.cache_type_v = params_base.cache_type_v;
} else {
params_dft.cache_type_k = params_spec.cache_type_k;
params_dft.cache_type_v = params_spec.cache_type_v;
}
if (params_spec.cpuparams.n_threads > 0) {
params_dft.cpuparams.n_threads = params_spec.cpuparams.n_threads;
@@ -938,9 +948,6 @@ private:
auto cparams = common_context_params_to_llama(params_dft);
const bool spec_mtp = std::find(params_base.speculative.types.begin(),
params_base.speculative.types.end(),
COMMON_SPECULATIVE_TYPE_DRAFT_MTP) != params_base.speculative.types.end();
if (spec_mtp) {
cparams.ctx_type = LLAMA_CONTEXT_TYPE_MTP;
}