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