[NVIDIA] Support flashinfer Mega Moe (#31470)
Co-authored-by: djns99 <40156487+djns99@users.noreply.github.com> Co-authored-by: 云挚 <ningyunxiao.nyx@antgroup.com> Co-authored-by: Yangmin Li <yangminl@nvidia.com> Co-authored-by: Po-Han Huang (NVIDIA) <53919306+nvpohanh@users.noreply.github.com>
This commit is contained in:
co-authored by
djns99
云挚
Yangmin Li
Po-Han Huang
parent
c0b790cf7f
commit
1b77f498a0
@@ -0,0 +1,62 @@
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatcher
|
||||
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||
from sglang.srt.layers.moe.utils import FlashinferA2ADispatchType
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def test_empty_mxfp8_dispatch_uses_same_payload_dtype_as_nonempty_rank():
|
||||
class FakeMoeAlltoAll:
|
||||
def dispatch(self, _topk_ids, payloads, *_args, **_kwargs):
|
||||
self.payload_dtypes = [payload.dtype for payload in payloads]
|
||||
return payloads
|
||||
|
||||
dispatcher = object.__new__(FlashinferDispatcher)
|
||||
dispatcher.dispatch_type = FlashinferA2ADispatchType.MXFP8
|
||||
dispatcher.hidden_size = 128
|
||||
dispatcher.max_num_tokens = 0
|
||||
dispatcher.ep_size = 1
|
||||
dispatcher.invalid_token_expert_id = 8
|
||||
dispatcher.payload_in_workspace = False
|
||||
dispatcher.quant_config = {"use_mxfp8": True}
|
||||
dispatcher.moe_a2a = FakeMoeAlltoAll()
|
||||
|
||||
hidden_states = torch.empty((0, 128), dtype=torch.bfloat16)
|
||||
topk_output = StandardTopKOutput(
|
||||
topk_weights=torch.empty((0, 1), dtype=torch.float32),
|
||||
topk_ids=torch.empty((0, 1), dtype=torch.int32),
|
||||
router_logits=None,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.layers.moe.token_dispatcher.flashinfer.get_dp_global_num_tokens",
|
||||
return_value=None,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.layers.moe.token_dispatcher.flashinfer.is_dp_attention_enabled",
|
||||
return_value=False,
|
||||
),
|
||||
):
|
||||
output = dispatcher.dispatch(hidden_states, topk_output)
|
||||
|
||||
assert dispatcher.moe_a2a.payload_dtypes == [
|
||||
torch.float8_e4m3fn,
|
||||
torch.uint8,
|
||||
torch.int32,
|
||||
torch.float32,
|
||||
]
|
||||
assert output.hidden_states.dtype == torch.float8_e4m3fn
|
||||
assert output.hidden_states_scale.dtype == torch.uint8
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pytest
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
@@ -0,0 +1,220 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _load_megamoe_module(monkeypatch):
|
||||
"""Load the adapter with only its small import-time dependencies stubbed."""
|
||||
|
||||
class MoeQuantInfo:
|
||||
pass
|
||||
|
||||
class MoeRunnerConfig:
|
||||
pass
|
||||
|
||||
def register_fused_func(*_args, **_kwargs):
|
||||
return lambda fn: fn
|
||||
|
||||
fake_modules = {
|
||||
"sglang": types.ModuleType("sglang"),
|
||||
"sglang.srt": types.ModuleType("sglang.srt"),
|
||||
"sglang.srt.environ": types.ModuleType("sglang.srt.environ"),
|
||||
"sglang.srt.layers": types.ModuleType("sglang.srt.layers"),
|
||||
"sglang.srt.layers.moe": types.ModuleType("sglang.srt.layers.moe"),
|
||||
"sglang.srt.layers.moe.moe_runner": types.ModuleType(
|
||||
"sglang.srt.layers.moe.moe_runner"
|
||||
),
|
||||
"sglang.srt.layers.moe.moe_runner.base": types.ModuleType(
|
||||
"sglang.srt.layers.moe.moe_runner.base"
|
||||
),
|
||||
"sglang.srt.layers.moe.token_dispatcher": types.ModuleType(
|
||||
"sglang.srt.layers.moe.token_dispatcher"
|
||||
),
|
||||
"sglang.srt.runtime_context": types.ModuleType("sglang.srt.runtime_context"),
|
||||
"deep_gemm": types.ModuleType("deep_gemm"),
|
||||
"deep_gemm.utils": types.ModuleType("deep_gemm.utils"),
|
||||
"deep_gemm.utils.math": types.ModuleType("deep_gemm.utils.math"),
|
||||
}
|
||||
fake_modules["sglang.srt.environ"].envs = types.SimpleNamespace(
|
||||
SGLANG_FLASHINFER_MEGAMOE_MAX_TOKENS_PER_RANK=types.SimpleNamespace(
|
||||
get=lambda: 0
|
||||
),
|
||||
SGLANG_FLASHINFER_MEGAMOE_COMBINE_DTYPE=types.SimpleNamespace(
|
||||
get=lambda: "bf16"
|
||||
),
|
||||
SGLANG_FLASHINFER_MEGAMOE_IN_KERNEL_FC2_REDUCE=types.SimpleNamespace(
|
||||
get=lambda: False
|
||||
),
|
||||
)
|
||||
runtime_context = fake_modules["sglang.srt.runtime_context"]
|
||||
runtime_context.cutedsl_moe_max_num_tokens = lambda: 2048
|
||||
base = fake_modules["sglang.srt.layers.moe.moe_runner.base"]
|
||||
base.MoeQuantInfo = MoeQuantInfo
|
||||
base.MoeRunnerConfig = MoeRunnerConfig
|
||||
base.register_fused_func = register_fused_func
|
||||
token_dispatcher = fake_modules["sglang.srt.layers.moe.token_dispatcher"]
|
||||
|
||||
class StandardCombineInput:
|
||||
def __init__(self, *, hidden_states):
|
||||
self.hidden_states = hidden_states
|
||||
|
||||
token_dispatcher.StandardCombineInput = StandardCombineInput
|
||||
for name, module in fake_modules.items():
|
||||
monkeypatch.setitem(sys.modules, name, module)
|
||||
|
||||
module_path = (
|
||||
Path(__file__).resolve().parents[5]
|
||||
/ "python/sglang/srt/layers/moe/flashinfer_megamoe.py"
|
||||
)
|
||||
module_name = "sglang_flashinfer_megamoe_adapter_test"
|
||||
spec = importlib.util.spec_from_file_location(module_name, module_path)
|
||||
assert spec is not None and spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
monkeypatch.setitem(sys.modules, module_name, module)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def test_max_tokens_uses_runtime_context_accessor(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
assert module._resolve_max_tokens_per_rank() == 2048
|
||||
|
||||
runtime_context = sys.modules["sglang.srt.runtime_context"]
|
||||
runtime_context.cutedsl_moe_max_num_tokens = lambda: 0
|
||||
assert module._resolve_max_tokens_per_rank() == 1024
|
||||
|
||||
|
||||
def test_adapter_keeps_router_ids_int32(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
class FakeMoEEpTensors:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
fake_moe_ep = types.ModuleType("flashinfer.moe_ep")
|
||||
fake_moe_ep.MoEEpTensors = FakeMoEEpTensors
|
||||
fake_flashinfer = types.ModuleType("flashinfer")
|
||||
fake_flashinfer.moe_ep = fake_moe_ep
|
||||
monkeypatch.setitem(sys.modules, "flashinfer", fake_flashinfer)
|
||||
monkeypatch.setitem(sys.modules, "flashinfer.moe_ep", fake_moe_ep)
|
||||
|
||||
hidden_states = torch.randn((3, 4), dtype=torch.bfloat16)
|
||||
topk_ids = torch.tensor([[0, 1], [1, 0], [0, 1]], dtype=torch.int32)
|
||||
topk_weights = torch.randn((3, 2), dtype=torch.float32)
|
||||
output = torch.randn_like(hidden_states)
|
||||
|
||||
class Mega:
|
||||
_workspace = object()
|
||||
|
||||
def forward(self, tensors):
|
||||
self.tensors = tensors
|
||||
return output
|
||||
|
||||
mega = Mega()
|
||||
dispatch_output = types.SimpleNamespace(
|
||||
hidden_states=hidden_states,
|
||||
topk_output=types.SimpleNamespace(
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
),
|
||||
)
|
||||
quant_info = module.FlashInferMegaMoeQuantInfo(mega=mega)
|
||||
runner_config = types.SimpleNamespace(routed_scaling_factor=1.0)
|
||||
|
||||
result = module.run_flashinfer_megamoe(
|
||||
dispatch_output,
|
||||
quant_info,
|
||||
runner_config,
|
||||
)
|
||||
|
||||
assert mega.tensors.topk_ids.data_ptr() == topk_ids.data_ptr()
|
||||
assert mega.tensors.topk_ids.dtype == torch.int32
|
||||
assert result.hidden_states is output
|
||||
|
||||
|
||||
def test_adapter_requests_workspace_output_view(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
class FakeMoEEpTensors:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
fake_moe_ep = types.ModuleType("flashinfer.moe_ep")
|
||||
fake_moe_ep.MoEEpTensors = FakeMoEEpTensors
|
||||
fake_flashinfer = types.ModuleType("flashinfer")
|
||||
fake_flashinfer.moe_ep = fake_moe_ep
|
||||
monkeypatch.setitem(sys.modules, "flashinfer", fake_flashinfer)
|
||||
monkeypatch.setitem(sys.modules, "flashinfer.moe_ep", fake_moe_ep)
|
||||
|
||||
hidden_states = torch.randn((2, 4), dtype=torch.bfloat16)
|
||||
topk_ids = torch.tensor([[0, 1], [1, 0]], dtype=torch.int32)
|
||||
topk_weights = torch.ones((2, 2), dtype=torch.float32)
|
||||
output = torch.randn_like(hidden_states)
|
||||
|
||||
class Mega:
|
||||
supports_output_view = True
|
||||
_workspace = object()
|
||||
|
||||
def forward(self, tensors, *, return_workspace_view=False):
|
||||
self.tensors = tensors
|
||||
self.return_workspace_view = return_workspace_view
|
||||
return output
|
||||
|
||||
mega = Mega()
|
||||
dispatch_output = types.SimpleNamespace(
|
||||
hidden_states=hidden_states,
|
||||
topk_output=types.SimpleNamespace(
|
||||
topk_ids=topk_ids,
|
||||
topk_weights=topk_weights,
|
||||
),
|
||||
)
|
||||
|
||||
result = module.run_flashinfer_megamoe(
|
||||
dispatch_output,
|
||||
module.FlashInferMegaMoeQuantInfo(mega=mega),
|
||||
types.SimpleNamespace(routed_scaling_factor=1.0),
|
||||
)
|
||||
|
||||
assert result.hidden_states is output
|
||||
assert mega.tensors.topk_ids.data_ptr() == topk_ids.data_ptr()
|
||||
assert mega.tensors.topk_ids.dtype == torch.int32
|
||||
assert mega.return_workspace_view is True
|
||||
|
||||
|
||||
def test_capture_safe_ue8m0_pack_is_scoped(monkeypatch):
|
||||
module = _load_megamoe_module(monkeypatch)
|
||||
|
||||
dgm = sys.modules["deep_gemm.utils.math"]
|
||||
|
||||
def original(value):
|
||||
return value
|
||||
|
||||
dgm.pack_ue8m0_to_int = original
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: False)
|
||||
|
||||
with module._capture_safe_ue8m0_pack():
|
||||
assert dgm.pack_ue8m0_to_int is original
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "is_current_stream_capturing", lambda: True)
|
||||
|
||||
with module._capture_safe_ue8m0_pack():
|
||||
assert dgm.pack_ue8m0_to_int is not original
|
||||
packed = dgm.pack_ue8m0_to_int(torch.ones(4, dtype=torch.float32))
|
||||
assert packed.dtype == torch.int32
|
||||
|
||||
assert dgm.pack_ue8m0_to_int is original
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import pytest
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user