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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user