diff --git a/megatron/core/utils.py b/megatron/core/utils.py index 169aebc27f9..14c5478f212 100644 --- a/megatron/core/utils.py +++ b/megatron/core/utils.py @@ -484,7 +484,7 @@ def is_flashinfer_min_version(version, check_equality=True): return False if check_equality: return flashinfer_version >= PkgVersion(version) - return flashinver_version > PkgVersion(version) + return flashinfer_version > PkgVersion(version) def accepts_parameter(func: Callable, name: str) -> bool: diff --git a/tests/unit_tests/test_utils.py b/tests/unit_tests/test_utils.py index 94ac440d8e0..b9db75c2fcb 100644 --- a/tests/unit_tests/test_utils.py +++ b/tests/unit_tests/test_utils.py @@ -52,6 +52,24 @@ def test_divide_improperly(): util.divide(4, 5) +@pytest.mark.skipif(not util.HAVE_PACKAGING, reason="packaging is not installed") +@pytest.mark.parametrize("check_equality", [True, False]) +def test_is_flashinfer_min_version(check_equality): + from packaging.version import Version as PkgVersion + + with patch.object(util, "get_flashinfer_version", return_value=PkgVersion("0.6.5")): + # check_equality=False exercised the path that used to reference an + # undefined name and raise NameError instead of returning a bool. + assert util.is_flashinfer_min_version("0.6.4", check_equality=check_equality) is True + assert util.is_flashinfer_min_version("0.7.0", check_equality=check_equality) is False + assert ( + util.is_flashinfer_min_version("0.6.5", check_equality=check_equality) is check_equality + ) + + with patch.object(util, "get_flashinfer_version", return_value=None): + assert util.is_flashinfer_min_version("0.6.4", check_equality=check_equality) is False + + def test_experimental_cls_init(): with patch.object(config, 'ENABLE_EXPERIMENTAL', True): # Check that initialization works