From 6304f9671568ade4ae3cba5bd8e1dbd15c82a545 Mon Sep 17 00:00:00 2001 From: Raiden-Makoto Date: Mon, 10 Aug 2026 17:57:56 -0700 Subject: [PATCH 1/3] perf(amd): add tuned MXFP4 GLM MLA absorbed BMM Preserve packed MXFP4 absorbed weights and route current ROCm MLA K/V BMMs through shape-aware A16WFP4 kernels without changing decode tiles. --- python/sglang/srt/environ.py | 3 + .../attention_forward_methods/forward_mla.py | 139 +++++ .../forward_mla_rocm.py | 13 +- .../deepseek_common/deepseek_weight_loader.py | 26 +- .../models/test_glm_mxfp4_absorbed_bmm.py | 492 ++++++++++++++++++ 5 files changed, 658 insertions(+), 15 deletions(-) create mode 100644 test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 5d0f844b2d79..8528f2b92000 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1386,6 +1386,9 @@ class Envs: SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False) SGLANG_DSA_TOPK_BROADCAST = EnvBool(False) SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False) + # Opt-in GLM MLA absorbed-BMM backend that keeps w_kc/w_vc in packed + # MXFP4 and dispatches the matching AITER FP4 kernels. + SGLANG_USE_MXFP4_MLA_BMM = EnvBool(False) # Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: fuse the # absorbed q bmm with the nope/rope concat + fp8 cast so q is written # directly in fp8 ("born fp8") and the standalone concat-cast kernel diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index e140a117db48..56fd1b693b1c 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -60,6 +60,7 @@ _is_cuda, _is_hip, _is_musa, + _use_aiter_gfx95, ) from sglang.srt.runtime_context import get_exec, get_parallel from sglang.srt.state_capturer.indexer_topk import ( @@ -110,6 +111,114 @@ def is_mla_dcp_lse_base_on_e(attention_backend: Optional[str]) -> bool: from sglang.kernels.ops.gemm import bmm_fp8 +if _use_aiter_gfx95: + from aiter.ops.triton._triton_kernels.gemm.batched.batched_gemm_a16wfp4 import ( + _get_config as _get_mxfp4_bmm_config, + ) + from aiter.ops.triton.batched_gemm_a16wfp4 import batched_gemm_a16wfp4 + + from sglang.srt.layers.quantization.rocm_mxfp4_utils import ( + batched_gemm_afp4wfp4_pre_quant, + ) + + +def _get_single_split_mxfp4_bmm_config(x: torch.Tensor, weight: torch.Tensor) -> dict: + config, _ = _get_mxfp4_bmm_config(x.shape[1], weight.shape[1], x.shape[2]) + config = config.copy() + config["NUM_KSPLIT"] = 1 + return config + + +def _get_glm_mxfp4_k_bmm_config(x: torch.Tensor, weight: torch.Tensor) -> dict: + config = _get_single_split_mxfp4_bmm_config(x, weight) + # GLM's K-up has K=192. Larger blocks over-read its six E8M0 scale groups. + config["BLOCK_SIZE_K"] = 64 + return config + + +def _get_glm_mxfp4_v_bmm_config(x: torch.Tensor, weight: torch.Tensor) -> dict: + config = _get_single_split_mxfp4_bmm_config(x, weight) + # Keep AITER's small-M decode buckets; use the profiled prefill tiles. + if x.shape[1] > 256: + config["BLOCK_SIZE_M"] = 128 + config["BLOCK_SIZE_K"] = 128 + return config + + +def _run_tuned_mxfp4_bmm( + x: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + output: torch.Tensor, + config: dict, + *, + transpose_bm: bool, +) -> torch.Tensor: + return batched_gemm_a16wfp4( + x, + weight, + weight_scale, + y=output, + config=config, + transpose_bm=transpose_bm, + prequant=True, + y_scale=None, + dtype=torch.bfloat16, + ) + + +def _run_mxfp4_k_bmm( + x: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + output: torch.Tensor, +) -> None: + if x.shape[2] == 192 and weight.shape[1] == 512: + _run_tuned_mxfp4_bmm( + x, + weight, + weight_scale, + output, + _get_glm_mxfp4_k_bmm_config(x, weight), + transpose_bm=False, + ) + return + + batched_gemm_afp4wfp4_pre_quant( + x, + weight, + weight_scale, + torch.bfloat16, + output, + ) + + +def _run_mxfp4_v_bmm( + x: torch.Tensor, + weight: torch.Tensor, + weight_scale: torch.Tensor, + output: torch.Tensor, +) -> torch.Tensor: + if x.shape[2] == 512 and weight.shape[1] == 256: + return _run_tuned_mxfp4_bmm( + x, + weight, + weight_scale, + output, + _get_glm_mxfp4_v_bmm_config(x, weight), + transpose_bm=True, + ) + + batched_gemm_afp4wfp4_pre_quant( + x, + weight, + weight_scale, + torch.bfloat16, + output.transpose(0, 1), + ) + return output + + def should_defer_dsa_cp_kv_gather( *, dsa_prefill_cp: bool, @@ -516,6 +625,21 @@ def forward_absorb_prepare( expected_m, ) q_nope_out = q_nope_out[:, :expected_m, :] + elif _use_aiter_gfx95 and self.w_kc.dtype == torch.uint8: + x = q_nope.transpose(0, 1) + q_nope_out = torch.empty( + x.shape[0], + x.shape[1], + self.w_kc.shape[2], + device=x.device, + dtype=torch.bfloat16, + ) + _run_mxfp4_k_bmm( + x, + self.w_kc.transpose(-2, -1), + self.w_scale_k.transpose(-2, -1), + q_nope_out, + ) elif self.w_kc.dtype == torch.float8_e4m3fn: if _is_cpu: q_nope_out = torch.bmm( @@ -820,6 +944,21 @@ def forward_absorb_core( attn_bmm_output = ( attn_bmm_output[:, :expected_m, :].transpose(0, 1).flatten(1, 2) ) + elif _use_aiter_gfx95 and self.w_vc.dtype == torch.uint8: + x = attn_output.transpose(0, 1) + bmm_output = torch.empty( + x.shape[1], + x.shape[0], + self.w_vc.shape[2], + device=x.device, + dtype=torch.bfloat16, + ) + attn_bmm_output = _run_mxfp4_v_bmm( + x, + self.w_vc.transpose(-2, -1), + self.w_scale_v.transpose(-2, -1), + bmm_output, + ).flatten(1, 2) elif self.w_vc.dtype == torch.float8_e4m3fn: if _is_cpu: attn_bmm_output = torch.bmm( diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py index 928eb6352fc3..5d080c168248 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py @@ -46,6 +46,8 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_context import get_token_to_kv_pool from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mla import ( + _run_mxfp4_k_bmm, + _run_mxfp4_v_bmm, _select_local_dcp_heads_for_autotune, is_dcp_mla_decode_phase, is_mla_dcp_lse_base_on_e, @@ -116,7 +118,6 @@ def fused_qk_rmsnorm_bf16(q, q_weight, q_eps, k, k_weight, k_eps): ) from sglang.srt.layers.quantization.rocm_mxfp4_utils import ( - batched_gemm_afp4wfp4_pre_quant, fused_flatten_mxfp4_quant, fused_rms_mxfp4_quant, ) @@ -140,11 +141,10 @@ def rocm_absorb_q_bmm( device=x.device, dtype=torch.bfloat16, ) - batched_gemm_afp4wfp4_pre_quant( + _run_mxfp4_k_bmm( x, attn.w_kc.transpose(-2, -1), attn.w_scale_k.transpose(-2, -1), - torch.bfloat16, q_nope_out, ) else: @@ -192,14 +192,13 @@ def rocm_absorb_v_bmm( device=x.device, dtype=torch.bfloat16, ) - attn_bmm_output = _bmm_buf.transpose(0, 1) - batched_gemm_afp4wfp4_pre_quant( + _bmm_buf = _run_mxfp4_v_bmm( x, attn.w_vc.transpose(-2, -1), attn.w_scale_v.transpose(-2, -1), - torch.bfloat16, - attn_bmm_output, + _bmm_buf, ) + attn_bmm_output = _bmm_buf else: _bmm_buf = None if _use_aiter_gfx95 and attn.w_kc.dtype == torch.float8_e4m3fn: diff --git a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py index 9899f9fb8f53..fa55e86df46f 100644 --- a/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py +++ b/python/sglang/srt/models/deepseek_common/deepseek_weight_loader.py @@ -640,19 +640,29 @@ def post_load_weights( torch.bfloat16 ) - # GLM ships kv_b_proj as bf16, which falls back to torch.bmm. Quantize to - # per-tensor e4m3fn (not fnuz) to match forward_mla_rocm's dtype gate. - if ( + # GLM absorbed weights load as bf16. Keep the current per-tensor FP8 + # conversion as the default rollback, or preserve packed MXFP4 weights + # and per-head scales for the opt-in AITER absorbed-BMM backend. + is_glm_bf16_absorbed_weight = ( _use_aiter_gfx95 and self.config.architectures and self.config.architectures[0] == "GlmMoeDsaForCausalLM" and w.dtype == torch.bfloat16 - ): - w, self_attn.w_scale = input_to_float8(w, dtype=torch.float8_e4m3fn) + ) + use_mxfp4_mla_bmm = ( + is_glm_bf16_absorbed_weight and envs.SGLANG_USE_MXFP4_MLA_BMM.get() + ) + if use_mxfp4_mla_bmm: + w_kc, self_attn.w_scale_k, w_vc, self_attn.w_scale_v = ( + quark_post_load_weights(self_attn, w, "mxfp4") + ) + else: + if is_glm_bf16_absorbed_weight: + w, self_attn.w_scale = input_to_float8(w, dtype=torch.float8_e4m3fn) - w_kc, w_vc = w.unflatten( - 0, (-1, self_attn.qk_nope_head_dim + self_attn.v_head_dim) - ).split([self_attn.qk_nope_head_dim, self_attn.v_head_dim], dim=1) + w_kc, w_vc = w.unflatten( + 0, (-1, self_attn.qk_nope_head_dim + self_attn.v_head_dim) + ).split([self_attn.qk_nope_head_dim, self_attn.v_head_dim], dim=1) if ( _use_aiter_gfx95 diff --git a/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py b/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py new file mode 100644 index 000000000000..628374c1e49b --- /dev/null +++ b/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py @@ -0,0 +1,492 @@ +"""Unit tests for GLM packed-MXFP4 MLA absorbed-BMM selection.""" + +import unittest +from types import SimpleNamespace +from unittest import mock + +import torch + +from sglang.srt.environ import envs +from sglang.srt.models.deepseek_common import deepseek_weight_loader as weight_loader +from sglang.srt.models.deepseek_common.attention_forward_methods import ( + forward_mla, + forward_mla_rocm, +) +from sglang.test.ci.ci_register import register_amd_ci, register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") +register_amd_ci(est_time=10, suite="stage-b-test-1-gpu-small-amd-mi35x") + + +def _make_loader(): + self_attn = SimpleNamespace( + kv_b_proj=SimpleNamespace(weight=torch.randn(12, 8, dtype=torch.bfloat16)), + qk_nope_head_dim=4, + v_head_dim=2, + w_kc=None, + w_vc=None, + w_scale=None, + w_scale_k=None, + w_scale_v=None, + ) + loader = object.__new__(weight_loader.DeepseekV2WeightLoaderMixin) + loader.model = SimpleNamespace( + start_layer=0, + end_layer=1, + layers=[SimpleNamespace(self_attn=self_attn)], + ) + loader.config = SimpleNamespace( + architectures=["GlmMoeDsaForCausalLM"], num_hidden_layers=1 + ) + loader.quant_config = None + return loader, self_attn + + +class TestGlmMxfp4AbsorbedWeightSelection(CustomTestCase): + def test_toggle_defaults_off(self): + self.assertFalse(envs.SGLANG_USE_MXFP4_MLA_BMM.default) + + def test_flag_off_preserves_fp8_rollback(self): + loader, self_attn = _make_loader() + fp8_weight = torch.empty_like( + self_attn.kv_b_proj.weight, dtype=torch.float8_e4m3fn + ) + fp8_scale = torch.tensor(0.25) + with ( + envs.SGLANG_USE_MXFP4_MLA_BMM.override(False), + mock.patch.object(weight_loader, "_use_aiter_gfx95", True), + mock.patch( + "sglang.srt.layers.quantization.fp8_utils.input_to_float8", + return_value=(fp8_weight, fp8_scale), + ) as input_to_float8, + mock.patch.object( + weight_loader, "quark_post_load_weights", create=True + ) as quark_post_load_weights, + ): + loader.post_load_weights( + weight_names=["model.layers.0.self_attn.kv_b_proj"] + ) + input_to_float8.assert_called_once_with( + self_attn.kv_b_proj.weight, dtype=torch.float8_e4m3fn + ) + quark_post_load_weights.assert_not_called() + self.assertEqual(self_attn.w_kc.dtype, torch.float8_e4m3fn) + self.assertEqual(self_attn.w_vc.dtype, torch.float8_e4m3fn) + self.assertIs(self_attn.w_scale, fp8_scale) + self.assertEqual(self_attn.w_kc.stride(), (32, 1, 4)) + self.assertEqual(self_attn.w_vc.stride(), (16, 1, 8)) + + def test_flag_on_assigns_packed_weights_and_scales(self): + loader, self_attn = _make_loader() + w_kc = torch.arange(32, dtype=torch.uint8).view(2, 2, 8) + w_scale_k = torch.arange(16, dtype=torch.uint8).view(2, 1, 8) + w_vc = torch.arange(16, dtype=torch.uint8).view(2, 2, 4) + w_scale_v = torch.arange(4, dtype=torch.uint8).view(2, 2, 1) + with ( + envs.SGLANG_USE_MXFP4_MLA_BMM.override(True), + mock.patch.object(weight_loader, "_use_aiter_gfx95", True), + mock.patch.object( + weight_loader, + "quark_post_load_weights", + create=True, + return_value=(w_kc, w_scale_k, w_vc, w_scale_v), + ) as quark_post_load_weights, + mock.patch( + "sglang.srt.layers.quantization.fp8_utils.input_to_float8" + ) as input_to_float8, + ): + loader.post_load_weights( + weight_names=["model.layers.0.self_attn.kv_b_proj"] + ) + quark_post_load_weights.assert_called_once_with( + self_attn, self_attn.kv_b_proj.weight, "mxfp4" + ) + input_to_float8.assert_not_called() + self.assertTrue(torch.equal(self_attn.w_kc, w_kc)) + self.assertEqual(self_attn.w_kc.stride(), (16, 1, 2)) + self.assertTrue(torch.equal(self_attn.w_vc, w_vc.transpose(1, 2))) + self.assertEqual(self_attn.w_vc.stride(), (8, 1, 4)) + self.assertIs(self_attn.w_scale_k, w_scale_k) + self.assertIs(self_attn.w_scale_v, w_scale_v) + torch.testing.assert_close( + self_attn.w_kc.transpose(-2, -1), w_kc.transpose(-2, -1) + ) + torch.testing.assert_close( + self_attn.w_scale_k.transpose(-2, -1), w_scale_k.transpose(-2, -1) + ) + torch.testing.assert_close(self_attn.w_vc.transpose(-2, -1), w_vc) + torch.testing.assert_close( + self_attn.w_scale_v.transpose(-2, -1), w_scale_v.transpose(-2, -1) + ) + + +class TestMxfp4KDispatch(CustomTestCase): + def test_non_glm_uint8_geometry_uses_prequant_fallback(self): + x = torch.randn(2, 3, 8, dtype=torch.bfloat16) + weight = torch.zeros(2, 4, 4, dtype=torch.uint8) + scale = torch.zeros(2, 4, 1, dtype=torch.uint8) + output = torch.empty(2, 3, 4, dtype=torch.bfloat16) + + with ( + mock.patch.object( + forward_mla, "batched_gemm_afp4wfp4_pre_quant", create=True + ) as prequant_bmm, + mock.patch.object(forward_mla, "_run_tuned_mxfp4_bmm") as tuned_bmm, + ): + result = forward_mla._run_mxfp4_k_bmm(x, weight, scale, output) + + prequant_bmm.assert_called_once_with(x, weight, scale, torch.bfloat16, output) + tuned_bmm.assert_not_called() + self.assertIsNone(result) + + def test_glm_geometry_uses_safe_k_block_without_split_k(self): + x = torch.randn(2, 3, 192, dtype=torch.bfloat16) + weight = torch.zeros(2, 512, 96, dtype=torch.uint8) + scale = torch.zeros(2, 512, 6, dtype=torch.uint8) + output = torch.empty(2, 3, 512, dtype=torch.bfloat16) + tuned_config = {"BLOCK_SIZE_K": 256, "NUM_KSPLIT": 4} + + with ( + mock.patch.object( + forward_mla, + "_get_mxfp4_bmm_config", + create=True, + return_value=(tuned_config, None), + ), + mock.patch.object(forward_mla, "_run_tuned_mxfp4_bmm") as tuned_bmm, + mock.patch.object( + forward_mla, "batched_gemm_afp4wfp4_pre_quant", create=True + ) as prequant_bmm, + ): + result = forward_mla._run_mxfp4_k_bmm(x, weight, scale, output) + + args, kwargs = tuned_bmm.call_args + self.assertIs(args[0], x) + self.assertIs(args[1], weight) + self.assertIs(args[2], scale) + self.assertIs(args[3], output) + self.assertEqual(args[4]["NUM_KSPLIT"], 1) + self.assertEqual(args[4]["BLOCK_SIZE_K"], 64) + self.assertEqual(tuned_config["NUM_KSPLIT"], 4) + self.assertEqual(tuned_config["BLOCK_SIZE_K"], 256) + self.assertFalse(kwargs["transpose_bm"]) + prequant_bmm.assert_not_called() + self.assertIsNone(result) + + def test_glm_geometry_output_matches_bf16_reference(self): + x = torch.randn(2, 3, 192, dtype=torch.bfloat16) + weight = torch.randn(2, 512, 192, dtype=torch.bfloat16) + scale = torch.zeros(2, 512, 6, dtype=torch.uint8) + output = torch.empty(2, 3, 512, dtype=torch.bfloat16) + expected = torch.bmm(x, weight.transpose(-2, -1)) + + def reference_bmm(x, weight, _scale, output, _config, *, transpose_bm): + result = torch.bmm(x, weight.transpose(-2, -1)) + if transpose_bm: + result = result.transpose(0, 1) + output.copy_(result) + return output + + with ( + mock.patch.object( + forward_mla, + "_get_mxfp4_bmm_config", + create=True, + return_value=({"NUM_KSPLIT": 4}, None), + ), + mock.patch.object( + forward_mla, + "_run_tuned_mxfp4_bmm", + side_effect=reference_bmm, + ), + ): + result = forward_mla._run_mxfp4_k_bmm(x, weight, scale, output) + + self.assertIsNone(result) + torch.testing.assert_close(output, expected) + + +class TestRocmMxfp4AbsorbedBmmRoute(CustomTestCase): + def test_q_route_dispatches_transposed_tensors_to_mxfp4_helper(self): + q_nope = torch.randn(3, 2, 8, dtype=torch.bfloat16) + w_kc = torch.zeros(2, 4, 5, dtype=torch.uint8) + w_scale_k = torch.zeros(2, 1, 5, dtype=torch.uint8) + attn = SimpleNamespace(w_kc=w_kc, w_scale_k=w_scale_k) + + with ( + mock.patch.object(forward_mla_rocm, "_use_aiter_gfx95", True), + mock.patch.object(forward_mla_rocm, "_run_mxfp4_k_bmm") as run_mxfp4_k_bmm, + ): + result = forward_mla_rocm.rocm_absorb_q_bmm( + attn, q_nope, is_capture_mode=False + ) + + args, kwargs = run_mxfp4_k_bmm.call_args + self.assertEqual(kwargs, {}) + self.assertEqual(args[0].shape, (2, 3, 8)) + self.assertEqual(args[0].stride(), q_nope.transpose(0, 1).stride()) + self.assertEqual(args[0].data_ptr(), q_nope.data_ptr()) + self.assertEqual(args[1].shape, (2, 5, 4)) + self.assertEqual(args[1].stride(), w_kc.transpose(-2, -1).stride()) + self.assertEqual(args[1].data_ptr(), w_kc.data_ptr()) + self.assertEqual(args[2].shape, (2, 5, 1)) + self.assertEqual(args[2].stride(), w_scale_k.transpose(-2, -1).stride()) + self.assertEqual(args[2].data_ptr(), w_scale_k.data_ptr()) + self.assertIs(args[3], result) + self.assertEqual(result.shape, (2, 3, 5)) + self.assertEqual(result.dtype, torch.bfloat16) + + def test_v_route_flattens_batch_major_mxfp4_output_as_view(self): + attn_output = torch.randn(3, 2, 8, dtype=torch.bfloat16) + w_vc = torch.zeros(2, 4, 5, dtype=torch.uint8) + w_scale_v = torch.zeros(2, 1, 5, dtype=torch.uint8) + attn = SimpleNamespace( + w_vc=w_vc, + w_scale_v=w_scale_v, + o_proj=SimpleNamespace(weight=torch.empty(1, dtype=torch.bfloat16)), + ) + + with ( + mock.patch.object(forward_mla_rocm, "_use_aiter_gfx95", True), + mock.patch.object( + forward_mla_rocm, + "_run_mxfp4_v_bmm", + side_effect=lambda _x, _weight, _scale, output: output, + ) as run_mxfp4_v_bmm, + ): + result = forward_mla_rocm.rocm_absorb_v_bmm(attn, attn_output) + + args, kwargs = run_mxfp4_v_bmm.call_args + self.assertEqual(kwargs, {}) + self.assertEqual(args[0].shape, (2, 3, 8)) + self.assertEqual(args[0].stride(), attn_output.transpose(0, 1).stride()) + self.assertEqual(args[0].data_ptr(), attn_output.data_ptr()) + self.assertEqual(args[1].shape, (2, 5, 4)) + self.assertEqual(args[1].stride(), w_vc.transpose(-2, -1).stride()) + self.assertEqual(args[1].data_ptr(), w_vc.data_ptr()) + self.assertEqual(args[2].shape, (2, 5, 1)) + self.assertEqual(args[2].stride(), w_scale_v.transpose(-2, -1).stride()) + self.assertEqual(args[2].data_ptr(), w_scale_v.data_ptr()) + bmm_output = args[3] + self.assertEqual(bmm_output.shape, (3, 2, 5)) + self.assertTrue(bmm_output.is_contiguous()) + self.assertEqual(result.shape, (3, 10)) + self.assertEqual(result.dtype, torch.bfloat16) + self.assertEqual(result.data_ptr(), bmm_output.data_ptr()) + + +class TestTunedMxfp4Bmm(CustomTestCase): + def test_calls_aiter_public_wrapper_with_tuned_config(self): + cases = ( + { + "name": "k", + "packed_k": 96, + "config": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 64, + "NUM_KSPLIT": 1, + }, + "transpose_bm": False, + }, + { + "name": "v", + "packed_k": 256, + "config": { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 256, + "BLOCK_SIZE_K": 128, + "NUM_KSPLIT": 1, + }, + "transpose_bm": True, + }, + ) + + for case in cases: + with self.subTest(case=case["name"]): + batch, m, n = 2, 3, 4 + packed_k = case["packed_k"] + x = torch.empty(batch, m, 2 * packed_k, dtype=torch.bfloat16) + weight = torch.empty(batch, n, packed_k, dtype=torch.uint8) + scale = torch.empty(batch, n, packed_k // 32, dtype=torch.uint8) + output_shape = (m, batch, n) if case["transpose_bm"] else (batch, m, n) + output = torch.empty(output_shape, dtype=torch.bfloat16) + config = case["config"] + original_config = config.copy() + wrapper_result = mock.sentinel.wrapper_result + + with mock.patch.object( + forward_mla, + "batched_gemm_a16wfp4", + return_value=wrapper_result, + ) as batched_gemm: + result = forward_mla._run_tuned_mxfp4_bmm( + x, + weight, + scale, + output, + config, + transpose_bm=case["transpose_bm"], + ) + + self.assertEqual(config, original_config) + batched_gemm.assert_called_once_with( + x, + weight, + scale, + y=output, + config=original_config, + transpose_bm=case["transpose_bm"], + prequant=True, + y_scale=None, + dtype=torch.bfloat16, + ) + self.assertIs(batched_gemm.call_args.kwargs["config"], config) + self.assertIs(result, wrapper_result) + + +class TestMxfp4VDispatch(CustomTestCase): + def _fallback_inputs(self): + return ( + torch.randn(2, 3, 8, dtype=torch.bfloat16), + torch.zeros(2, 4, 4, dtype=torch.uint8), + torch.zeros(2, 4, 1, dtype=torch.uint8), + torch.empty(3, 2, 4, dtype=torch.bfloat16), + ) + + def _glm_inputs(self): + return ( + torch.randn(2, 3, 512, dtype=torch.bfloat16), + torch.zeros(2, 256, 256, dtype=torch.uint8), + torch.zeros(2, 256, 16, dtype=torch.uint8), + torch.empty(3, 2, 256, dtype=torch.bfloat16), + ) + + def test_non_glm_uint8_geometry_uses_prequant_fallback(self): + x, weight, scale, output = self._fallback_inputs() + with ( + mock.patch.object( + forward_mla, "batched_gemm_afp4wfp4_pre_quant", create=True + ) as prequant_bmm, + mock.patch.object(forward_mla, "_run_tuned_mxfp4_bmm") as tuned_bmm, + ): + result = forward_mla._run_mxfp4_v_bmm(x, weight, scale, output) + args, kwargs = prequant_bmm.call_args + self.assertEqual(kwargs, {}) + self.assertIs(args[0], x) + self.assertIs(args[1], weight) + self.assertIs(args[2], scale) + self.assertIs(args[3], torch.bfloat16) + self.assertEqual(args[4].shape, (2, 3, 4)) + self.assertEqual(args[4].stride(), output.transpose(0, 1).stride()) + self.assertEqual(args[4].data_ptr(), output.data_ptr()) + tuned_bmm.assert_not_called() + self.assertIs(result, output) + + def test_glm_geometry_uses_atom_batch_major_dispatch(self): + x, weight, scale, output = self._glm_inputs() + tuned_config = {"BLOCK_SIZE_K": 256, "NUM_KSPLIT": 4} + with ( + mock.patch.object( + forward_mla, + "_get_mxfp4_bmm_config", + create=True, + return_value=(tuned_config, None), + ), + mock.patch.object( + forward_mla, + "_run_tuned_mxfp4_bmm", + return_value=output, + ) as tuned_bmm, + mock.patch.object( + forward_mla, "batched_gemm_afp4wfp4_pre_quant", create=True + ) as prequant_bmm, + ): + result = forward_mla._run_mxfp4_v_bmm(x, weight, scale, output) + args, kwargs = tuned_bmm.call_args + self.assertIs(args[0], x) + self.assertIs(args[1], weight) + self.assertIs(args[2], scale) + self.assertIs(args[3], output) + self.assertEqual(args[4]["NUM_KSPLIT"], 1) + self.assertEqual(args[4]["BLOCK_SIZE_K"], 256) + self.assertTrue(kwargs["transpose_bm"]) + self.assertEqual(tuned_config["NUM_KSPLIT"], 4) + self.assertEqual(tuned_config["BLOCK_SIZE_K"], 256) + prequant_bmm.assert_not_called() + self.assertIs(result, output) + self.assertTrue(result.is_contiguous()) + + def test_glm_geometry_batch_major_output_matches_bf16_reference(self): + x = torch.randn(2, 3, 512, dtype=torch.bfloat16) + weight = torch.randn(2, 256, 512, dtype=torch.bfloat16) + scale = torch.zeros(2, 256, 16, dtype=torch.uint8) + output = torch.empty(3, 2, 256, dtype=torch.bfloat16) + expected = torch.bmm(x, weight.transpose(-2, -1)).transpose(0, 1) + + def reference_bmm(x, weight, _scale, output, _config, *, transpose_bm): + result = torch.bmm(x, weight.transpose(-2, -1)) + if transpose_bm: + result = result.transpose(0, 1) + output.copy_(result) + return output + + with ( + mock.patch.object( + forward_mla, + "_get_mxfp4_bmm_config", + create=True, + return_value=({"NUM_KSPLIT": 4}, None), + ), + mock.patch.object( + forward_mla, + "_run_tuned_mxfp4_bmm", + side_effect=reference_bmm, + ), + ): + result = forward_mla._run_mxfp4_v_bmm(x, weight, scale, output) + + self.assertIs(result, output) + torch.testing.assert_close(result, expected) + + +class TestMxfp4VConfig(CustomTestCase): + def _config(self, tokens): + x = torch.empty(16, tokens, 512, dtype=torch.bfloat16) + weight = torch.empty(16, 256, 256, dtype=torch.uint8) + tuned_config = { + "BLOCK_SIZE_M": 256, + "BLOCK_SIZE_N": 256, + "BLOCK_SIZE_K": 256, + "NUM_KSPLIT": 4, + } + with mock.patch.object( + forward_mla, + "_get_mxfp4_bmm_config", + create=True, + return_value=(tuned_config, None), + ): + config = forward_mla._get_glm_mxfp4_v_bmm_config(x, weight) + return config, tuned_config + + def test_large_prefill_uses_m128_k128_single_split(self): + config, tuned_config = self._config(8192) + self.assertEqual(config["BLOCK_SIZE_M"], 128) + self.assertEqual(config["BLOCK_SIZE_N"], 256) + self.assertEqual(config["BLOCK_SIZE_K"], 128) + self.assertEqual(config["NUM_KSPLIT"], 1) + self.assertEqual(tuned_config["BLOCK_SIZE_M"], 256) + self.assertEqual(tuned_config["BLOCK_SIZE_K"], 256) + self.assertEqual(tuned_config["NUM_KSPLIT"], 4) + + def test_short_shapes_preserve_aiter_bucket(self): + for tokens in (1, 16, 64, 128, 256): + with self.subTest(tokens=tokens): + config, _ = self._config(tokens) + self.assertEqual(config["BLOCK_SIZE_M"], 256) + self.assertEqual(config["NUM_KSPLIT"], 1) + + +if __name__ == "__main__": + unittest.main() From ecdc6fe2afe79b502ba7e96f718e899df7d941f7 Mon Sep 17 00:00:00 2001 From: Raiden-Makoto Date: Mon, 17 Aug 2026 10:12:45 -0700 Subject: [PATCH 2/3] test(amd): make ROCm wrapper mock CPU-safe Create the conditionally imported AITER symbol in the mock so CPU CI can exercise wrapper argument forwarding. --- test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py | 1 + 1 file changed, 1 insertion(+) diff --git a/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py b/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py index 628374c1e49b..8e2a5a4e2809 100644 --- a/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py +++ b/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py @@ -320,6 +320,7 @@ def test_calls_aiter_public_wrapper_with_tuned_config(self): forward_mla, "batched_gemm_a16wfp4", return_value=wrapper_result, + create=True, ) as batched_gemm: result = forward_mla._run_tuned_mxfp4_bmm( x, From e675b4707dd3febd3c706ab0e7ab86723dcd9f82 Mon Sep 17 00:00:00 2001 From: Raiden-Makoto Date: Mon, 17 Aug 2026 10:59:40 -0700 Subject: [PATCH 3/3] test(amd): patch the active weight-loader FP8 binding Mock the symbol imported into deepseek_weight_loader so CPU CI observes the GLM rollback conversion call after the upstream import refactor. --- .../unit/models/test_glm_mxfp4_absorbed_bmm.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py b/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py index 8e2a5a4e2809..1707a1c7d659 100644 --- a/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py +++ b/test/registered/unit/models/test_glm_mxfp4_absorbed_bmm.py @@ -56,8 +56,9 @@ def test_flag_off_preserves_fp8_rollback(self): with ( envs.SGLANG_USE_MXFP4_MLA_BMM.override(False), mock.patch.object(weight_loader, "_use_aiter_gfx95", True), - mock.patch( - "sglang.srt.layers.quantization.fp8_utils.input_to_float8", + mock.patch.object( + weight_loader, + "input_to_float8", return_value=(fp8_weight, fp8_scale), ) as input_to_float8, mock.patch.object( @@ -92,9 +93,7 @@ def test_flag_on_assigns_packed_weights_and_scales(self): create=True, return_value=(w_kc, w_scale_k, w_vc, w_scale_v), ) as quark_post_load_weights, - mock.patch( - "sglang.srt.layers.quantization.fp8_utils.input_to_float8" - ) as input_to_float8, + mock.patch.object(weight_loader, "input_to_float8") as input_to_float8, ): loader.post_load_weights( weight_names=["model.layers.0.self_attn.kv_b_proj"]