flashinfer.sparse

用于块稀疏 flashattention 的内核。

class flashinfer.sparse.BlockSparseAttentionWrapper(float_workspace_buffer: Tensor, backend: str = 'auto')

用于使用块稀疏矩阵作为注意力掩码的注意力计算包装类。块稀疏矩阵的定义可以在 SciPy 的 bsr_matrix 中找到。

此 API 支持任何块大小 (R, C)

示例

>>> import torch
>>> import flashinfer
>>> num_qo_heads = 32
>>> num_kv_heads = 8
>>> head_dim = 128
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> bsr_wrapper = flashinfer.BlockSparseAttentionWrapper(workspace_buffer)
>>> # sparse mask: [[0, 0, 1], [1, 0, 1], [0, 1, 1]]
>>> M = 3
>>> N = 3
>>> indptr = torch.tensor([0, 1, 3, 5], dtype=torch.int32, device="cuda:0")
>>> indices = torch.tensor([2, 0, 2, 1, 2], dtype=torch.int32, device="cuda:0")
>>> bsr_wrapper.plan(
...     indptr,
...     indices,
...     M,
...     N,
...     1, # R(block_rows)=1
...     1, # C(block_columns)=1
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
... )
>>> q = torch.randn((M, num_qo_heads, head_dim), dtype=torch.float16, device="cuda:0")
>>> k = torch.randn((N, num_kv_heads, head_dim), dtype=torch.float16, device="cuda:0")
>>> v = torch.randn((N, num_kv_heads, head_dim), dtype=torch.float16, device="cuda:0")
>>> o = bsr_wrapper.run(q, k, v)
>>> # use dense implementation with attention mask for comparison
>>> mask = torch.tensor([[0, 0, 1], [1, 0, 1], [0, 1, 1]], dtype=torch.bool, device="cuda:0")
>>> o_ref = flashinfer.single_prefill_with_kv_cache(q, k, v, custom_mask=mask)
>>> torch.allclose(o, o_ref)
True
__init__(float_workspace_buffer: Tensor, backend: str = 'auto') None

构造 BlockSparseAttentionWrapper

参数:
  • float_workspace_buffer (torch.Tensor) – 用于在split-k算法中存储中间注意力结果的用户保留的浮点工作区缓冲区。推荐大小为128MB,工作区缓冲区的设备应与输入张量的设备相同。

  • backend (str) – 实现后端,可以是 auto/fa2fa3。默认为 auto。如果设置为 auto,该函数将根据设备架构和内核可用性自动选择后端。

plan(indptr: Tensor, indices: Tensor, M: int, N: int, R: int, C: int, num_qo_heads: int, num_kv_heads: int, head_dim: int, mask: Tensor | None = None, packed_mask: Tensor | None = None, causal: bool = False, pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, logits_soft_cap: float | None = None, sm_scale: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, q_data_type: str | dtype = 'float16', kv_data_type: str | dtype | None = None, o_data_type: str | dtype = 'float16', non_blocking: bool = True) None

为块稀疏注意力创建辅助数据结构。

