# tensorplay.nn.attention Source: https://www.tensorplay.cn/docs/nn.attention.html This module contains functions and classes that alter the behavior of tensorplay.nn.functional.scaled_dot_product_attention tensorplay.nn.attention collects the pieces around [tensorplay.nn.functional.scaled_dot_product_attention()](/docs/generated/tensorplay.nn.functional.scaled_dot_product_attention.html#tensorplay.nn.functional.scaled_dot_product_attention): the backend-selection context, the flash-attention implementation registry, helpers for the causal and variable-length variants, and omni_attention — the programmable form of attention that takes user-written score-modification and block-mask functions. scaled_dot_product_attention routes each call to one of a small set of kernels — the math fallback, flash attention, or memory-efficient attention. Which one is eligible depends on the inputs (dtype, head dim, mask type, whether gradients are needed), and [tensorplay.nn.attention.sdpa_kernel()](/docs/generated/tensorplay.nn.attention.sdpa_kernel.html#tensorplay.nn.attention.sdpa_kernel) lets you restrict that choice from the outside, either to force a specific kernel or to study how the candidates behave. ``` import tensorplay as tp import tensorplay.nn.functional as F from tensorplay.nn.attention import SDPBackend, sdpa_kernel q = k = v = tp.randn(1, 2, 8, 16) # restrict routing to the math implementation while experimenting with sdpa_kernel(SDPBackend.MATH): out = F.scaled_dot_product_attention(q, k, v) # a list means "any of these, in this preference order" with sdpa_kernel([SDPBackend.MATH, SDPBackend.EFFICIENT_ATTENTION]): out = F.scaled_dot_product_attention(q, k, v) ``` The backend names are [tensorplay.nn.attention.SDPBackend](/docs/generated/tensorplay.nn.attention.SDPBackend.html#tensorplay.nn.attention.SDPBackend) members: MATH (the reference implementation), FLASH_ATTENTION, EFFICIENT_ATTENTION, and CUDNN_ATTENTION. The registry functions manage which flash-attention implementation backs the FLASH_ATTENTION route. [tensorplay.nn.attention.list_flash_attention_impls()](/docs/generated/tensorplay.nn.attention.list_flash_attention_impls.html#tensorplay.nn.attention.list_flash_attention_impls) reports the implementations this build knows about ("FA3" and "FA4"). [tensorplay.nn.attention.activate_flash_attention_impl()](/docs/generated/tensorplay.nn.attention.activate_flash_attention_impl.html#tensorplay.nn.attention.activate_flash_attention_impl) switches to one of them and returns a handle; [tensorplay.nn.attention.restore_flash_attention_impl()](/docs/generated/tensorplay.nn.attention.restore_flash_attention_impl.html#tensorplay.nn.attention.restore_flash_attention_impl) returns to the previous state. Activating an implementation requires its kernel package to be importable — the loader raises ModuleNotFoundError naming the missing module when it is not. [tensorplay.nn.attention.register_flash_attention_impl()](/docs/generated/tensorplay.nn.attention.register_flash_attention_impl.html#tensorplay.nn.attention.register_flash_attention_impl) registers a custom implementation under a new name, which is how additional flash-attention backends plug in. By default no implementation is active: [tensorplay.nn.attention.current_flash_attention_impl()](/docs/generated/tensorplay.nn.attention.current_flash_attention_impl.html#tensorplay.nn.attention.current_flash_attention_impl) reports None. ## Utils | sdpa_kernel |Context manager to select which backend to use for scaled dot product attention. | | --- | --- | | SDPBackend | An enum-like class that contains the different backends for scaled dot product attention. | | BlockMask | BlockMask is our format for representing a block-sparse attention mask. | | create_block_mask | This function creates a block mask tuple from a mask_mod function. | | register_flash_attention_impl | Register the callable that activates a flash attention impl. | | activate_flash_attention_impl | Activate into the dispatcher a previously registered flash attention impl. | | list_flash_attention_impls | Return the names of all available flash attention implementations. | | current_flash_attention_impl | Return the currently activated flash attention impl name, if any. | | restore_flash_attention_impl | Restore the default FA2 implementation | | create_mask | This function creates a mask tensor from a mod_fn function. | | and_masks | Returns a mask_mod that's the intersection of provided mask_mods | | or_masks | Returns a mask_mod that's the union of provided mask_mods | | noop_mask | Returns a noop mask_mod | tensorplay.nn.attention.can_use_flash_attention() and tensorplay.nn.attention.can_use_efficient_attention() answer the routing question for one concrete call: they take an SDPAParams record (query/key/value tensors, dropout probability, and mask) and report whether the corresponding kernel would accept it. They are the same probes the dispatcher consults. ## Mask helpers Four functions turn mask predicates — plain functions of (b, h, q_idx, kv_idx) returning a boolean tensor — into mask material you can pass around. [tensorplay.nn.attention.create_mask()](/docs/generated/tensorplay.nn.attention.create_mask.html#tensorplay.nn.attention.create_mask) compiles a predicate into a dense boolean mask, and [tensorplay.nn.attention.and_masks()](/docs/generated/tensorplay.nn.attention.and_masks.html#tensorplay.nn.attention.and_masks) / [tensorplay.nn.attention.or_masks()](/docs/generated/tensorplay.nn.attention.or_masks.html#tensorplay.nn.attention.or_masks) compose predicates before compiling; [tensorplay.nn.attention.noop_mask()](/docs/generated/tensorplay.nn.attention.noop_mask.html#tensorplay.nn.attention.noop_mask) is the allow-everything predicate: ``` import tensorplay as tp from tensorplay.nn.attention import and_masks, create_mask def causal(b, h, q_idx, kv_idx): return q_idx >= kv_idx def window(b, h, q_idx, kv_idx): return q_idx - kv_idx < 4 mask = create_mask(and_masks(causal, window), B=1, H=1, Q_LEN=6, KV_LEN=6) print(mask.shape, mask.dtype) # (1, 1, 6, 6) bool ``` ## Submodules | omni_attention |This function implements scaled dot product attention with an arbitrary attention score modification function. | | --- | --- | | bias | Defines bias subclasses that work with scaled_dot_product_attention | | varlen | Variable-length attention implementation using Flash Attention. |