latest (dev)
Copy
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:
- 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)
Help improve this page
Found an error, an unclear step, or a missing example?

