Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 8 additions & 3 deletions cosmos_framework/model/attention/flash2/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__

Expand Down
44 changes: 44 additions & 0 deletions cosmos_framework/model/attention/flash2/support_test.py
Original file line number Diff line number Diff line change
@@ -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()
Loading