# tensorplay.nn.attention API Source: https://www.tensorplay.cn/docs/api/tensorplay.nn.attention.html ## Functions 11 [#](#api-tensorplay.nn.attention.activate_flash_attention_impl) ### activate_flash_attention_impl function[Full reference ↗](/docs/generated/tensorplay.nn.attention.activate_flash_attention_impl.html) ```python 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()](/docs/generated/tensorplay.nn.attention.list_flash_attention_impls.html#tensorplay.nn.attention.list_flash_attention_impls) for available implementations. If the backend’s [register_flash_attention_impl()](/docs/generated/tensorplay.nn.attention.register_flash_attention_impl.html#tensorplay.nn.attention.register_flash_attention_impl) callable returns a FlashAttentionHandle, the registry keeps that handle alive for the lifetime of the process (until explicit uninstall support exists). Example ``` >>> activate_flash_attention_impl("FA4") ``` [#](#api-tensorplay.nn.attention.and_masks) ### and_masks function[Full reference ↗](/docs/generated/tensorplay.nn.attention.and_masks.html) ```python tensorplay.nn.attention.and_masks(*mask_mods: Callable[[Tensor, Tensor, Tensor, Tensor], Tensor]) → Callable[[Tensor, Tensor, Tensor, Tensor], Tensor] ``` Returns a mask_mod that’s the intersection of provided mask_mods [#](#api-tensorplay.nn.attention.create_block_mask) ### create_block_mask function[Full reference ↗](/docs/generated/tensorplay.nn.attention.create_block_mask.html) ```python 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 ``` 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](https://docs.python.org/3/builtins/functions.html#int)) – Batch size. - H ([int](https://docs.python.org/3/builtins/functions.html#int)) – Number of query heads. - Q_LEN ([int](https://docs.python.org/3/builtins/functions.html#int)) – Sequence length of query. - KV_LEN ([int](https://docs.python.org/3/builtins/functions.html#int)) – Sequence length of key/value. - device ([str](https://docs.python.org/3/builtins/stdtypes.html#str)) – Device to run the mask creation on. - BLOCK_SIZE ([int](https://docs.python.org/3/builtins/functions.html#int) or [tuple](https://docs.python.org/3/builtins/stdtypes.html#tuple)[[int](https://docs.python.org/3/builtins/functions.html#int), [int](https://docs.python.org/3/builtins/functions.html#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](https://docs.python.org/3/builtins/functions.html#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](https://docs.python.org/3/builtins/functions.html#bool)) – If True, precompute dQ write-order metadata needed by deterministic block-sparse backward. - dq_kv_order ([bool](https://docs.python.org/3/builtins/functions.html#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: [BlockMask](/docs/generated/tensorplay.nn.attention.BlockMask.html#tensorplay.nn.attention.BlockMask) 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) ``` [#](#api-tensorplay.nn.attention.create_mask) ### create_mask function[Full reference ↗](/docs/generated/tensorplay.nn.attention.create_mask.html) ```python 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 ``` This function creates a mask tensor from a mod_fn function. Parameters: - mod_fn (Union[_score_mod_signature, _mask_mod_signature]) – Function to modify attention scores. - B ([int](https://docs.python.org/3/builtins/functions.html#int)) – Batch size. - H ([int](https://docs.python.org/3/builtins/functions.html#int)) – Number of query heads. - Q_LEN ([int](https://docs.python.org/3/builtins/functions.html#int)) – Sequence length of query. - KV_LEN ([int](https://docs.python.org/3/builtins/functions.html#int)) – Sequence length of key/value. - device ([str](https://docs.python.org/3/builtins/stdtypes.html#str)) – Device to run the mask creation on. Returns: A mask tensor with shape (B, H, M, N). Return type: mask ([Tensor](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor)) [#](#api-tensorplay.nn.attention.current_flash_attention_impl) ### current_flash_attention_impl function[Full reference ↗](/docs/generated/tensorplay.nn.attention.current_flash_attention_impl.html) ```python tensorplay.nn.attention.current_flash_attention_impl() → str | None ``` Return the currently activated flash attention impl name, if any. None indicates that no custom impl has been activated. [#](#api-tensorplay.nn.attention.list_flash_attention_impls) ### list_flash_attention_impls function[Full reference ↗](/docs/generated/tensorplay.nn.attention.list_flash_attention_impls.html) ```python tensorplay.nn.attention.list_flash_attention_impls() → list[str] ``` Return the names of all available flash attention implementations. [#](#api-tensorplay.nn.attention.noop_mask) ### noop_mask function[Full reference ↗](/docs/generated/tensorplay.nn.attention.noop_mask.html) ```python tensorplay.nn.attention.noop_mask(batch: Tensor, head: Tensor, token_q: Tensor, token_kv: Tensor) → Tensor ``` Returns a noop mask_mod [#](#api-tensorplay.nn.attention.or_masks) ### or_masks function[Full reference ↗](/docs/generated/tensorplay.nn.attention.or_masks.html) ```python tensorplay.nn.attention.or_masks(*mask_mods: Callable[[Tensor, Tensor, Tensor, Tensor], Tensor]) → Callable[[Tensor, Tensor, Tensor, Tensor], Tensor] ``` Returns a mask_mod that’s the union of provided mask_mods [#](#api-tensorplay.nn.attention.register_flash_attention_impl) ### register_flash_attention_impl function[Full reference ↗](/docs/generated/tensorplay.nn.attention.register_flash_attention_impl.html) ```python 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()](/docs/generated/tensorplay.nn.attention.activate_flash_attention_impl.html#tensorplay.nn.attention.activate_flash_attention_impl) and should register custom kernels with the dispatcher. It may optionally return a handle implementing FlashAttentionHandle to 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 ... ) ``` [#](#api-tensorplay.nn.attention.restore_flash_attention_impl) ### restore_flash_attention_impl function[Full reference ↗](/docs/generated/tensorplay.nn.attention.restore_flash_attention_impl.html) ```python tensorplay.nn.attention.restore_flash_attention_impl(_raise_warn: bool = True) → None ``` Restore the default FA2 implementation [#](#api-tensorplay.nn.attention.sdpa_kernel) ### sdpa_kernel function[Full reference ↗](/docs/generated/tensorplay.nn.attention.sdpa_kernel.html) ```python tensorplay.nn.attention.sdpa_kernel(backends: list[_SDPBackend] | _SDPBackend, set_priority: bool = False) ``` 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](/docs/generated/tensorplay.nn.attention.SDPBackend.html#tensorplay.nn.attention.SDPBackend)], [SDPBackend](/docs/generated/tensorplay.nn.attention.SDPBackend.html#tensorplay.nn.attention.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 [#](#api-tensorplay.nn.attention.BlockMask) ### BlockMask class[Full reference ↗](/docs/generated/tensorplay.nn.attention.BlockMask.html) ```python 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] = , *, 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) ``` 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]] = 1 ``` Notably, 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. ```python 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]] ``` ```python 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](https://docs.python.org/3/builtins/functions.html#bool)) – If True, it will flatten the tuple of (KV_BLOCK_SIZE, Q_BLOCK_SIZE) ```python 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 ``` Creates a BlockMask instance from key-value block information. Parameters: - kv_num_blocks ([Tensor](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor)) – Number of kv_blocks in each Q_BLOCK_SIZE row tile. - kv_indices ([Tensor](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor)) – Indices of key-value blocks in each Q_BLOCK_SIZE row tile. - full_kv_num_blocks (Optional[[Tensor](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor)]) – Number of full kv_blocks in each Q_BLOCK_SIZE row tile. - full_kv_indices (Optional[[Tensor](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor)]) – Indices of full key-value blocks in each Q_BLOCK_SIZE row tile. - BLOCK_SIZE (Union[[int](https://docs.python.org/3/builtins/functions.html#int), [tuple](https://docs.python.org/3/builtins/stdtypes.html#tuple)[[int](https://docs.python.org/3/builtins/functions.html#int), [int](https://docs.python.org/3/builtins/functions.html#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](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor)]) – Precomputed deterministic dQ write-order metadata. - dq_write_order_full (Optional[[Tensor](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor)]) – Precomputed deterministic dQ write-order metadata for full blocks. - dq_kv_order (Optional[Union[[Tensor](/docs/generated/tensorplay.Tensor.html#tensorplay.Tensor), [bool](https://docs.python.org/3/builtins/functions.html#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: [BlockMask](#tensorplay.nn.attention.BlockMask) Raises: - [RuntimeError](https://docs.python.org/3/builtins/exceptions.html#RuntimeError) – If kv_indices has ```python numel() → int ``` Returns the number of elements (not accounting for sparsity) in the mask. ```python sparsity() → float ``` Computes the percentage of blocks that are sparse (i.e. not computed) ```python to(device: device | str) → BlockMask ``` Moves the BlockMask to the specified device. Parameters: device ([tensorplay.device](/docs/generated/tensorplay.Device.html#tensorplay.Device) or [str](https://docs.python.org/3/builtins/stdtypes.html#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: [BlockMask](#tensorplay.nn.attention.BlockMask) > **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. ```python to_dense() → Tensor ``` Returns a dense block that is equivalent to the block mask. ```python to_string(grid_size: int | tuple[int, int] = (20, 20), limit: int = 4) → str ``` Returns a string representation of the block mask. Quite nifty. If grid_size is -1, prints out an uncompressed version. Warning, it can be quite big! [#](#api-tensorplay.nn.attention.SDPBackend) ### SDPBackend class[Full reference ↗](/docs/generated/tensorplay.nn.attention.SDPBackend.html) ```python 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 ```python SDPBackend.name -> str ```