uint32_t n_pad,
uint32_t n_swa,
llama_swa_type swa_type,
- const layer_filter_cb & filter,
+ const layer_filter_cb & filter_mla,
+ const layer_filter_cb & filter_lid,
const layer_reuse_cb & reuse) :
hparams_lid(model.hparams), n_stream(unified ? 1 : n_seq_max) {
kv_mla = std::make_unique<llama_kv_cache>(
model, model.hparams, type_k, type_v,
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
- n_swa, swa_type, nullptr, filter, reuse, nullptr);
+ n_swa, swa_type, nullptr, filter_mla, reuse, nullptr);
// we use llama_kv_cache for caching indexer keys
// by hand-tweaking some hparams we fool it to create
kv_lid = std::make_unique<llama_kv_cache>(
model, hparams_lid, type_k, type_v,
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
- n_swa, swa_type, nullptr, filter, reuse, nullptr);
+ n_swa, swa_type, nullptr, filter_lid, reuse, nullptr);
}
void llama_kv_cache_dsa::clear(bool data) {
} else {
// Main context: DSA cache for the trunk layers only - the nextn
// layer(s) are never attended by the trunk graph.
- llama_kv_cache::layer_filter_cb filter = nullptr;
+ llama_kv_cache::layer_filter_cb filter_mla = nullptr;
if (hparams.n_layer_nextn > 0) {
- filter = [&](uint32_t il) { return il < hparams.n_layer(); };
+ filter_mla = [&](uint32_t il) { return il < hparams.n_layer(); };
}
+ llama_kv_cache::layer_filter_cb filter_lid = [&](uint32_t il) { return il < hparams.n_layer() && (arch != LLM_ARCH_GLM_DSA || hparams.is_indexer_full(il)); };
res = new llama_kv_cache_dsa(
*this,
1,
hparams.n_swa,
hparams.swa_type,
- filter,
+ filter_mla,
+ filter_lid,
nullptr);
}
} break;