TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

tensorplay.nn.attention.create_block_mask

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)

On this page

Ask DeepWiki