}
void llm_graph_input_attn_kv_iswa::set_input(const llama_ubatch * ubatch) {
- mctx->get_base()->set_input_k_idxs(self_k_idxs, ubatch);
- mctx->get_base()->set_input_v_idxs(self_v_idxs, ubatch);
+ // base tensors may not be allocated if there are no non-SWA attention layers
+ if (self_k_idxs && self_k_idxs->buffer) {
+ mctx->get_base()->set_input_k_idxs(self_k_idxs, ubatch);
+ mctx->get_base()->set_input_v_idxs(self_v_idxs, ubatch);
- mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);
+ mctx->get_base()->set_input_kq_mask(self_kq_mask, ubatch, cparams.causal_attn);
+ }
- mctx->get_swa()->set_input_k_idxs(self_k_idxs_swa, ubatch);
- mctx->get_swa()->set_input_v_idxs(self_v_idxs_swa, ubatch);
+ // swa tensors may not be allocated if there are no SWA attention layers
+ if (self_k_idxs_swa && self_k_idxs_swa->buffer) {
+ mctx->get_swa()->set_input_k_idxs(self_k_idxs_swa, ubatch);
+ mctx->get_swa()->set_input_v_idxs(self_v_idxs_swa, ubatch);
- mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);
+ mctx->get_swa()->set_input_kq_mask(self_kq_mask_swa, ubatch, cparams.causal_attn);
+ }
if (self_k_rot) {
mctx->get_base()->set_input_k_rot(self_k_rot);
bool res = true;
- res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;
- //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there
+ // base tensors may not be allocated if there are no non-SWA attention layers
+ if (self_k_idxs && self_k_idxs->buffer) {
+ res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;
+ //res &= self_v_idxs->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there
- res &= self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;
- //res &= self_v_idxs_swa->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there
+ res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams);
+ }
- res &= can_reuse_kq_mask(self_kq_mask, mctx->get_base(), params.ubatch, params.cparams);
- res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams);
+ // swa tensors may not be allocated if there are no SWA attention layers
+ if (self_k_idxs_swa && self_k_idxs_swa->buffer) {
+ res &= self_k_idxs_swa->ne[0] == params.ubatch.n_tokens;
+ //res &= self_v_idxs_swa->ne[0] == params.ubatch.n_tokens; // TODO: need to move this to the unified cache and check there
+
+ res &= can_reuse_kq_mask(self_kq_mask_swa, mctx->get_swa(), params.ubatch, params.cparams);
+ }
return res;
}