flashinfer.cascade.merge_states

flashinfer.cascade.merge_states(v: Tensor, s: Tensor) Tuple[Tensor, Tensor]

合并多个注意力状态 (v, s)。

参数:
  • v (torch.Tensor) – 来自 KV 分段的注意力输出,形状:[seq_len, num_states, num_heads, head_dim]

  • s (torch.Tensor) – 来自 KV 分段的 logsumexp 值,形状:[seq_len, num_states, num_heads],预期为 float32 张量。

返回值:

  • V (torch.Tensor) – 合并后的注意力输出,形状:[seq_len, num_heads, head_dim]

  • S (torch.Tensor) – 合并后的 KV 分段的 logsumexp 值,形状:[seq_len, num_heads]

示例

>>> import torch
>>> import flashinfer
>>> seq_len = 2048
>>> num_heads = 32
>>> head_dim = 128
>>> num_states = 100
>>> v = torch.randn(seq_len, num_states, num_heads, head_dim).half().to("cuda:0")
>>> s = torch.randn(seq_len, num_states, num_heads, dtype=torch.float32).to("cuda:0")
>>> v_merged, s_merged = flashinfer.merge_states(v, s)
>>> v_merged.shape
torch.Size([2048, 32, 128])
>>> s_merged.shape
torch.Size([2048, 32])