flashinfer.cascade.merge_state

flashinfer.cascade.merge_state(v_a: Tensor, s_a: Tensor, v_b: Tensor, s_b: Tensor) Tuple[Tensor, Tensor]

合并来自两个 KV 分段的注意力输出 V 和 logsumexp 值 S。有关数学细节,请查看 我们的教程

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

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

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

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

返回值:

  • V (torch.Tensor) – 合并后的注意力输出(等效于使用合并的 KV 分段 [A: B] 的注意力),形状:[seq_len, num_heads, head_dim]

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

示例

>>> import torch
>>> import flashinfer
>>> seq_len = 2048
>>> num_heads = 32
>>> head_dim = 128
>>> va = torch.randn(seq_len, num_heads, head_dim).half().to("cuda:0")
>>> sa = torch.randn(seq_len, num_heads, dtype=torch.float32).to("cuda:0")
>>> vb = torch.randn(seq_len, num_heads, head_dim).half().to("cuda:0")
>>> sb = torch.randn(seq_len, num_heads, dtype=torch.float32).to("cuda:0")
>>> v_merged, s_merged = flashinfer.merge_state(va, sa, vb, sb)
>>> v_merged.shape
torch.Size([2048, 32, 128])
>>> s_merged.shape
torch.Size([2048, 32])