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