参数:
  • indptr (torch.Tensor) – 块稀疏矩阵在行维度上的块索引指针,形状为 (MB + 1,),其中 MB 是行维度中的块数。

  • indices (torch.Tensor) – 块稀疏矩阵在列维度上的块索引,形状为 (nnz,),其中 nnz 是非零块的数量。 indices 数组中的元素应小于 NB:列维度中的块数。

  • M (int) – 块稀疏矩阵的行数,MB = ceil_div(M, R)

  • N (int) – 块稀疏矩阵的列数,NB = N // CN 应该可以被 C 整除。

  • R (int) – 每个块的行数。

  • C (int) – 每个块的列数。

  • num_qo_heads (int) – 查询/输出张量中的头数。

  • num_kv_heads (int) – 键/值张量中的头数。

  • head_dim (int) – 每个头的维度。

  • mask (torch.Tensor, optional) – 形状为 (nnz, R, C,) 的掩码张量,其中 nnz 是非零块的数量。 如果每个块都是满的,则不需要提供掩码张量。

  • packed_mask (torch.Tensor, optional) – 1D 压缩掩码张量,如果提供,则 custom_mask 将被忽略。 压缩掩码张量由 flashinfer.quantization.packbits() 生成。

  • causal (bool) – 是否将因果掩码应用于注意力矩阵。 只有在 plan() 中未提供 custom_mask 时,此设置才有效。

  • pos_encoding_mode (str, optional) – 注意力内核内部应用的定位编码,可以是 NONE/ROPE_LLAMA(LLAMA 风格旋转嵌入)/ALIBI。 默认值为 NONE

  • use_fp16_qk_reduction (bool) – 是否使用 fp16 进行 qk 缩减(速度更快,但精度略有损失)。

  • logits_soft_cap (Optional[float]) – 注意力 logits 的软上限值(用于 Gemini、Grok 和 Gemma-2 等),如果未提供,将设置为 0。如果大于 0,logits 将根据公式进行截断:\(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\),其中 \(x\) 是输入 logits。

  • sm_scale (Optional[float]) – softmax 中使用的缩放比例,如果未提供,则设置为 1.0 / sqrt(head_dim)

  • rope_scale (Optional[float]) – RoPE 插值中使用的比例,如果未提供,将设置为 1.0

  • rope_theta (Optional[float]) – RoPE 中使用的 theta 值,如果未提供,则设置为 1e4

  • q_data_type (str, optional) – 查询张量的数据类型。

  • kv_data_type (Optional[Union[str, torch.dtype]]) – key/value 张量的数据类型。如果为 None,则设置为 q_data_type

  • o_data_type (str, optional) – 输出张量的数据类型。 默认值为 half。 由于量化中无法通过输入 dtype 推断输出 dtype

  • non_blocking (bool) – 是否异步将输入张量复制到设备,默认为 True

在调用任何 run()run_return_lse() 调用之前,应调用 plan() 方法,在此调用期间将创建并缓存辅助数据结构以供多次内核运行使用。

num_qo_heads 必须是 num_kv_heads 的倍数。如果 num_qo_heads 不等于 num_kv_heads,该函数将使用 分组查询注意力

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer: Tensor) None

重置工作区缓冲区。

参数:
  • float_workspace_buffer (torch.Tensor) – 新的 float 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

  • int_workspace_buffer (torch.Tensor) – 新的 int 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

run(q: Tensor, k: Tensor, v: Tensor, scale_q: Tensor | None = None, scale_k: Tensor | None = None, scale_v: Tensor | None = None, out: Tensor | None = None, lse: Tensor | None = None, return_lse: bool = False, enable_pdl: bool | None = None) Tensor | Tuple[Tensor, Tensor]

计算 Q/K/V 张量之间的块稀疏注意力。

参数:
  • q (torch.Tensor) – 查询张量,形状为 (M, num_qo_heads, head_dim)

  • k (torch.Tensor) – 键张量,形状为 (N, num_kv_heads, head_dim)

  • v (torch.Tensor) – 值张量,形状为 (N, num_kv_heads, head_dim)

  • scale_q (Optional[torch.Tensor]) – 查询的缩放张量,每头量化,形状:[num_qo_heads]。用于 FP8 量化。如果未提供,将设置为 1.0

  • scale_k (Optional[torch.Tensor]) – 键的缩放张量,每头量化,形状:[num_kv_heads]。用于 FP8 量化。如果未提供,将设置为 1.0

  • scale_v (Optional[torch.Tensor]) – 值的缩放张量,每头量化,形状:[num_kv_heads]。用于 FP8 量化。如果未提供,将设置为 1.0

  • out (Optional[torch.Tensor]) – 输出张量,如果未提供,将在内部分配。

  • lse (Optional[torch.Tensor]) – 注意力 logits 的对数和指数,如果未提供,将在内部分配。

  • return_lse (bool) – 是否返回注意力 logits 的对数和指数

  • enable_pdl (bool) – 是否启用程序依赖启动 (PDL)。请参阅 https://docs.nvda.net.cn/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization 仅支持 >= sm90,并且当前仅支持 FA2 和 CUDA 核心解码。

