TensorPlay
latest (dev)
Copy
View Markdown

Latest development documentation · Updated 2026-10-08

Source code for tensorplay.nn.attention

# mypy: allow-untyped-defs
"""This module contains functions and classes that alter the behavior of tensorplay.nn.functional.scaled_dot_product_attention"""

import contextlib
from collections.abc import Iterable
from contextvars import ContextVar
from typing import Union
from warnings import warn

import tensorplay.backends.cuda
from tensorplay._C import _SDPBackend as SDPBackend
from tensorplay.backends.cuda import (
    can_use_efficient_attention,
    can_use_flash_attention,
    SDPAParams,
)


__all__: list[str] = [
    "SDPBackend",
    "sdpa_kernel",
    "WARN_FOR_UNFUSED_KERNELS",
    "register_flash_attention_impl",
    "activate_flash_attention_impl",
    "list_flash_attention_impls",
    "current_flash_attention_impl",
    "restore_flash_attention_impl",
]


# Note: [SDPA warnings]
# This only affects users of bias subclasses
# If this is set to True, we will warn the user if they are not using the fused kernels
# As well, it will raise warnings for all the reasons why the fused kernels can't be run.
# To set this to True, run
# tensorplay.nn.attention.WARN_FOR_UNFUSED_KERNELS = True
WARN_FOR_UNFUSED_KERNELS = False


r"""An enum-like class that contains the different backends for scaled dot product attention.
    This backend class is designed to be used with the sdpa_kernel context manager.

    The following Enums are available:
        - ERROR: An error occurred when trying to determine the backend.
        - MATH: The math backend for scaled dot product attention.
        - FLASH_ATTENTION: The flash attention backend for scaled dot product attention.
        - EFFICIENT_ATTENTION: The efficient attention backend for scaled dot product attention.
        - CUDNN_ATTENTION: The cuDNN backend for scaled dot product attention.
        - OVERRIDEABLE: The overridable backend for extension.

    See :func:`tensorplay.nn.attention.sdpa_kernel` for more details.

    .. warning:: This class is in beta and subject to change.
"""
SDPBackend.__module__ = __name__
SDPBackend.__name__ = "SDPBackend"


def _raise_kernel_warnings(params: SDPAParams) -> None:
    """
    If WARN_FOR_UNFUSED_KERNELS is set to True, this will raise warnings
    for all the reasons why the fused kernels can't be run. If using subclasses
    """
    if WARN_FOR_UNFUSED_KERNELS:
        if not can_use_efficient_attention(params):
            warn("Efficient attention can't be used because:", stacklevel=2)
            can_use_efficient_attention(params, True)
        if not can_use_flash_attention(params):
            warn("Flash attention can't be used because:", stacklevel=2)
            can_use_flash_attention(params, True)


_backend_names = {
    "cudnn": "CUDNN_ATTENTION",
    "flash": "FLASH_ATTENTION",
    "mem_efficient": "EFFICIENT_ATTENTION",
    "math": "MATH",
    "overrideable": "OVERRIDEABLE",
}
_backend_enabled = {
    name: getattr(tensorplay._C, f"_get_{name}_sdp_enabled")
    for name in _backend_names
}
_backend_setter = {
    name: getattr(tensorplay._C, f"_set_sdp_use_{name}")
    for name in _backend_names
}
_backend_value = {name: getattr(SDPBackend, val) for name, val in _backend_names.items()}
_sdpa_kernel_uses_priority = ContextVar("sdpa_kernel_uses_priority", default=False)


def _is_sdp_priority_order_active() -> bool:
    """Return whether an enclosing sdpa_kernel context set backend priority."""
    return _sdpa_kernel_uses_priority.get()


def _backend_from_string(name: str):
    return getattr(SDPBackend, name)


_backend_enabled = {
    name: getattr(tensorplay._C, f"_get_{name}_sdp_enabled")
    for name in _backend_names
}


def _cur_sdpa_kernel_backends(with_priority: bool = False):
    # The enabled-set read is a C++ round trip per backend; the member lookups
    # are hoisted to import time so a hot call pays bound calls rather than
    # attribute walks from the module root.
    backends = []
    for name, val in _backend_names.items():
        if _backend_enabled[name]():
            backends.append(_backend_value[name])
    if with_priority:
        curr_priority = tensorplay._C._get_sdp_priority_order()
        backends = sorted(
            backends, key=lambda backend: curr_priority.index(int(backend))
        )
    return backends


def _sdpa_kernel(backends: Iterable, set_priority: bool = False) -> None:
    for name, val in _backend_names.items():
        enabled = _backend_value[name] in backends
        _backend_setter[name](enabled)
    if set_priority:
        # backends should be a unique list
        user_priority = [int(backend) for backend in backends]
        previous_priority = tensorplay._C._get_sdp_priority_order()
        for backend in previous_priority:
            if backend not in user_priority:
                user_priority.append(int(backend))
        tensorplay._C._set_sdp_priority_order(user_priority)



[docs]
@contextlib.contextmanager
def sdpa_kernel(backends: list[SDPBackend] | SDPBackend, set_priority: bool = False):
    r"""
    Context manager to select which backend to use for scaled dot product attention.

    .. warning:: This function is beta and subject to change.

    Args:
        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:

    .. code-block:: python

        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.
    """
    if not isinstance(backends, (list, SDPBackend)):
        raise AssertionError(
            f"Backend must be an instance of SDPBackend or a list of SDPBackend instances, got {type(backends).__name__}"
        )

    if isinstance(backends, SDPBackend):
        backends = [backends]

    backends = list(dict.fromkeys(backends))

    previous_backends = _cur_sdpa_kernel_backends(with_priority=set_priority)
    priority_token = _sdpa_kernel_uses_priority.set(
        set_priority or _sdpa_kernel_uses_priority.get()
    )
    try:
        _sdpa_kernel(backends, set_priority)
        yield {}
    finally:
        _sdpa_kernel_uses_priority.reset(priority_token)
        _sdpa_kernel(previous_backends, set_priority)



# variadic version of sdpa_kernel for the tracer to use while reconstructing
@contextlib.contextmanager
def _sdpa_kernel_variadic(*backends: SDPBackend):
    with sdpa_kernel(list(backends)):
        yield


def _get_flash_version() -> str:
    """This returns the closest matching tag for the flash attention backend"""
    return "2.5.7"


from .omni_attention import (
    AuxOutput,
    AuxRequest,
    BlockMask,
    and_masks,
    create_block_mask,
    create_mask,
    omni_attention,
    noop_mask,
    or_masks,
)


__all__ += [
    "AuxOutput",
    "AuxRequest",
    "BlockMask",
    "and_masks",
    "create_block_mask",
    "create_mask",
    "omni_attention",
    "noop_mask",
    "or_masks",
]

from . import _registry


# Re-export registry types and functions for public API
_FlashAttentionImpl = _registry._FlashAttentionImpl
_RegisterFn = _registry._RegisterFn
register_flash_attention_impl = _registry.register_flash_attention_impl
activate_flash_attention_impl = _registry.activate_flash_attention_impl
list_flash_attention_impls = _registry.list_flash_attention_impls
current_flash_attention_impl = _registry.current_flash_attention_impl
restore_flash_attention_impl = _registry.restore_flash_attention_impl

register_flash_attention_impl.__module__ = __name__
activate_flash_attention_impl.__module__ = __name__
list_flash_attention_impls.__module__ = __name__
current_flash_attention_impl.__module__ = __name__
restore_flash_attention_impl.__module__ = __name__

# Import built-in implementations to trigger self-registration
from . import _fa3, _fa4
Ask DeepWiki