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
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
@@ -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()