flashinfer.page.append_paged_kv_cache

flashinfer.page.append_paged_kv_cache(append_key: Tensor, append_value: Tensor, batch_indices: Tensor, positions: Tensor, paged_kv_cache: Tensor | Tuple[Tensor, Tensor], kv_indices: Tensor, kv_indptr: Tensor, kv_last_page_len: Tensor, kv_layout: str = 'NHD') None

将一批键值对附加到分页的键值缓存。

参数:
  • append_key (torch.Tensor) – 要附加的键张量,采用稀疏张量格式,形状为:[append_indptr[-1], num_kv_heads, head_dim]

  • append_value (torch.Tensor) – 要附加的值张量,采用稀疏张量格式,形状为:[append_indptr[-1], num_kv_heads, head_dim]

  • batch_indices (torch.Tensor) – 附加键值对中每个条目的批次索引,形状为:[append_indptr[-1]]

  • positions (torch.Tensor) – 附加键值对中每个条目的位置,形状为:[append_indptr[-1]]

  • paged_kv_cache (Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]) –

    存储的分页 KV 缓存,作为张量元组或单个张量

    • 一个元组 (k_cache, v_cache),包含 4D 张量,每个张量的形状为:[max_num_pages, page_size, num_kv_heads, head_dim],如果 kv_layoutNHD,以及 [max_num_pages, num_kv_heads, page_size, head_dim],如果 kv_layoutHND

    • 单个 5D 张量,形状为:[max_num_pages, 2, page_size, num_kv_heads, head_dim],如果 kv_layoutNHD,形状为 [max_num_pages, 2, num_kv_heads, page_size, head_dim],如果 kv_layoutHND。其中 paged_kv_cache[:, 0] 是键缓存,paged_kv_cache[:, 1] 是值缓存。

  • 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]

  • kv_layout (str) – 分页 kv 缓存的布局,可以是 NHDHND

示例

>>> import torch
>>> import flashinfer
>>> nnz_kv = 100
>>> num_kv_heads = 32
>>> head_dim = 128
>>> k_append = torch.randn(nnz_kv, num_kv_heads, head_dim).half().to(0)
>>> v_append = torch.randn(nnz_kv, num_kv_heads, head_dim).half().to(0)
>>> # 45 + 8 + 25 + 22 = nnz_kv
>>> kv_append_length = torch.tensor([45, 8, 25, 22], dtype=torch.int32, device="cuda:0")
>>> kv_append_indptr = torch.cat(
...     [torch.zeros(1).int().to(0), torch.cumsum(kv_append_length, dim=0)]
... ).int()  # [0, 45, 53, 78, 100]
>>> max_num_pages = 1000
>>> page_size = 16
>>> paged_kv_cache = torch.randn(max_num_pages, 2, page_size, num_kv_heads, head_dim).half().to(0)
>>> num_pages_per_req = torch.tensor([3, 1, 2, 2], dtype=torch.int32, device="cuda:0")
>>> kv_page_indptr = torch.cat(
...     [torch.zeros(1).int().to(0), torch.cumsum(num_pages_per_req, dim=0)]
... ).int()
>>> # use first 8 pages in the paged-kv
>>> kv_page_indices = torch.arange(8, dtype=torch.int32, device="cuda:0")
>>> # 45 = (3 - 1) * 16 + 13
>>> # 8 = (1 - 1) * 16 + 8
>>> # 25 = (2 - 1) * 16 + 9
>>> # 22 = (2 - 1) * 16 + 6
>>> kv_last_page_len = torch.tensor([13, 8, 9, 6], dtype=torch.int32, device="cuda:0")
>>> batch_indices, positions = flashinfer.get_batch_indices_positions(
...     kv_append_indptr, flashinfer.get_seq_lens(kv_page_indptr, kv_last_page_len, page_size), nnz_kv
... )
>>> batch_indices
tensor([0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
        0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1,
        1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
        2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3,
        3, 3, 3, 3], device='cuda:0', dtype=torch.int32)
>>> positions
tensor([ 0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15, 16, 17,
        18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35,
        36, 37, 38, 39, 40, 41, 42, 43, 44,  0,  1,  2,  3,  4,  5,  6,  7,  0,
        1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11, 12, 13, 14, 15, 16, 17, 18,
        19, 20, 21, 22, 23, 24,  0,  1,  2,  3,  4,  5,  6,  7,  8,  9, 10, 11,
        12, 13, 14, 15, 16, 17, 18, 19, 20, 21], device='cuda:0',
    dtype=torch.int32)
>>> flashinfer.append_paged_kv_cache(
...     k_append,
...     v_append,
...     batch_indices,
...     positions,
...     paged_kv_cache,
...     kv_page_indices,
...     kv_page_indptr,
...     kv_last_page_len
... )

注意

该函数假定已分配附加 k/v 的空间,这意味着 kv_indiceskv_indptrkv_last_page_len 已经包含了附加的 k/v。