From 095e45100b3b78e00464108becad1d58c45037c3 Mon Sep 17 00:00:00 2001 From: Raiden Makoto <81530826+Raiden-Makoto@users.noreply.github.com> Date: Mon, 21 Sep 2026 22:24:29 -0700 Subject: [PATCH] [AMD] [GLM-5.3-Flash Day 0] Route mHC through AITER on gfx950 (#38545) Co-authored-by: Raiden-Makoto Co-authored-by: Thomas Wang Co-authored-by: Kevin Mi <45493463+kevin-mii@users.noreply.github.com> Co-authored-by: Kevin Mi Co-authored-by: Claude Fable 5.1 --- python/sglang/kernels/ops/layernorm/mhc.py | 143 +++++++++++- .../ops/layernorm/test_mhc_aiter_hip.py | 213 ++++++++++++++++++ 2 files changed, 352 insertions(+), 4 deletions(-) create mode 100644 test/registered/kernels/ops/layernorm/test_mhc_aiter_hip.py diff --git a/python/sglang/kernels/ops/layernorm/mhc.py b/python/sglang/kernels/ops/layernorm/mhc.py index 9b069fb3f..a48fe316e 100644 --- a/python/sglang/kernels/ops/layernorm/mhc.py +++ b/python/sglang/kernels/ops/layernorm/mhc.py @@ -18,10 +18,99 @@ from sglang.srt.layers.attention.dsa.utils import is_dsa_prefill_cp_interleave from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.utils.common import strict_contiguous from sglang.srt.runtime_context import get_parallel, get_platform +from sglang.srt.utils import is_gfx95_supported, is_hip from sglang.srt.utils.common import is_gfx1250_supported logger = logging.getLogger(__name__) +_AITER_MHC_RUNTIME_DISABLED = False +_AITER_MHC_ACTIVE_LOGGED = False + + +def _use_aiter_mhc() -> bool: + return ( + not _AITER_MHC_RUNTIME_DISABLED + and is_gfx95_supported() + and envs.SGLANG_USE_AITER.get() + ) + + +def _try_aiter_mhc_pre( + residual: torch.Tensor, + fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + rms_eps: float, + hc_pre_eps: float, + hc_sinkhorn_eps: float, + hc_post_mult_value: float, + sinkhorn_repeat: int, + norm_weight: torch.Tensor | None, + norm_eps: float | None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None: + global _AITER_MHC_RUNTIME_DISABLED, _AITER_MHC_ACTIVE_LOGGED + + try: + from aiter.ops.mhc import mhc_pre as aiter_mhc_pre + except Exception as err: + logger.warning("AITER mHC pre is unavailable, falling back: %s", err) + _AITER_MHC_RUNTIME_DISABLED = True + return None + + kwargs = {} + if norm_weight is not None: + kwargs["norm_weight"] = norm_weight + kwargs["norm_eps"] = norm_eps if norm_eps is not None else rms_eps + + try: + result = aiter_mhc_pre( + residual, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_mult_value, + sinkhorn_repeat, + **kwargs, + ) + except Exception as err: + logger.warning("AITER mHC pre failed, disabling fast path: %s", err) + _AITER_MHC_RUNTIME_DISABLED = True + return None + + if not _AITER_MHC_ACTIVE_LOGGED: + logger.info("Using AITER gfx950 mHC pre/post kernels") + _AITER_MHC_ACTIVE_LOGGED = True + return result + + +def _try_aiter_mhc_post( + x: torch.Tensor, + residual: torch.Tensor, + post_layer_mix: torch.Tensor, + comb_res_mix: torch.Tensor, +) -> torch.Tensor | None: + global _AITER_MHC_RUNTIME_DISABLED + + try: + from aiter.ops.mhc import mhc_post as aiter_mhc_post + except Exception as err: + logger.warning("AITER mHC post is unavailable, falling back: %s", err) + _AITER_MHC_RUNTIME_DISABLED = True + return None + + out = torch.empty_like(residual) + try: + aiter_mhc_post(out, x, residual, post_layer_mix, comb_res_mix) + except Exception as err: + logger.warning("AITER mHC post failed, disabling fast path: %s", err) + _AITER_MHC_RUNTIME_DISABLED = True + return None + return out + + # This module is imported during model-registry discovery. Do not import the real # TileLang package here: it loads native CUDA stubs. The proxy below lets # module-level @tilelang.jit declarations parse, then imports and applies real @@ -119,6 +208,24 @@ pass_configs = { tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, } + +def _use_deep_gemm_hc_prenorm() -> bool: + if is_hip() or not envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): + return False + + from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM + + return ENABLE_JIT_DEEPGEMM + + +def _use_tilelang_mhc_pre() -> bool: + return envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get() and not is_hip() + + +def _use_tilelang_mhc_post() -> bool: + return envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get() and not is_hip() + + FP8 = "float8_e4m3" BF16 = "bfloat16" FP32 = "float32" @@ -1041,7 +1148,7 @@ def mhc_pre( num_tokens, hidden_size, dtype=torch.bfloat16, device=residual.device ) - if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): + if _use_deep_gemm_hc_prenorm(): n_splits = _compute_num_split_for_mhc_pre(num_tokens, hc_hidden_size) gemm_out_mul = torch.empty( @@ -1653,7 +1760,7 @@ def mhc_fused_post_pre( hidden_size, ) - if envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.get(): + if _use_deep_gemm_hc_prenorm(): import deep_gemm deep_gemm.tf32_hc_prenorm_gemm( @@ -1847,7 +1954,25 @@ def _mhc_pre_dispatch( norm_eps: float | None = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, bool]: assert residual.dim() == 3, f"residual must be (s, n, h); got {residual.shape}" - if not envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.get(): + if _use_aiter_mhc(): + result = _try_aiter_mhc_pre( + residual=residual, + fn=fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=rms_eps, + hc_pre_eps=hc_pre_eps, + hc_sinkhorn_eps=hc_sinkhorn_eps, + hc_post_mult_value=hc_post_mult_value, + sinkhorn_repeat=sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + if result is not None: + post_mix, comb_mix, layer_input = result + return post_mix, comb_mix, layer_input, norm_weight is not None + + if not _use_tilelang_mhc_pre(): post_mix, comb_mix, layer_input = _mhc_pre_torch( residual=residual, fn=fn, @@ -1886,7 +2011,17 @@ def _mhc_post_dispatch( ) -> torch.Tensor: assert x.dim() == 2 and residual.dim() == 3 assert post_layer_mix.dim() == 3 and comb_res_mix.dim() == 3 - if not envs.SGLANG_OPT_USE_TILELANG_MHC_POST.get(): + if _use_aiter_mhc(): + result = _try_aiter_mhc_post( + x=x, + residual=residual, + post_layer_mix=post_layer_mix, + comb_res_mix=comb_res_mix, + ) + if result is not None: + return result + + if not _use_tilelang_mhc_post(): return _mhc_post_torch(x, residual, post_layer_mix, comb_res_mix) return mhc_post(x, residual, post_layer_mix, comb_res_mix) diff --git a/test/registered/kernels/ops/layernorm/test_mhc_aiter_hip.py b/test/registered/kernels/ops/layernorm/test_mhc_aiter_hip.py new file mode 100644 index 000000000..612e5e07f --- /dev/null +++ b/test/registered/kernels/ops/layernorm/test_mhc_aiter_hip.py @@ -0,0 +1,213 @@ +"""The AITER mHC route on gfx950: gate, fallback latch, and kernel numerics vs the Torch oracle.""" + +import sys +import types +import unittest +from unittest.mock import patch + +import torch + +from sglang.kernels.ops.layernorm import mhc +from sglang.srt.environ import envs +from sglang.srt.utils import is_gfx95_supported, is_hip +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.test_utils import CustomTestCase + +register_amd_ci(est_time=120, suite="stage-b-test-1-gpu-small-amd-mi35x") + + +@unittest.skipUnless( + torch.cuda.is_available() and is_hip() and is_gfx95_supported(), + "requires one gfx950 GPU", +) +class TestAiterMHCGLM53Flash(CustomTestCase): + hidden_size = 4096 + hc_mult = 4 + rms_eps = 1e-6 + hc_eps = 1e-6 + + def setUp(self): + mhc._AITER_MHC_RUNTIME_DISABLED = False + + def _inputs(self, tokens: int, seed: int = 0): + torch.manual_seed(seed) + device = torch.device("cuda") + mix_size = 2 * self.hc_mult + self.hc_mult**2 + residual = ( + torch.randn( + tokens, + self.hc_mult, + self.hidden_size, + device=device, + dtype=torch.bfloat16, + ) + * 0.1 + ) + fn = ( + torch.randn( + mix_size, + self.hc_mult * self.hidden_size, + device=device, + dtype=torch.float32, + ) + * 0.01 + ) + scale = torch.tensor([0.5, 0.25, 0.25], device=device, dtype=torch.float32) + base = torch.zeros(mix_size, device=device, dtype=torch.float32) + return residual, fn, scale, base + + def _rmsnorm(self, x, weight): + return ( + x.float() + * torch.rsqrt(x.float().square().mean(dim=-1, keepdim=True) + self.rms_eps) + * weight.float() + ).to(x.dtype) + + def test_gate_selects_aiter_on_gfx950(self): + """A gate that resolves False on real hardware silently serves the Torch path.""" + with envs.SGLANG_USE_AITER.override(True): + self.assertTrue(mhc._use_aiter_mhc()) + + def test_hip_without_aiter_stays_on_torch_and_never_loads_tilelang(self): + """The TileLang/DeepGEMM flags default on; only the HIP gate keeps them off this device.""" + residual, fn, scale, base = self._inputs(8) + x = residual.reshape(8, self.hc_mult * self.hidden_size) + _, _, layer_ref = mhc._mhc_pre_torch( + residual, fn, scale, base, self.rms_eps, self.hc_eps, self.hc_eps, 2.0, 4 + ) + with ( + envs.SGLANG_USE_AITER.override(False), + envs.SGLANG_OPT_USE_TILELANG_MHC_PRE.override(True), + envs.SGLANG_OPT_USE_TILELANG_MHC_POST.override(True), + envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.override(True), + patch.object( + mhc, "_load_tilelang", side_effect=AssertionError("TileLang imported") + ), + ): + self.assertFalse(mhc._use_aiter_mhc()) + self.assertFalse(mhc._use_tilelang_mhc_pre()) + self.assertFalse(mhc._use_tilelang_mhc_post()) + self.assertFalse(mhc._use_deep_gemm_hc_prenorm()) + layer_input, h_res, h_post, norm_fused = mhc.hc_pre( + x, fn, scale, base, self.hc_mult, self.rms_eps, self.hc_eps, 4 + ) + out = mhc.hc_post(layer_input, x, h_post, h_res, self.hc_mult) + self.assertFalse(norm_fused) + torch.testing.assert_close(layer_input, layer_ref) + self.assertTrue(torch.isfinite(out).all()) + + def test_aiter_import_and_runtime_failures_latch_to_torch(self): + """A missing symbol or a raising kernel must disable the route once, not fail the request.""" + residual, fn, scale, base = self._inputs(8) + x = residual.reshape(8, self.hc_mult * self.hidden_size) + modules = { + "aiter": types.ModuleType("aiter"), + "aiter.ops": types.ModuleType("aiter.ops"), + "aiter.ops.mhc": types.ModuleType("aiter.ops.mhc"), + } + with patch.dict(sys.modules, modules): + result = mhc._try_aiter_mhc_pre( + residual, + fn, + scale, + base, + self.rms_eps, + self.hc_eps, + self.hc_eps, + 2.0, + 4, + None, + None, + ) + self.assertIsNone(result) + self.assertTrue(mhc._AITER_MHC_RUNTIME_DISABLED) + + mhc._AITER_MHC_RUNTIME_DISABLED = False + + def fail_post(*_args, **_kwargs): + raise RuntimeError("synthetic failure") + + failing = types.ModuleType("aiter.ops.mhc") + failing.mhc_post = fail_post + modules["aiter.ops.mhc"] = failing + with envs.SGLANG_USE_AITER.override(False): + layer_input, h_res, h_post, _ = mhc.hc_pre( + x, fn, scale, base, self.hc_mult, self.rms_eps, self.hc_eps, 4 + ) + with ( + patch.dict(sys.modules, modules), + envs.SGLANG_USE_AITER.override(True), + ): + out = mhc.hc_post(layer_input, x, h_post, h_res, self.hc_mult) + self.assertTrue(mhc._AITER_MHC_RUNTIME_DISABLED) + self.assertTrue(torch.isfinite(out).all()) + + def test_aiter_pre_post_match_torch_oracle(self): + """A positional or kwarg mixup in the AITER call shows up only against the real kernel.""" + norm_weight = torch.linspace( + 0.75, 1.25, self.hidden_size, device="cuda", dtype=torch.bfloat16 + ) + for tokens in (1, 8, 17, 32, 64, 128): + for sinkhorn_iters in (2, 20): + for fused_norm in (False, True): + with self.subTest( + tokens=tokens, sinkhorn_iters=sinkhorn_iters, norm=fused_norm + ): + residual, fn, scale, base = self._inputs(tokens) + post_ref, comb_ref, layer_ref = mhc._mhc_pre_torch( + residual, + fn, + scale, + base, + self.rms_eps, + self.hc_eps, + self.hc_eps, + 2.0, + sinkhorn_iters, + ) + result = mhc._try_aiter_mhc_pre( + residual, + fn, + scale, + base, + self.rms_eps, + self.hc_eps, + self.hc_eps, + 2.0, + sinkhorn_iters, + norm_weight if fused_norm else None, + self.rms_eps if fused_norm else None, + ) + self.assertIsNotNone(result, "AITER mHC pre fell back") + post_out, comb_out, layer_out = result + if fused_norm: + layer_ref = self._rmsnorm(layer_ref, norm_weight) + torch.cuda.synchronize() + + torch.testing.assert_close( + post_out, post_ref, atol=2e-3, rtol=2e-3 + ) + torch.testing.assert_close( + comb_out, comb_ref, atol=2e-3, rtol=2e-3 + ) + torch.testing.assert_close( + layer_out, layer_ref, atol=2e-2, rtol=2e-2 + ) + + x = (layer_ref.float() * 0.75).to(layer_ref.dtype) + post_ref_out = mhc._mhc_post_torch( + x, residual, post_ref, comb_ref + ) + post_out_actual = mhc._try_aiter_mhc_post( + x, residual, post_out, comb_out + ) + self.assertIsNotNone( + post_out_actual, "AITER mHC post fell back" + ) + torch.testing.assert_close( + post_out_actual, post_ref_out, atol=2e-2, rtol=2e-2 + ) + + +if __name__ == "__main__": + unittest.main()