返回值:

如果 return_lseFalse,则注意力输出,形状:[M, num_qo_heads, head_dim]。如果 return_lseTrue,则返回一个包含两个张量的元组

  • 注意力输出,形状:[M, num_qo_heads, head_dim]

  • 注意力输出的对数和指数,形状:[M, num_qo_heads]

返回值类型:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]

class flashinfer.sparse.VariableBlockSparseAttentionWrapper(float_workspace_buffer: Tensor, backend: str = 'auto')

带有块稀疏矩阵作为注意力掩码的注意力计算的包装类。此 API 支持由 block_row_szblock_col_sz 提供的可变块大小。此外,每个 kv_head_idx 都可以指定自己的稀疏模式,而无需使用相同的掩码。

示例

>>> import torch
>>> import flashinfer
>>> num_qo_heads = 1
>>> num_kv_heads = 1
>>> head_dim = 128
>>> seq_len = 6 # This corresponds to the `block_row_sz` and `block_col_sz`
>>> # allocate 128MB workspace buffer
>>> workspace_buffer = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device="cuda:0")
>>> wrapper = flashinfer.VariableBlockSparseAttentionWrapper(workspace_buffer)
>>> block_mask_map = torch.tensor([[[0, 0, 1], [1, 0, 1], [0, 1, 1]]], dtype=torch.bool, device="cuda:0")
>>> block_row_sz = torch.tensor([[1, 2, 3]], dtype=torch.int32, device="cuda:0")
>>> block_col_sz = torch.tensor([[3, 1, 2]], dtype=torch.int32, device="cuda:0")
>>> wrapper.plan(
...     block_mask_map,
...     block_row_sz,
...     block_col_sz,
...     num_qo_heads,
...     num_kv_heads,
...     head_dim,
... )
>>> q = torch.randn((num_qo_heads, seq_len, head_dim), dtype=torch.float16, device="cuda:0")
>>> k = torch.randn((num_kv_heads, seq_len, head_dim), dtype=torch.float16, device="cuda:0")
>>> v = torch.randn((num_kv_heads, seq_len, head_dim), dtype=torch.float16, device="cuda:0")
>>> o = wrapper.run(q, k, v)
variable block sparse attention plan function diagram
__init__(float_workspace_buffer: Tensor, backend: str = 'auto') None

VariableBlockSparseAttentionWrapper 的构造函数。

参数:
  • float_workspace_buffer (torch.Tensor) – 用于在split-k算法中存储中间注意力结果的用户保留的浮点工作区缓冲区。推荐大小为128MB,工作区缓冲区的设备应与输入张量的设备相同。

  • backend (str) – 实现后端,可以是 auto/fa2fa3。默认为 auto。如果设置为 auto,该函数将根据设备架构和内核可用性自动选择后端。

plan(block_mask_map: Tensor, block_row_sz: Tensor, block_col_sz: Tensor, num_qo_heads: int, num_kv_heads: int, head_dim: int, causal: bool = False, pos_encoding_mode: str = 'NONE', use_fp16_qk_reduction: bool = False, logits_soft_cap: float | None = None, sm_scale: float | None = None, rope_scale: float | None = None, rope_theta: float | None = None, non_blocking: bool = True, q_data_type: str | dtype = 'float16', kv_data_type: str | dtype | None = None) None

为块稀疏注意力创建辅助数据结构。

