diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 2c13f6f51e0a..4014e4e6c89b 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -153,6 +153,7 @@ log_info_on_rank0, make_layers, ) +from sglang.srt.utils.common import is_sm120_supported from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.hf_transformers_utils import get_rope_config @@ -204,12 +205,13 @@ def _get_mhc_ops() -> MhcOps: def _is_fused_mhc_post_pre_enabled() -> bool: - # The fused path directly reuses TileLang mhc_post/mhc_pre kernels and their - # tensor layout assumptions, so keep it disabled when either dependency is off. + # SM120 disables the standalone TileLang pre path. mhc_fused_post_pre does + # not read that flag and dispatches independently for both small and large + # token batches, so the standalone pre flag must not veto the fused opt-in. return ( envs.SGLANG_OPT_FUSE_MHC_POST_PRE.get() - and envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get() and envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get() + and (envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get() or is_sm120_supported()) ) diff --git a/test/registered/unit/models/test_deepseek_v4_fused_mhc_policy.py b/test/registered/unit/models/test_deepseek_v4_fused_mhc_policy.py new file mode 100644 index 000000000000..d5e1c3518cae --- /dev/null +++ b/test/registered/unit/models/test_deepseek_v4_fused_mhc_policy.py @@ -0,0 +1,79 @@ +"""Unit tests for the DeepSeek-V4 fused-MHC enable policy.""" + +import unittest +from unittest.mock import patch + +import sglang.srt.models.deepseek_v4 as deepseek_v4 +from sglang.srt.environ import envs +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestDeepseekV4FusedMHCPolicy(CustomTestCase): + def _is_enabled( + self, + *, + fuse: bool, + tilelang_pre: bool, + tilelang_post: bool, + sm120: bool, + ) -> bool: + with ( + envs.SGLANG_OPT_FUSE_MHC_POST_PRE.override(fuse), + envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.override(tilelang_pre), + envs.SGLANG_OPT_USE_TILELANG_MHC_POST.override(tilelang_post), + patch.object(deepseek_v4, "is_sm120_supported", return_value=sm120), + ): + return deepseek_v4._is_fused_mhc_post_pre_enabled() + + def test_sm120_allows_fused_opt_in_with_standalone_pre_disabled(self): + self.assertTrue( + self._is_enabled( + fuse=True, + tilelang_pre=False, + tilelang_post=True, + sm120=True, + ) + ) + + def test_other_platform_still_requires_tilelang_pre(self): + self.assertFalse( + self._is_enabled( + fuse=True, + tilelang_pre=False, + tilelang_post=True, + sm120=False, + ) + ) + self.assertTrue( + self._is_enabled( + fuse=True, + tilelang_pre=True, + tilelang_post=True, + sm120=False, + ) + ) + + def test_fusion_opt_in_and_tilelang_post_remain_required(self): + self.assertFalse( + self._is_enabled( + fuse=False, + tilelang_pre=False, + tilelang_post=True, + sm120=True, + ) + ) + self.assertFalse( + self._is_enabled( + fuse=True, + tilelang_pre=False, + tilelang_post=False, + sm120=True, + ) + ) + + +if __name__ == "__main__": + unittest.main()