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:
co-authored by
Aurick Qiao
parent
99060191e7
commit
5dde6e8f02
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user