参数:
  • block_mask_map (torch.Tensor) – 块掩码映射(布尔型),形状 (num_kv_heads, MB, NB),其中 MB 是行维度的块数,NB 是列维度的块数。

  • block_row_sz (torch.Tensor) – 块行大小,形状 (num_kv_heads, MB,)

  • block_col_sz (torch.Tensor) – 块列大小,形状 (num_kv_heads, NB,)

  • num_qo_heads (int) – 查询/输出张量中的头数。

  • num_kv_heads (int) – key/value 张量中的头数。请注意,qo_heads 的一组共享 kv_heads 的相同稀疏模式。

  • head_dim (int) – 每个头的维度。

  • causal (bool) – 是否将因果掩码应用于注意力矩阵。

  • pos_encoding_mode (str, optional) – 注意力内核内部应用的定位编码,可以是 NONE/ROPE_LLAMA(LLAMA 风格旋转嵌入)/ALIBI。 默认值为 NONE

  • use_fp16_qk_reduction (bool) – 是否使用 fp16 进行 qk 缩减(速度更快,但精度略有损失)。

  • logits_soft_cap (Optional[float]) – 注意力 logits 的软上限值(用于 Gemini、Grok 和 Gemma-2 等),如果未提供,将设置为 0。如果大于 0,logits 将根据公式进行截断:\(\texttt{logits_soft_cap} \times \mathrm{tanh}(x / \texttt{logits_soft_cap})\),其中 \(x\) 是输入 logits。

  • sm_scale (Optional[float]) – softmax 中使用的缩放比例,如果未提供,则设置为 1.0 / sqrt(head_dim)

  • rope_scale (Optional[float]) – RoPE 插值中使用的比例,如果未提供,将设置为 1.0

  • rope_theta (Optional[float]) – RoPE 中使用的 theta 值,如果未提供,则设置为 1e4

  • non_blocking (bool) – 是否异步将输入张量复制到设备,默认为 True

在任何 run()run_return_lse() 调用之前,应调用 plan() 方法,在此调用期间将创建辅助数据结构并缓存用于多次内核运行。

num_qo_heads 必须是 num_kv_heads 的倍数。如果 num_qo_heads 不等于 num_kv_heads,该函数将使用 分组查询注意力

reset_workspace_buffer(float_workspace_buffer: Tensor, int_workspace_buffer: Tensor) None

重置工作区缓冲区。

参数:
  • float_workspace_buffer (torch.Tensor) – 新的 float 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

  • int_workspace_buffer (torch.Tensor) – 新的 int 工作区缓冲区,该缓冲区的设备应与输入张量的设备相同。

run(q: Tensor, k: Tensor, v: Tensor, out: Tensor | None = None, lse: Tensor | None = None, return_lse: bool = False, enable_pdl: bool | None = None) Tensor | Tuple[Tensor, Tensor]

计算 Q/K/V 张量之间的块稀疏注意力。

参数:
  • q (torch.Tensor) – 查询张量,形状 (num_qo_heads, qo_len, head_dim)

  • k (torch.Tensor) – key 张量,形状 (num_kv_heads, kv_len, head_dim)

  • v (torch.Tensor) – value 张量,形状 (num_kv_heads, kv_len, head_dim)

  • out (Optional[torch.Tensor]) – 输出张量,如果未提供,将在内部分配。

  • lse (Optional[torch.Tensor]) – 注意力 logits 的对数和指数,如果未提供,将在内部分配。

  • return_lse (bool) – 是否返回注意力 logits 的对数和指数

  • enable_pdl (bool) – 是否启用程序依赖启动 (PDL)。请参阅 https://docs.nvda.net.cn/cuda/cuda-c-programming-guide/#programmatic-dependent-launch-and-synchronization 仅支持 >= sm90,并且当前仅支持 FA2 和 CUDA 核心解码。

返回值:

如果 return_lseFalse,则注意力输出,形状:[M, num_qo_heads, head_dim]。如果 return_lseTrue,则返回一个包含两个张量的元组

  • 注意力输出,形状:[M, num_qo_heads, head_dim]

  • 注意力输出的对数和指数,形状:[M, num_qo_heads]

返回值类型:

Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]