Uh oh!
There was an error while loading. Please reload this page.
- Notifications
You must be signed in to change notification settings - Fork 7.3k
[core] use kernels to support _flash_3_hub attention backend#12236
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Uh oh!
There was an error while loading. Please reload this page.
Changes from all commits
827fc15a0177ebac43e84bc4097187d08792bb37964e69d42595ae6b548f56e0e7eac06d5c247c2a5aff0097c57943b4a866a681193c3eb9a1e1faf25c701d8dada049168e62648c9dcFile filter
Filter by extension
Conversations
Uh oh!
There was an error while loading. Please reload this page.
Jump to
Uh oh!
There was an error while loading. Please reload this page.
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -26,6 +26,7 @@ | ||
| is_flash_attn_3_available, | ||
| is_flash_attn_available, | ||
| is_flash_attn_version, | ||
| is_kernels_available, | ||
| is_sageattention_available, | ||
| is_sageattention_version, | ||
| is_torch_npu_available, | ||
| @@ -35,7 +36,7 @@ | ||
| is_xformers_available, | ||
| is_xformers_version, | ||
| ) | ||
| from ..utils.constants import DIFFUSERS_ATTN_BACKEND, DIFFUSERS_ATTN_CHECKS | ||
| from ..utils.constants import DIFFUSERS_ATTN_BACKEND, DIFFUSERS_ATTN_CHECKS, DIFFUSERS_ENABLE_HUB_KERNELS | ||
| _REQUIRED_FLASH_VERSION = "2.6.3" | ||
| @@ -67,6 +68,17 @@ | ||
| flash_attn_3_func = None | ||
| flash_attn_3_varlen_func = None | ||
| if DIFFUSERS_ENABLE_HUB_KERNELS: | ||
| if not is_kernels_available(): | ||
| raise ImportError( | ||
| "To use FA3 kernel for your hardware from the Hub, the `kernels` library must be installed. Install with `pip install kernels`." | ||
| ) | ||
| from ..utils.kernels_utils import _get_fa3_from_hub | ||
| flash_attn_interface_hub = _get_fa3_from_hub() | ||
sayakpaul marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| flash_attn_3_func_hub = flash_attn_interface_hub.flash_attn_func | ||
| else: | ||
| flash_attn_3_func_hub = None | ||
| if _CAN_USE_SAGE_ATTN: | ||
| from sageattention import ( | ||
| @@ -153,6 +165,8 @@ class AttentionBackendName(str, Enum): | ||
| FLASH_VARLEN = "flash_varlen" | ||
| _FLASH_3 = "_flash_3" | ||
| _FLASH_VARLEN_3 = "_flash_varlen_3" | ||
| _FLASH_3_HUB = "_flash_3_hub" | ||
| # _FLASH_VARLEN_3_HUB = "_flash_varlen_3_hub" # not supported yet. | ||
| # PyTorch native | ||
| FLEX = "flex" | ||
| @@ -351,6 +365,17 @@ def _check_attention_backend_requirements(backend: AttentionBackendName) -> None | ||
| f"Flash Attention 3 backend '{backend.value}' is not usable because of missing package or the version is too old. Please build FA3 beta release from source." | ||
| ) | ||
| # TODO: add support Hub variant of FA3 varlen later | ||
| elif backend in [AttentionBackendName._FLASH_3_HUB]: | ||
| if not DIFFUSERS_ENABLE_HUB_KERNELS: | ||
| raise RuntimeError( | ||
| f"Flash Attention 3 Hub backend '{backend.value}' is not usable because the `DIFFUSERS_ENABLE_HUB_KERNELS` env var isn't set. Please set it like `export DIFFUSERS_ENABLE_HUB_KERNELS=yes`." | ||
| ) | ||
| if not is_kernels_available(): | ||
| raise RuntimeError( | ||
| f"Flash Attention 3 Hub backend '{backend.value}' is not usable because the `kernels` package isn't available. Please install it with `pip install kernels`." | ||
| ) | ||
| elif backend in [ | ||
| AttentionBackendName.SAGE, | ||
| AttentionBackendName.SAGE_VARLEN, | ||
| @@ -657,6 +682,44 @@ def _flash_attention_3( | ||
| return (out, lse) if return_attn_probs else out | ||
| @_AttentionBackendRegistry.register( | ||
| AttentionBackendName._FLASH_3_HUB, | ||
| constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], | ||
| ) | ||
| def _flash_attention_3_hub( | ||
| query: torch.Tensor, | ||
| key: torch.Tensor, | ||
| value: torch.Tensor, | ||
| scale: Optional[float] = None, | ||
| is_causal: bool = False, | ||
| window_size: Tuple[int, int] = (-1, -1), | ||
| softcap: float = 0.0, | ||
| deterministic: bool = False, | ||
| return_attn_probs: bool = False, | ||
| ) -> torch.Tensor: | ||
| out = flash_attn_3_func_hub( | ||
MemberAuthor There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Follow (internal) this link | ||
| q=query, | ||
| k=key, | ||
| v=value, | ||
| softmax_scale=scale, | ||
| causal=is_causal, | ||
| qv=None, | ||
| q_descale=None, | ||
| k_descale=None, | ||
| v_descale=None, | ||
| window_size=window_size, | ||
| softcap=softcap, | ||
| num_splits=1, | ||
| pack_gqa=None, | ||
| deterministic=deterministic, | ||
| sm_margin=0, | ||
| return_attn_probs=return_attn_probs, | ||
| ) | ||
| # When `return_attn_probs` is True, the above returns a tuple of | ||
| # actual outputs and lse. | ||
| return (out[0], out[1]) if return_attn_probs else out | ||
| @_AttentionBackendRegistry.register( | ||
| AttentionBackendName._FLASH_VARLEN_3, | ||
| constraints=[_check_device, _check_qkv_dtype_bf16_or_fp16, _check_shape], | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,23 @@ | ||
| from ..utils import get_logger | ||
| from .import_utils import is_kernels_available | ||
| logger = get_logger(__name__) | ||
| _DEFAULT_HUB_ID_FA3 = "kernels-community/flash-attn3" | ||
| def _get_fa3_from_hub(): | ||
| if not is_kernels_available(): | ||
| return None | ||
| else: | ||
| from kernels import get_kernel | ||
| try: | ||
| # TODO: temporary revision for now. Remove when merged upstream into `main`. | ||
| flash_attn_3_hub = get_kernel(_DEFAULT_HUB_ID_FA3, revision="fake-ops-return-probs") | ||
| return flash_attn_3_hub | ||
| except Exception as e: | ||
| logger.error(f"An error occurred while fetching kernel '{_DEFAULT_HUB_ID_FA3}' from the Hub: {e}") | ||
| raise |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Made it a constant in
constants.pyas I think it will be shared across modules.