[AMD] Fix multimodal diffusion test crash on ROCm by falling back to SDPA (#22335)
This commit is contained in:
@@ -78,6 +78,8 @@ def _is_fa3_supported(device=None) -> bool:
|
||||
# https://docs.nvidia.com/cuda/cuda-c-programming-guide/#shared-memory-8-x
|
||||
# And for sgl-kernel right now, we can build fa3 on sm80/sm86/sm89/sm90a.
|
||||
# That means if you use A100/A*0/L20/L40/L40s/4090 you can use fa3.
|
||||
if torch.version.cuda is None:
|
||||
return False
|
||||
return (torch.version.cuda >= "12.3") and (
|
||||
torch.cuda.get_device_capability(device)[0] == 9
|
||||
or torch.cuda.get_device_capability(device)[0] == 8
|
||||
|
||||
@@ -149,17 +149,26 @@ class RocmPlatform(Platform):
|
||||
try:
|
||||
import flash_attn # noqa: F401
|
||||
|
||||
from sglang.jit_kernel.flash_attention_v3 import _is_fa3_supported
|
||||
from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import ( # noqa: F401
|
||||
FlashAttentionBackend,
|
||||
)
|
||||
|
||||
supported_sizes = FlashAttentionBackend.get_supported_head_sizes()
|
||||
if head_size not in supported_sizes:
|
||||
if not _is_fa3_supported():
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for head size %d.",
|
||||
head_size,
|
||||
"FlashAttention backend now dispatches through FA3 "
|
||||
"(CUDA-only). Using Torch SDPA backend on ROCm."
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if target_backend == AttentionBackendEnum.FA:
|
||||
supported_sizes = FlashAttentionBackend.get_supported_head_sizes()
|
||||
if head_size not in supported_sizes:
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for head size %d.",
|
||||
head_size,
|
||||
)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
except ImportError:
|
||||
logger.info(
|
||||
"Cannot use FlashAttention backend because the "
|
||||
|
||||
Reference in New Issue
Block a user