TensorPlay
API reference
latest (dev)
Copy
View Markdown

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’s 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")
#

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:

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

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:
  • mod_fn (Union[_score_mod_signature, _mask_mod_signature]) – Function to modify attention scores.

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

Returns:

A mask tensor with shape (B, H, M, N).

Return type:

mask (Tensor)

#

current_flash_attention_impl

functionFull reference ↗
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.

#

list_flash_attention_impls

functionFull reference ↗
tensorplay.nn.attention.list_flash_attention_impls() → list[str]

Return the names of all available flash attention implementations.

#

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

restore_flash_attention_impl

functionFull reference ↗
tensorplay.nn.attention.restore_flash_attention_impl(_raise_warn: bool = True) → None

Restore the default FA2 implementation

#

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]] = 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.

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:

BlockMask

Raises:
numel() → int[source]

Returns the number of elements (not accounting for sparsity) in the mask.

sparsity() → float[source]

Computes the percentage of blocks that are sparse (i.e. not computed)

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:

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.

to_dense() → Tensor[source]

Returns a dense block that is equivalent to the block mask.

to_string(grid_size: int | tuple[int, int] = (20, 20), limit: int = 4) → str[source]

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!

#

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

On this page

Ask DeepWiki