diff --git a/cosmos_framework/model/attention/flash2/__init__.py b/cosmos_framework/model/attention/flash2/__init__.py index 85c7cbc4..837570b1 100644 --- a/cosmos_framework/model/attention/flash2/__init__.py +++ b/cosmos_framework/model/attention/flash2/__init__.py @@ -45,9 +45,14 @@ def flash2_supported() -> bool: flash2_version_str = None if not hasattr(flash_attn, "__version__"): - from importlib.metadata import version - - flash2_version_str = version("flash_attn") + from importlib.metadata import PackageNotFoundError, version + + try: + flash2_version_str = version("flash_attn") + except PackageNotFoundError: + # FlashAttention 4 also provides this namespace, without the FA2 distribution. + log.debug("Flash Attention v2 is not supported because its distribution was not found.") + return False else: flash2_version_str = flash_attn.__version__ diff --git a/cosmos_framework/model/attention/flash2/support_test.py b/cosmos_framework/model/attention/flash2/support_test.py new file mode 100644 index 00000000..adcafd81 --- /dev/null +++ b/cosmos_framework/model/attention/flash2/support_test.py @@ -0,0 +1,44 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +import sys +from importlib.metadata import PackageNotFoundError +from types import ModuleType +from unittest.mock import patch + +import pytest + +from cosmos_framework.model.attention.flash2 import flash2_supported + + +@pytest.fixture +def flash_attn_namespace(monkeypatch): + """Model the namespace shared by FA2 and FA4 without requiring CUDA kernels.""" + namespace = ModuleType("flash_attn") + monkeypatch.setitem(sys.modules, "flash_attn", namespace) + monkeypatch.setattr("torch.cuda.is_available", lambda: True) + return namespace + + +def test_missing_distribution(flash_attn_namespace): + with patch("importlib.metadata.version", side_effect=PackageNotFoundError("flash_attn")): + assert not flash2_supported() + + +@pytest.mark.parametrize("version, supported", [("2.7.4.post1", True), ("2.6.0", False)]) +def test_distribution_version(flash_attn_namespace, version, supported): + with patch("importlib.metadata.version", return_value=version): + assert flash2_supported() is supported + + +@pytest.mark.parametrize("version, supported", [("2.7.4.post1", True), ("2.6.0", False)]) +def test_module_version(flash_attn_namespace, version, supported): + flash_attn_namespace.__version__ = version + with patch("importlib.metadata.version", side_effect=PackageNotFoundError("flash_attn")): + assert flash2_supported() is supported + + +def test_no_cuda(flash_attn_namespace, monkeypatch): + monkeypatch.setattr("torch.cuda.is_available", lambda: False) + flash_attn_namespace.__version__ = "2.7.4.post1" + assert not flash2_supported()