latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
tensorplay.nn.attention API
Functions 11
activate_flash_attention_impl
functionFull reference ↗- tensorplay.nn.attention.activate_flash_attention_impl(impl: str | Literal['FA3', 'FA4']) None
Activate into the dispatcher a previously registered flash attention impl.
Note
Backend providers should NOT automatically activate their implementation on import. Users should explicitly opt-in by calling this function or via environment variables to ensure multiple provider libraries can coexist.
- Parameters:
impl – Implementation identifier to activate. See
list_flash_attention_impls()for available implementations. If the backend’sregister_flash_attention_impl()callable returns aFlashAttentionHandle, the registry keeps that handle alive for the lifetime of the process (until explicit uninstall support exists).
Example
>>> activate_flash_attention_impl("FA4")
and_masks
functionFull reference ↗create_block_mask
functionFull reference ↗- tensorplay.nn.attention.create_block_mask(mask_mod: _mask_mod_signature, B: int | None, H: int | None, Q_LEN: int, KV_LEN: int, device: DeviceLikeType | None = None, BLOCK_SIZE: int | tuple[int, int] = 128, _compile=False, separate_full_blocks: bool = True, compute_dq_write_order: bool = False, dq_kv_order: bool = True) BlockMask[source]
This function creates a block mask tuple from a mask_mod function.
- Parameters:
mask_mod (Callable) – mask_mod function. This is a callable that defines the masking pattern for the attention mechanism. It takes four arguments: b (batch size), h (number of heads), q_idx (query index), and kv_idx (key/value index). It should return a boolean tensor indicating which attention connections are allowed (True) or masked out (False).
B (int) – Batch size.
H (int) – Number of query heads.
Q_LEN (int) – Sequence length of query.
KV_LEN (int) – Sequence length of key/value.
device (str) – Device to run the mask creation on.
BLOCK_SIZE (int or tuple[int, int]) – Block size for the block mask. If a single int is provided it is used for both query and key/value.
separate_full_blocks (bool) – If True, fully unmasked blocks are stored separately so kernels can skip mask_mod on those blocks. If False, all non-empty blocks are stored as partial blocks and mask_mod is applied to every block.
compute_dq_write_order (bool) – If True, precompute dQ write-order metadata needed by deterministic block-sparse backward.
dq_kv_order (bool) – KV-column scheduler order used for deterministic dQ accumulation when compute_dq_write_order is True. False means ascending n-block order and True means descending/SPT order. Explicit tensor schedules are not supported by create_block_mask yet; they are supported by BlockMask.from_kv_blocks for callers that provide precomputed write-order metadata directly.
- Returns:
A BlockMask object that contains the block mask information.
- Return type:
- Example Usage:
def causal_mask(b, h, q_idx, kv_idx): return q_idx >= kv_idx block_mask = create_block_mask(causal_mask, 1, 1, 8192, 8192, device="cuda") query = tensorplay.randn(1, 1, 8192, 64, device="cuda", dtype=tensorplay.float16) key = tensorplay.randn(1, 1, 8192, 64, device="cuda", dtype=tensorplay.float16) value = tensorplay.randn(1, 1, 8192, 64, device="cuda", dtype=tensorplay.float16) output = omni_attention(query, key, value, block_mask=block_mask)
create_mask
functionFull reference ↗- tensorplay.nn.attention.create_mask(mod_fn: _score_mod_signature | _mask_mod_signature, B: int | None, H: int | None, Q_LEN: int, KV_LEN: int, device: DeviceLikeType | None = None) Tensor[source]
This function creates a mask tensor from a mod_fn function.
- Parameters:
- Returns:
A mask tensor with shape (B, H, M, N).
- Return type:
mask (Tensor)
current_flash_attention_impl
functionFull reference ↗list_flash_attention_impls
functionFull reference ↗noop_mask
functionFull reference ↗or_masks
functionFull reference ↗register_flash_attention_impl
functionFull reference ↗- tensorplay.nn.attention.register_flash_attention_impl(impl: str | Literal['FA3', 'FA4'], *, register_fn: Callable[[...], FlashAttentionHandle | None]) None
Register the callable that activates a flash attention impl.
Note
This function is intended for SDPA backend providers to register their implementations. End users should use
activate_flash_attention_impl()to activate a registered implementation.- Parameters:
impl – Implementation identifier (e.g.,
"FA4").register_fn – Callable that performs the actual dispatcher registration. This function will be invoked by
activate_flash_attention_impl()and should register custom kernels with the dispatcher. It may optionally return a handle implementingFlashAttentionHandleto keep any necessary state alive.
Example
>>> def my_impl_register(module_path: str = "my_flash_impl"): ... # Register custom kernels with the dispatcher ... pass >>> register_flash_attention_impl( ... "MyImpl", register_fn=my_impl_register ... )
restore_flash_attention_impl
functionFull reference ↗sdpa_kernel
functionFull reference ↗- tensorplay.nn.attention.sdpa_kernel(backends: list[_SDPBackend] | _SDPBackend, set_priority: bool = False)[source]
Context manager to select which backend to use for scaled dot product attention.
Warning
This function is beta and subject to change.
- Parameters:
backends (Union[List[SDPBackend], SDPBackend]) – A backend or list of backends for scaled dot product attention.
set_priority (bool=False) – Whether the ordering of the backends is interpreted as their priority order.
Example:
from tensorplay.nn.functional import scaled_dot_product_attention from tensorplay.nn.attention import SDPBackend, sdpa_kernel # Only enable flash attention backend with sdpa_kernel(SDPBackend.FLASH_ATTENTION): scaled_dot_product_attention(...) # Enable the Math or Efficient attention backends with sdpa_kernel([SDPBackend.MATH, SDPBackend.EFFICIENT_ATTENTION]): scaled_dot_product_attention(...) # Enable the cuDNN or flash attention backends, and in that order with sdpa_kernel( [SDPBackend.CUDNN_ATTENTION, SDPBackend.FLASH_ATTENTION], set_priority=True ): scaled_dot_product_attention(...)This context manager can be used to select which backend to use for scaled dot product attention. Upon exiting the context manager, the previous state of the flags will be restored, enabling all backends.
Classes 2
BlockMask
classFull reference ↗- class tensorplay.nn.attention.BlockMask(seq_lengths: tuple[int, int], kv_num_blocks: ~tensorplay.Tensor, kv_indices: ~tensorplay.Tensor, full_kv_num_blocks: ~tensorplay.Tensor | None, full_kv_indices: ~tensorplay.Tensor | None, q_num_blocks: ~tensorplay.Tensor | None, q_indices: ~tensorplay.Tensor | None, full_q_num_blocks: ~tensorplay.Tensor | None, full_q_indices: ~tensorplay.Tensor | None, BLOCK_SIZE: tuple[int, int] = (128, 128), mask_mod: ~collections.abc.Callable[[~tensorplay.Tensor, ~tensorplay.Tensor, ~tensorplay.Tensor, ~tensorplay.Tensor], ~tensorplay.Tensor] = <function noop_mask>, *, dq_write_order: ~tensorplay.Tensor | None = None, dq_write_order_full: ~tensorplay.Tensor | None = None, dq_kv_order: ~tensorplay.Tensor | None = None, dq_kv_order_spt: bool | None = None)[source]
BlockMask is our format for representing a block-sparse attention mask. It is somewhat of a cross in-between BCSR and a non-sparse format.
Basics
A block-sparse mask means that instead of representing the sparsity of individual elements in the mask, a KV_BLOCK_SIZE x Q_BLOCK_SIZE block is considered sparse only if every element within that block is sparse. This aligns well with hardware, which generally expects to perform contiguous loads and computation.
This format is primarily optimized for 1. simplicity, and 2. kernel efficiency. Notably, it is not optimized for size, as this mask is always reduced by a factor of KV_BLOCK_SIZE * Q_BLOCK_SIZE. If the size is a concern, the tensors can be reduced in size by increasing the block size.
The essentials of our format are:
num_blocks_in_row: Tensor[ROWS]: Describes the number of blocks present in each row.
col_indices: Tensor[ROWS, MAX_BLOCKS_IN_COL]: col_indices[i] is the sequence of block positions for row i. The values of this row after col_indices[i][num_blocks_in_row[i]] are undefined.
For example, to reconstruct the original tensor from this format:
dense_mask = tensorplay.zeros(ROWS, COLS) for row in range(ROWS): for block_idx in range(num_blocks_in_row[row]): dense_mask[row, col_indices[row, block_idx]] = 1Notably, this format makes it easier to implement a reduction along the rows of the mask.
Details
The basics of our format require only kv_num_blocks and kv_indices. The primary block-sparse layout is represented by up to 4 tensor pairs:
1. (kv_num_blocks, kv_indices): Used for the forwards pass of attention, as we reduce along the KV dimension.
2. [OPTIONAL] (full_kv_num_blocks, full_kv_indices): This is optional and purely an optimization. As it turns out, applying masking to every block is quite expensive! If we specifically know which blocks are “full” and don’t require masking at all, then we can skip applying mask_mod to these blocks. This requires the user to split out a separate mask_mod from the score_mod. For causal masks, this is about a 15% speedup.
3. [GENERATED] (q_num_blocks, q_indices): Required for the backwards pass, as computing dKV requires iterating along the mask along the Q dimension. These are autogenerated from 1.
4. [GENERATED] (full_q_num_blocks, full_q_indices): Same as above, but for the backwards pass. These are autogenerated from 2.
Additional optional tensors may carry deterministic dQ metadata for block-sparse backward:
5. [OPTIONAL] dq_write_order: Write-order metadata for partial blocks. This is produced by create_block_mask when compute_dq_write_order=True, or passed directly to BlockMask.from_kv_blocks by callers that precompute it.
6. [OPTIONAL] dq_write_order_full: Write-order metadata for full blocks, produced or passed the same way as dq_write_order.
7. [OPTIONAL] dq_kv_order: Explicit KV scheduler order used to produce the write-order metadata. create_block_mask currently accepts a boolean dq_kv_order; BlockMask.from_kv_blocks also accepts a tensor for callers that provide precomputed write-order metadata directly.
- as_tuple(flatten: Literal[True] = True) tuple[int, int, Tensor, Tensor, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, bool | None, int, int, Callable[[Tensor, Tensor, Tensor, Tensor], Tensor]][source]
- as_tuple(flatten: Literal[False]) tuple[tuple[int, int], Tensor, Tensor, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, Tensor | None, bool | None, tuple[int, int], Callable[[Tensor, Tensor, Tensor, Tensor], Tensor]]
Returns a tuple of the attributes of the BlockMask.
- Parameters:
flatten (bool) – If True, it will flatten the tuple of (KV_BLOCK_SIZE, Q_BLOCK_SIZE)
- classmethod from_kv_blocks(kv_num_blocks: Tensor, kv_indices: Tensor, full_kv_num_blocks: Tensor | None = None, full_kv_indices: Tensor | None = None, BLOCK_SIZE: int | tuple[int, int] = 128, mask_mod: Callable[[Tensor, Tensor, Tensor, Tensor], Tensor] | None = None, seq_lengths: tuple[int, int] | None = None, compute_q_blocks: bool = True, *, dq_write_order: Tensor | None = None, dq_write_order_full: Tensor | None = None, dq_kv_order: Tensor | bool | None = None) Self[source]
Creates a BlockMask instance from key-value block information.
- Parameters:
kv_num_blocks (Tensor) – Number of kv_blocks in each Q_BLOCK_SIZE row tile.
kv_indices (Tensor) – Indices of key-value blocks in each Q_BLOCK_SIZE row tile.
full_kv_num_blocks (Optional[Tensor]) – Number of full kv_blocks in each Q_BLOCK_SIZE row tile.
full_kv_indices (Optional[Tensor]) – Indices of full key-value blocks in each Q_BLOCK_SIZE row tile.
BLOCK_SIZE (Union[int, tuple[int, int]]) – Size of KV_BLOCK_SIZE x Q_BLOCK_SIZE tiles.
mask_mod (Optional[Callable]) – Function to modify the mask.
dq_write_order (Optional[Tensor]) – Precomputed deterministic dQ write-order metadata.
dq_write_order_full (Optional[Tensor]) – Precomputed deterministic dQ write-order metadata for full blocks.
dq_kv_order (Optional[Union[Tensor, bool]]) – KV-column scheduler order used to produce dq_write_order. A bool selects a built-in order; a tensor gives an explicit scheduler-rank to n-block permutation.
- Returns:
Instance with full Q information generated via _transposed_ordered
- Return type:
- Raises:
RuntimeError – If kv_indices has < 2 dimensions.
AssertionError – If only one of full_kv_* args is provided.
- to(device: device | str) BlockMask[source]
Moves the BlockMask to the specified device.
- Parameters:
device (tensorplay.device or str) – The target device to move the BlockMask to. Can be a device object or a string (e.g., ‘cpu’, ‘cuda:0’).
- Returns:
A new BlockMask instance with all tensor components moved to the specified device.
- Return type:
Note
This method does not modify the original BlockMask in-place. Instead, it returns a new BlockMask instance where individual tensor attributes may or may not be moved to the specified device, depending on their current device placement.
SDPBackend
classFull reference ↗- class tensorplay.nn.attention.SDPBackend
An enum-like class that contains the different backends for scaled dot product attention.
… warning:: This class is in beta and subject to change.
This backend class is designed to be used with the sdpa_kernel context manager.See :func: tensorplay.nn.attention.sdpa_kernel for more details.
Members:
ERROR
MATH
FLASH_ATTENTION
EFFICIENT_ATTENTION
CUDNN_ATTENTION
OVERRIDEABLE
- SDPBackend.name -> str
Help improve this page
Found an error, an unclear step, or a missing example?

