[NPU][bugfix] update low latency quantization input and update MXFP8 tests (#38831)
Co-authored-by: AndyLi429 <AndyLi429@noreply.gitcode.com>
This commit is contained in:
@@ -27,6 +27,7 @@ from sglang.srt.hardware_backend.npu.quantization.moe_methods import (
|
||||
w4a8_mxfp_gmm,
|
||||
)
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||
from sglang.srt.layers.moe.token_dispatcher import deepep
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8MoEMethod
|
||||
|
||||
|
||||
@@ -186,6 +187,435 @@ class TestPairPackMxfpActScale(unittest.TestCase):
|
||||
_pair_pack_mxfp_act_scale(torch.zeros(2, 3))
|
||||
|
||||
|
||||
class _LowLatencyBuffer:
|
||||
"""The MXFP8-era Buffer: bool flags, no quant_mode."""
|
||||
|
||||
def __init__(self):
|
||||
self.kwargs = None
|
||||
|
||||
def low_latency_dispatch(
|
||||
self,
|
||||
hidden_states,
|
||||
topk_ids,
|
||||
num_max_dispatch_tokens_per_rank,
|
||||
num_experts,
|
||||
*,
|
||||
use_fp8,
|
||||
use_mxfp4=False,
|
||||
use_mxfp8=False,
|
||||
**kwargs,
|
||||
):
|
||||
self.kwargs = {
|
||||
"use_fp8": use_fp8,
|
||||
"use_mxfp4": use_mxfp4,
|
||||
"use_mxfp8": use_mxfp8,
|
||||
**kwargs,
|
||||
}
|
||||
return torch.empty(0), torch.empty(0), object(), object(), object()
|
||||
|
||||
|
||||
class _LegacyLowLatencyBuffer:
|
||||
"""A DeepEP API version that predates the use_mxfp8 flag."""
|
||||
|
||||
def __init__(self):
|
||||
self.use_mxfp4 = None
|
||||
|
||||
def low_latency_dispatch(
|
||||
self,
|
||||
hidden_states,
|
||||
topk_ids,
|
||||
num_max_dispatch_tokens_per_rank,
|
||||
num_experts,
|
||||
*,
|
||||
use_fp8,
|
||||
use_mxfp4=False,
|
||||
topk_weights,
|
||||
async_finish,
|
||||
return_recv_hook,
|
||||
):
|
||||
self.use_mxfp4 = use_mxfp4
|
||||
return torch.empty(0), torch.empty(0), object(), object(), object()
|
||||
|
||||
|
||||
class _CudaLowLatencyBuffer:
|
||||
"""CUDA's Buffer API does not accept the NPU-only MXFP flags."""
|
||||
|
||||
def __init__(self):
|
||||
self.use_fp8 = None
|
||||
|
||||
def low_latency_dispatch(
|
||||
self,
|
||||
hidden_states,
|
||||
topk_ids,
|
||||
num_max_dispatch_tokens_per_rank,
|
||||
num_experts,
|
||||
*,
|
||||
use_fp8,
|
||||
round_scale=False,
|
||||
use_ue8m0=False,
|
||||
async_finish=False,
|
||||
return_recv_hook=False,
|
||||
):
|
||||
self.use_fp8 = use_fp8
|
||||
return torch.empty(0), torch.empty(0), object(), object(), object()
|
||||
|
||||
|
||||
class _CudaNormalBuffer:
|
||||
"""CUDA's normal Buffer API receives an already-quantized input tuple."""
|
||||
|
||||
def __init__(self):
|
||||
self.dispatched = False
|
||||
|
||||
def get_dispatch_layout(self, *args, **kwargs):
|
||||
return (
|
||||
torch.ones(1, dtype=torch.int32),
|
||||
None,
|
||||
torch.ones(2, dtype=torch.int32),
|
||||
torch.ones(1, 1, dtype=torch.bool),
|
||||
None,
|
||||
)
|
||||
|
||||
def dispatch(
|
||||
self,
|
||||
x,
|
||||
*,
|
||||
topk_idx,
|
||||
topk_weights,
|
||||
num_tokens_per_rank,
|
||||
num_tokens_per_rdma_rank,
|
||||
is_token_in_rank,
|
||||
num_tokens_per_expert,
|
||||
previous_event,
|
||||
async_finish,
|
||||
allocate_on_comm_stream,
|
||||
expert_alignment,
|
||||
config,
|
||||
):
|
||||
self.dispatched = True
|
||||
return torch.empty(0), torch.empty(0), torch.empty(0), [], object(), object()
|
||||
|
||||
|
||||
class _FlagNormalBuffer(_CudaNormalBuffer):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.quantization_kwargs = None
|
||||
|
||||
def dispatch(
|
||||
self,
|
||||
x,
|
||||
*,
|
||||
topk_idx,
|
||||
topk_weights,
|
||||
num_tokens_per_rank,
|
||||
num_tokens_per_rdma_rank,
|
||||
is_token_in_rank,
|
||||
num_tokens_per_expert,
|
||||
previous_event,
|
||||
async_finish,
|
||||
allocate_on_comm_stream,
|
||||
expert_alignment,
|
||||
config,
|
||||
use_fp8,
|
||||
use_mxfp4,
|
||||
use_mxfp8,
|
||||
):
|
||||
self.dispatched = True
|
||||
self.quantization_kwargs = {
|
||||
"use_fp8": use_fp8,
|
||||
"use_mxfp4": use_mxfp4,
|
||||
"use_mxfp8": use_mxfp8,
|
||||
}
|
||||
return torch.empty(0), torch.empty(0), torch.empty(0), [], object(), object()
|
||||
|
||||
|
||||
class _LegacyNormalBuffer:
|
||||
"""The pre-bool-flags DeepEP normal-dispatch API used by CI."""
|
||||
|
||||
def __init__(self):
|
||||
self.quant_mode = None
|
||||
|
||||
def get_dispatch_layout(self, *args, **kwargs):
|
||||
return (
|
||||
torch.ones(1, dtype=torch.int32),
|
||||
None,
|
||||
torch.ones(2, dtype=torch.int32),
|
||||
torch.ones(1, 1, dtype=torch.bool),
|
||||
None,
|
||||
)
|
||||
|
||||
def dispatch(
|
||||
self,
|
||||
x,
|
||||
*,
|
||||
topk_idx,
|
||||
topk_weights,
|
||||
num_tokens_per_rank,
|
||||
num_tokens_per_rdma_rank,
|
||||
is_token_in_rank,
|
||||
num_tokens_per_expert,
|
||||
previous_event,
|
||||
async_finish,
|
||||
allocate_on_comm_stream,
|
||||
expert_alignment,
|
||||
config,
|
||||
quant_mode,
|
||||
):
|
||||
self.quant_mode = quant_mode
|
||||
return torch.empty(0), torch.empty(0), torch.empty(0), [], object(), object()
|
||||
|
||||
|
||||
class _OpaqueNormalBuffer(_CudaNormalBuffer):
|
||||
"""The A3 pybind Buffer API whose dispatch signature hides quantization args."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.dispatch_kwargs = None
|
||||
|
||||
def dispatch(self, *args, **kwargs):
|
||||
self.dispatched = True
|
||||
self.dispatch_kwargs = kwargs
|
||||
return torch.empty(0), torch.empty(0), torch.empty(0), [], object(), object()
|
||||
|
||||
|
||||
class TestDeepEPLowLatencyMxfp8Dispatch(unittest.TestCase):
|
||||
def test_mxfp4_output_dtype_enables_only_mxfp4(self):
|
||||
dispatcher = object.__new__(deepep._DeepEPDispatcherImplBase)
|
||||
|
||||
with patch.object(
|
||||
deepep,
|
||||
"get_deepep_output_dtype",
|
||||
return_value=deepep.DispatcherOutputDtype.MXFP4,
|
||||
):
|
||||
dispatcher.set_deepep_dispatcher_dtype()
|
||||
|
||||
self.assertFalse(dispatcher.use_fp8)
|
||||
self.assertTrue(dispatcher.use_mxfp4)
|
||||
self.assertFalse(dispatcher.use_mxfp8)
|
||||
|
||||
@staticmethod
|
||||
def _dispatcher(quant_mode, buffer):
|
||||
dispatcher = object.__new__(deepep._DeepEPDispatcherImplLowLatency)
|
||||
dispatcher.quant_config = {}
|
||||
dispatcher.use_fp8 = False
|
||||
dispatcher.use_mxfp4 = False
|
||||
dispatcher.use_mxfp8 = quant_mode == "mxfp8"
|
||||
dispatcher.use_nvfp4 = False
|
||||
dispatcher.num_max_dispatch_tokens_per_rank = 2
|
||||
dispatcher.num_experts = 2
|
||||
dispatcher.return_recv_hook = False
|
||||
dispatcher._get_buffer = lambda: buffer
|
||||
return dispatcher
|
||||
|
||||
def test_mxfp8_passes_the_mxfp8_flag_without_ue8m0(self):
|
||||
buffer = _LowLatencyBuffer()
|
||||
dispatcher = self._dispatcher("mxfp8", buffer)
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertFalse(buffer.kwargs["use_fp8"])
|
||||
self.assertTrue(buffer.kwargs["use_mxfp8"])
|
||||
self.assertNotIn("use_ue8m0", buffer.kwargs)
|
||||
|
||||
def test_mxfp4_passes_the_mxfp4_flag(self):
|
||||
buffer = _LowLatencyBuffer()
|
||||
dispatcher = self._dispatcher("mxfp4", buffer)
|
||||
dispatcher.use_mxfp4 = True
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertFalse(buffer.kwargs["use_fp8"])
|
||||
self.assertTrue(buffer.kwargs["use_mxfp4"])
|
||||
self.assertFalse(buffer.kwargs["use_mxfp8"])
|
||||
|
||||
def test_bf16_omits_unsupported_mxfp8_flag_for_legacy_buffer(self):
|
||||
buffer = _LegacyLowLatencyBuffer()
|
||||
dispatcher = self._dispatcher("bf16", buffer)
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertFalse(buffer.use_mxfp4)
|
||||
|
||||
def test_normal_dispatch_passes_quantization_flags(self):
|
||||
dispatcher = object.__new__(deepep._DeepEPDispatcherImplNormal)
|
||||
dispatcher.num_experts = 2
|
||||
dispatcher.async_finish = False
|
||||
dispatcher.use_fp8 = False
|
||||
dispatcher.use_mxfp4 = False
|
||||
dispatcher.use_mxfp8 = True
|
||||
buffer = _FlagNormalBuffer()
|
||||
dispatcher._get_buffer = lambda: buffer
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
patch.object(
|
||||
deepep.DeepEPConfig,
|
||||
"get_instance",
|
||||
return_value=SimpleNamespace(normal_dispatch_config=None),
|
||||
),
|
||||
patch.object(
|
||||
deepep,
|
||||
"get_global_expert_distribution_recorder",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
None,
|
||||
)
|
||||
|
||||
self.assertTrue(buffer.quantization_kwargs["use_mxfp8"])
|
||||
self.assertFalse(buffer.quantization_kwargs["use_fp8"])
|
||||
self.assertFalse(buffer.quantization_kwargs["use_mxfp4"])
|
||||
|
||||
def test_normal_dispatch_uses_legacy_quant_mode_when_flags_are_unsupported(self):
|
||||
dispatcher = object.__new__(deepep._DeepEPDispatcherImplNormal)
|
||||
dispatcher.num_experts = 2
|
||||
dispatcher.async_finish = False
|
||||
dispatcher.use_fp8 = False
|
||||
dispatcher.use_mxfp4 = False
|
||||
dispatcher.use_mxfp8 = False
|
||||
buffer = _LegacyNormalBuffer()
|
||||
dispatcher._get_buffer = lambda: buffer
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
patch.object(
|
||||
deepep.DeepEPConfig,
|
||||
"get_instance",
|
||||
return_value=SimpleNamespace(normal_dispatch_config=None),
|
||||
),
|
||||
patch.object(
|
||||
deepep,
|
||||
"get_global_expert_distribution_recorder",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
None,
|
||||
)
|
||||
|
||||
self.assertEqual(buffer.quant_mode, "bf16")
|
||||
|
||||
def test_normal_dispatch_keeps_a3_legacy_path_for_opaque_signature(self):
|
||||
dispatcher = object.__new__(deepep._DeepEPDispatcherImplNormal)
|
||||
dispatcher.num_experts = 2
|
||||
dispatcher.async_finish = False
|
||||
dispatcher.use_fp8 = True
|
||||
dispatcher.use_mxfp4 = False
|
||||
dispatcher.use_mxfp8 = False
|
||||
buffer = _OpaqueNormalBuffer()
|
||||
dispatcher._get_buffer = lambda: buffer
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", True),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
patch.object(
|
||||
deepep.DeepEPConfig,
|
||||
"get_instance",
|
||||
return_value=SimpleNamespace(normal_dispatch_config=None),
|
||||
),
|
||||
patch.object(
|
||||
deepep,
|
||||
"get_global_expert_distribution_recorder",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
None,
|
||||
)
|
||||
|
||||
self.assertTrue(buffer.dispatched)
|
||||
self.assertNotIn("use_fp8", buffer.dispatch_kwargs)
|
||||
self.assertNotIn("use_mxfp4", buffer.dispatch_kwargs)
|
||||
self.assertNotIn("use_mxfp8", buffer.dispatch_kwargs)
|
||||
self.assertNotIn("quant_mode", buffer.dispatch_kwargs)
|
||||
|
||||
def test_cuda_normal_dispatch_omits_npu_quantization_flags(self):
|
||||
dispatcher = object.__new__(deepep._DeepEPDispatcherImplNormal)
|
||||
dispatcher.num_experts = 2
|
||||
dispatcher.async_finish = False
|
||||
dispatcher.use_fp8 = True
|
||||
dispatcher.use_mxfp4 = True
|
||||
dispatcher.use_mxfp8 = True
|
||||
buffer = _CudaNormalBuffer()
|
||||
dispatcher._get_buffer = lambda: buffer
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", False),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
patch.object(
|
||||
deepep.DeepEPConfig,
|
||||
"get_instance",
|
||||
return_value=SimpleNamespace(normal_dispatch_config=None),
|
||||
),
|
||||
patch.object(
|
||||
deepep,
|
||||
"get_global_expert_distribution_recorder",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
None,
|
||||
)
|
||||
|
||||
self.assertTrue(buffer.dispatched)
|
||||
|
||||
def test_cuda_low_latency_dispatch_omits_npu_mxfp_flags(self):
|
||||
buffer = _CudaLowLatencyBuffer()
|
||||
dispatcher = self._dispatcher("mxfp8", buffer)
|
||||
dispatcher.use_fp8 = True
|
||||
dispatcher.use_mxfp4 = True
|
||||
|
||||
with (
|
||||
patch.object(deepep, "_is_npu", False),
|
||||
patch.object(deepep, "_deepep_precompile_tp_barrier"),
|
||||
):
|
||||
dispatcher._dispatch_core(
|
||||
torch.zeros(1, 64),
|
||||
torch.zeros(1, 1, dtype=torch.int64),
|
||||
torch.ones(1, 1),
|
||||
)
|
||||
|
||||
self.assertTrue(buffer.use_fp8)
|
||||
|
||||
|
||||
class TestW4A8MxfpGmmInputScale(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.input = torch.randn(2, 64)
|
||||
|
||||
Reference in New Issue
Block a user