flashinfer.cute_dsl.sparse.bsa_attn_sm120.bsa_attn_sm120_blk64_fwd¶
- flashinfer.cute_dsl.sparse.bsa_attn_sm120.bsa_attn_sm120_blk64_fwd(q: Tensor, k: Tensor, v: Tensor, q2k_block_index: Tensor, block_sparse_num: int, block_sizes: Tensor | None = None, q2k_block_nums: Tensor | None = None, softmax_scale: float | None = None, return_lse: bool = False, out: Tensor | None = None, lse: Tensor | None = None) Tuple[Tensor, Tensor | None]¶
Forward pass for BSA block-sparse attention using the sm120_blk64 CuTe-DSL kernel (SM120/SM121 only).
- Parameters:
q – Query tensor (batch, seqlen_q, num_heads, head_dim), fp16/bf16.
k – Key tensor (batch, seqlen_k, num_kv_heads, head_dim).
v – Value tensor (batch, seqlen_k, num_kv_heads, head_dim).
q2k_block_index – (batch, num_heads, num_q_blocks, max_kv_blocks) int32.
block_sparse_num – Number of KV blocks per Q block. Ignored when q2k_block_nums is provided.
block_sizes – Actual token count per KV block, int32. Shape: (num_kv_blocks,) or (batch, num_kv_blocks) or (batch, num_heads, num_kv_blocks). Pass None to skip per-block padding masking.
q2k_block_nums – Per-(batch, head, q_block) KV block count, (batch, num_heads, num_q_blocks) int32. Optional.
softmax_scale – Softmax scale (default: 1/sqrt(head_dim)).
return_lse – Whether to return log-sum-exp.
out – Pre-allocated output tensor (batch, seqlen_q, num_heads, head_dim).
lse – Pre-allocated LSE tensor (batch, num_heads, seqlen_q).