#include "models.h"
+#include "llama-impl.h"
#include "llama-kv-cache.h"
#include "llama-kv-cache-iswa.h"
const auto * kv = is_swa ? inp_attn_iswa->mctx->get_swa() : inp_attn_iswa->mctx->get_base();
ggml_tensor * k_idxs = is_swa ? inp_attn_iswa->get_k_idxs_swa() : inp_attn_iswa->get_k_idxs();
ggml_tensor * v_idxs = is_swa ? inp_attn_iswa->get_v_idxs_swa() : inp_attn_iswa->get_v_idxs();
+ // rotate K/V into the cache's rotated space
+ ggml_tensor * k_rot = is_swa ? inp_attn_iswa->self_k_rot_swa : inp_attn_iswa->self_k_rot;
+ ggml_tensor * v_rot = is_swa ? inp_attn_iswa->self_v_rot_swa : inp_attn_iswa->self_v_rot;
+ if (k_rot) {
+ Kcur = llama_mul_mat_hadamard(ctx0, Kcur, k_rot);
+ }
+ if (v_rot) {
+ Vcur = llama_mul_mat_hadamard(ctx0, Vcur, v_rot);
+ }
ggml_build_forward_expand(gf, kv->cpy_k(ctx0, Kcur, k_idxs, il));
ggml_build_forward_expand(gf, kv->cpy_v(ctx0, Vcur, v_idxs, il));
} else {
+ // rotate K/V into the cache's rotated space
+ if (inp_attn->self_k_rot) {
+ Kcur = llama_mul_mat_hadamard(ctx0, Kcur, inp_attn->self_k_rot);
+ }
+ if (inp_attn->self_v_rot) {
+ Vcur = llama_mul_mat_hadamard(ctx0, Vcur, inp_attn->self_v_rot);
+ }
ggml_build_forward_expand(gf, inp_attn->mctx->cpy_k(ctx0, Kcur, inp_attn->get_k_idxs(), il));
ggml_build_forward_expand(gf, inp_attn->mctx->cpy_v(ctx0, Vcur, inp_attn->get_v_idxs(), il));
}