flashinfer.page.append_paged_mla_kv_cache¶
- flashinfer.page.append_paged_mla_kv_cache(append_ckv: Tensor, append_kpe: Tensor, batch_indices: Tensor, positions: Tensor, ckv_cache: Tensor | None, kpe_cache: Tensor | None, kv_indices: Tensor, kv_indptr: Tensor, kv_last_page_len: Tensor) None¶
将一批键值对附加到分页的键值缓存,注意:当前仅支持 ckv=512 和 kpe=64
- 参数:
append_ckv (torch.Tensor) – 要附加的压缩 kv 张量,采用稀疏张量格式,形状:
[append_indptr[-1], ckv_dim]。append_kpe (torch.Tensor) – 要附加的值张量,采用稀疏张量格式,形状:
[append_indptr[-1], kpe_dim]。batch_indices (torch.Tensor) – 附加的键值对中每个条目的批次索引,形状:
[append_indptr[-1]]。positions (torch.Tensor) – 附加的键值对中每个条目的位置,形状:
[append_indptr[-1]]。ckv_cache (压缩 kv 的缓存, torch.Tensor, 形状: [page_num, page_size, ckv_dim])
kpe_cache (关键位置嵌入的缓存, torch.Tensor, 形状: [page_num, page_size, kpe_dim])
kv_indices (torch.Tensor) – 分页 kv 缓存的页索引,形状:
[kv_indptr[-1]]。kv_indptr (torch.Tensor) – 分页 kv 缓存的 indptr,形状:
[batch_size + 1]。kv_last_page_len (torch.Tensor) – 分页 kv 缓存中每个请求的最后一页中的条目数,形状:
[batch_size]。