From 5dde6e8f02e0992525f914a52868b89c7d56f64a Mon Sep 17 00:00:00 2001 From: Aurick Qiao Date: Mon, 14 Sep 2026 19:24:02 -0700 Subject: [PATCH] Fix MegaMoE buffer allocation and caching for effective SM budgets (#39223) Co-authored-by: Aurick Qiao <6137920+aurickq@users.noreply.github.com> --- python/sglang/srt/layers/moe/mega_moe.py | 36 ++++---- .../layers/moe/test_mega_moe_deepgemm_api.py | 88 ++++++++++++++++++- 2 files changed, 104 insertions(+), 20 deletions(-) diff --git a/python/sglang/srt/layers/moe/mega_moe.py b/python/sglang/srt/layers/moe/mega_moe.py index 1f3f83344..20a94b338 100644 --- a/python/sglang/srt/layers/moe/mega_moe.py +++ b/python/sglang/srt/layers/moe/mega_moe.py @@ -97,28 +97,30 @@ def _get_mega_moe_symm_buffer( import deep_gemm mma_type = _mega_moe_mma_type() - key = ( - id(group), - num_max_tokens_per_rank, - num_experts, - num_topk, - hidden, - intermediate_hidden, - mma_type, - ) - buf = _MEGA_MOE_SYMM_BUFFER.get(key) - if buf is None: - buf = deep_gemm.get_symm_buffer_for_mega_moe( - group, - num_experts, + with _configure_mega_moe_deep_gemm_num_sms(deep_gemm): + key = ( + id(group), num_max_tokens_per_rank, + num_experts, num_topk, hidden, intermediate_hidden, - mma_type=mma_type, - activation="swiglu", + mma_type, + deep_gemm.get_num_sms(), ) - _MEGA_MOE_SYMM_BUFFER[key] = buf + buf = _MEGA_MOE_SYMM_BUFFER.get(key) + if buf is None: + buf = deep_gemm.get_symm_buffer_for_mega_moe( + group, + num_experts, + num_max_tokens_per_rank, + num_topk, + hidden, + intermediate_hidden, + mma_type=mma_type, + activation="swiglu", + ) + _MEGA_MOE_SYMM_BUFFER[key] = buf return buf diff --git a/test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py b/test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py index 17232b826..e345d00f1 100644 --- a/test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py +++ b/test/registered/unit/layers/moe/test_mega_moe_deepgemm_api.py @@ -4,7 +4,7 @@ import sys import unittest from contextlib import nullcontext from types import ModuleType, SimpleNamespace -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock, call, patch import torch @@ -22,13 +22,24 @@ class TestDeepGemmMegaMoeApi(CustomTestCase): def setUp(self): super().setUp() mega_moe._MEGA_MOE_SYMM_BUFFER.clear() + self.deep_gemm = ModuleType("deep_gemm") + self.deep_gemm.num_sms = 132 + self.deep_gemm.get_num_sms = MagicMock( + side_effect=lambda: self.deep_gemm.num_sms + ) + self.deep_gemm.set_num_sms = MagicMock( + side_effect=lambda num_sms: setattr(self.deep_gemm, "num_sms", num_sms) + ) + max_num_sms = patch.object(mega_moe, "_mega_moe_max_num_sms", return_value=130) + self.max_num_sms = max_num_sms.start() + self.addCleanup(max_num_sms.stop) def tearDown(self): mega_moe._MEGA_MOE_SYMM_BUFFER.clear() super().tearDown() def test_mxf4_buffer_uses_typed_api(self): - deep_gemm = ModuleType("deep_gemm") + deep_gemm = self.deep_gemm expected_buffer = object() deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(return_value=expected_buffer) group = object() @@ -66,7 +77,7 @@ class TestDeepGemmMegaMoeApi(CustomTestCase): self.assertEqual(mega_moe._mega_moe_mma_type(), expected) def test_buffer_cache_separates_mma_types(self): - deep_gemm = ModuleType("deep_gemm") + deep_gemm = self.deep_gemm expected_buffers = (object(), object()) deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(side_effect=expected_buffers) group = object() @@ -101,6 +112,63 @@ class TestDeepGemmMegaMoeApi(CustomTestCase): ["fp8xfp4", "mxf4xmxf4"], ) + def test_buffer_cache_uses_effective_sm_budget(self): + deep_gemm = self.deep_gemm + deep_gemm.get_symm_buffer_for_mega_moe = MagicMock( + side_effect=lambda *_args, **_kwargs: SimpleNamespace( + num_sms=deep_gemm.get_num_sms() + ) + ) + group = object() + buffers = [] + + for current_num_sms, expected_num_sms in ( + (132, 130), + (131, 130), + (96, 96), + (97, 96), + (132, 130), + ): + with self.subTest(current_num_sms=current_num_sms): + deep_gemm.set_num_sms(current_num_sms) + buf = self._get_test_buffer(group) + buffers.append(buf) + self.assertEqual(buf.num_sms, expected_num_sms) + self.assertEqual(deep_gemm.get_num_sms(), current_num_sms) + with mega_moe._configure_mega_moe_deep_gemm_num_sms(deep_gemm): + self.assertEqual(buf.num_sms, deep_gemm.get_num_sms()) + self.assertEqual(deep_gemm.get_num_sms(), current_num_sms) + + self.assertIs(buffers[0], buffers[1]) + self.assertIs(buffers[0], buffers[4]) + self.assertIs(buffers[2], buffers[3]) + self.assertIsNot(buffers[0], buffers[2]) + self.assertEqual(deep_gemm.get_symm_buffer_for_mega_moe.call_count, 2) + + def test_buffer_allocation_without_sm_override(self): + deep_gemm = self.deep_gemm + self.max_num_sms.return_value = None + deep_gemm.get_symm_buffer_for_mega_moe = MagicMock( + side_effect=lambda *_args, **_kwargs: deep_gemm.get_num_sms() + ) + + buf = self._get_test_buffer(object()) + self.assertEqual(buf, 132) + deep_gemm.set_num_sms.assert_not_called() + + def test_buffer_allocation_failure_restores_sm_budget(self): + deep_gemm = self.deep_gemm + deep_gemm.get_symm_buffer_for_mega_moe = MagicMock( + side_effect=RuntimeError("allocation failed") + ) + + with self.assertRaisesRegex(RuntimeError, "allocation failed"): + self._get_test_buffer(object()) + + self.assertEqual(deep_gemm.get_num_sms(), 132) + self.assertEqual(deep_gemm.set_num_sms.call_args_list, [call(130), call(132)]) + self.assertEqual(mega_moe._MEGA_MOE_SYMM_BUFFER, {}) + def test_mxf4_weight_transform_uses_matching_mma_type(self): from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod @@ -257,6 +325,20 @@ class TestDeepGemmMegaMoeApi(CustomTestCase): torch.testing.assert_close(actual, expected) + def _get_test_buffer(self, group): + with ( + patch.dict(sys.modules, {"deep_gemm": self.deep_gemm}), + patch.object(mega_moe, "_mega_moe_mma_type", return_value="fp8xfp4"), + ): + return mega_moe._get_mega_moe_symm_buffer( + group, + num_experts=8, + num_max_tokens_per_rank=64, + num_topk=2, + hidden=128, + intermediate_hidden=256, + ) + if __name__ == "__main__": unittest.main()