From c57d9c3b75fc5397b69aeec282bb99c6f2865cf5 Mon Sep 17 00:00:00 2001 From: Rahul Steiger Date: Mon, 21 Sep 2026 07:08:07 -0700 Subject: [PATCH] Fix FlashAttention 2 detection when distribution metadata is absent Signed-off-by: Rahul Steiger --- .../model/attention/flash2/__init__.py | 11 +++-- .../model/attention/flash2/support_test.py | 44 +++++++++++++++++++ 2 files changed, 52 insertions(+), 3 deletions(-) create mode 100644 cosmos_framework/model/attention/flash2/support_test.py diff --git a/cosmos_framework/model/attention/flash2/__init__.py b/cosmos_framework/model/attention/flash2/__init__.py index 85c7cbc4e..837570b19 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 000000000..adcafd817 --- /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()