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