Fix MegaMoE buffer allocation and caching for effective SM budgets (#39223)

Co-authored-by: Aurick Qiao <6137920+aurickq@users.noreply.github.com>
This commit is contained in:
Aurick Qiao
2026-09-15 10:24:02 +08:00
committed by GitHub
co-authored by Aurick Qiao
parent 99060191e7
commit 5dde6e8f02
2 changed files with 104 additions and 20 deletions
+19 -17
View File
@@ -97,28 +97,30 @@ def _get_mega_moe_symm_buffer(
import deep_gemm import deep_gemm
mma_type = _mega_moe_mma_type() mma_type = _mega_moe_mma_type()
key = ( with _configure_mega_moe_deep_gemm_num_sms(deep_gemm):
id(group), key = (
num_max_tokens_per_rank, id(group),
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,
num_max_tokens_per_rank, num_max_tokens_per_rank,
num_experts,
num_topk, num_topk,
hidden, hidden,
intermediate_hidden, intermediate_hidden,
mma_type=mma_type, mma_type,
activation="swiglu", 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 return buf
@@ -4,7 +4,7 @@ import sys
import unittest import unittest
from contextlib import nullcontext from contextlib import nullcontext
from types import ModuleType, SimpleNamespace from types import ModuleType, SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, call, patch
import torch import torch
@@ -22,13 +22,24 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
def setUp(self): def setUp(self):
super().setUp() super().setUp()
mega_moe._MEGA_MOE_SYMM_BUFFER.clear() 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): def tearDown(self):
mega_moe._MEGA_MOE_SYMM_BUFFER.clear() mega_moe._MEGA_MOE_SYMM_BUFFER.clear()
super().tearDown() super().tearDown()
def test_mxf4_buffer_uses_typed_api(self): def test_mxf4_buffer_uses_typed_api(self):
deep_gemm = ModuleType("deep_gemm") deep_gemm = self.deep_gemm
expected_buffer = object() expected_buffer = object()
deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(return_value=expected_buffer) deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(return_value=expected_buffer)
group = object() group = object()
@@ -66,7 +77,7 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
self.assertEqual(mega_moe._mega_moe_mma_type(), expected) self.assertEqual(mega_moe._mega_moe_mma_type(), expected)
def test_buffer_cache_separates_mma_types(self): def test_buffer_cache_separates_mma_types(self):
deep_gemm = ModuleType("deep_gemm") deep_gemm = self.deep_gemm
expected_buffers = (object(), object()) expected_buffers = (object(), object())
deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(side_effect=expected_buffers) deep_gemm.get_symm_buffer_for_mega_moe = MagicMock(side_effect=expected_buffers)
group = object() group = object()
@@ -101,6 +112,63 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
["fp8xfp4", "mxf4xmxf4"], ["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): def test_mxf4_weight_transform_uses_matching_mma_type(self):
from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod
@@ -257,6 +325,20 @@ class TestDeepGemmMegaMoeApi(CustomTestCase):
torch.testing.assert_close(actual, expected) 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__": if __name__ == "__main__":
unittest.main() unittest.main()