feat(humming): FP8 DeepEP dispatch for humming MoE backend (#31429)

This commit is contained in:
guzekai01
2026-08-24 18:59:07 +08:00
committed by GitHub
parent 46b92b22e2
commit 21258b7a35
6 changed files with 543 additions and 38 deletions
+1 -1
View File
@@ -35,7 +35,7 @@ dependencies = [
"flash-attn-4>=4.0.0b18",
"flashinfer_python[cu13]==0.6.17", # keep it aligned with jit-cache version in Dockerfile
"gguf",
"humming-kernels[cu13]==0.1.10",
"humming-kernels[cu13]==0.1.12",
"interegular",
"IPython",
"kernels>=0.14.1,<0.15",
+50 -13
View File
@@ -2021,22 +2021,12 @@ def fp8_per_token_to_per_tensor_quant_triton(
)
def moe_permute(
def _moe_permute_rows(
inputs: torch.Tensor,
topk_ids: torch.Tensor,
num_experts: int,
use_int64_offset: bool = False,
is_ep: bool = False,
src2dst: torch.Tensor,
outputs: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
from sglang.kernels.ops.moe.moe_permute_prepare import moe_permute_prepare
expert_offsets, src2dst = moe_permute_prepare(
topk_ids=topk_ids,
num_experts=num_experts,
use_int64_offset=use_int64_offset,
is_ep=is_ep,
)
) -> torch.Tensor:
output_shape = (topk_ids.nelement(), inputs.size(-1))
if outputs is None:
outputs = torch.empty(output_shape, dtype=inputs.dtype, device=inputs.device)
@@ -2056,9 +2046,56 @@ def moe_permute(
BLOCK_SIZE=512,
)
return outputs
def moe_permute(
inputs: torch.Tensor,
topk_ids: torch.Tensor,
num_experts: int,
use_int64_offset: bool = False,
is_ep: bool = False,
outputs: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
from sglang.kernels.ops.moe.moe_permute_prepare import moe_permute_prepare
expert_offsets, src2dst = moe_permute_prepare(
topk_ids=topk_ids,
num_experts=num_experts,
use_int64_offset=use_int64_offset,
is_ep=is_ep,
)
outputs = _moe_permute_rows(inputs, topk_ids, src2dst, outputs)
return outputs, src2dst, expert_offsets
def moe_permute_with_scale(
inputs: torch.Tensor,
input_scale: torch.Tensor,
topk_ids: torch.Tensor,
num_experts: int,
use_int64_offset: bool = False,
is_ep: bool = False,
outputs: torch.Tensor | None = None,
scale_outputs: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
if inputs.size(0) != input_scale.size(0):
raise ValueError("inputs and input_scale must have the same number of rows")
outputs, src2dst, expert_offsets = moe_permute(
inputs=inputs,
topk_ids=topk_ids,
num_experts=num_experts,
use_int64_offset=use_int64_offset,
is_ep=is_ep,
outputs=outputs,
)
scale_outputs = _moe_permute_rows(input_scale, topk_ids, src2dst, scale_outputs)
return outputs, scale_outputs, src2dst, expert_offsets
def moe_unpermute(
inputs: torch.Tensor,
src2dst: torch.Tensor,
@@ -9,7 +9,11 @@ from weakref import WeakValueDictionary
import torch
from sglang.kernels.ops.moe.ep_moe_kernels import moe_permute, moe_unpermute
from sglang.kernels.ops.moe.ep_moe_kernels import (
moe_permute,
moe_permute_with_scale,
moe_unpermute,
)
from sglang.kernels.ops.moe.moe_fused_mul_sum import moe_fused_mul_sum
from sglang.srt.environ import envs
from sglang.srt.layers.moe.moe_runner.base import (
@@ -73,6 +77,7 @@ class HummingRunnerInput(RunnerInput):
gemm_type: HummingGemmType
expert_num_tokens: torch.Tensor | None = None
expected_m: int | None = None
hidden_states_scale: torch.Tensor | None = None
apply_routed_scaling_factor: bool = True
@property
@@ -103,6 +108,7 @@ def humming_moe_runner_core_run(
topk_ids: torch.Tensor,
expert_num_tokens: torch.Tensor | None = None,
expected_m: int | None = None,
hidden_states_scale: torch.Tensor | None = None,
apply_routed_scaling_factor: bool = True,
) -> torch.Tensor:
runner = HummingRunnerCore.runner_cores[moe_runner_id]
@@ -111,6 +117,7 @@ def humming_moe_runner_core_run(
hidden_states=hidden_states,
topk_ids=topk_ids,
topk_weights=topk_weights,
hidden_states_scale=hidden_states_scale,
apply_routed_scaling_factor=apply_routed_scaling_factor,
)
elif gemm_type == "grouped_contiguous":
@@ -118,6 +125,7 @@ def humming_moe_runner_core_run(
hidden_states=hidden_states,
topk_ids=topk_ids,
topk_weights=topk_weights,
hidden_states_scale=hidden_states_scale,
apply_routed_scaling_factor=apply_routed_scaling_factor,
)
elif gemm_type == "grouped_masked":
@@ -128,6 +136,7 @@ def humming_moe_runner_core_run(
topk_weights=topk_weights,
expected_m=expected_m,
expert_num_tokens=expert_num_tokens,
hidden_states_scale=hidden_states_scale,
)
else:
raise ValueError(f"Unknown gemm type: {gemm_type}")
@@ -396,6 +405,64 @@ class HummingRunnerCore(MoeRunnerCore):
else:
raise ValueError(f"Unsupported activation: {self.activation}")
def _grouped_masked_act_quant(
self,
gate_up: torch.Tensor,
expert_num_tokens: torch.Tensor,
buffers: dict[str, torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor | None]:
num_experts, max_tokens, two_i = gate_up.shape
intermediate = two_i // 2
w2_meta = self.layer.humming_metas["w2"]
groups = intermediate // 128
num_threads = intermediate // 8
use_fused_masked_act_quant = (
w2_meta.a_dtype == dtypes.float8e4m3
and w2_meta.input_scale_group_size == 128
and self.activation == "silu"
and gate_up.dtype == torch.bfloat16
and intermediate % 256 == 0
and num_threads <= 1024
and num_experts <= min(256, num_threads)
)
if use_fused_masked_act_quant:
from sglang.kernels.ops.attention.dsv4.moe import (
silu_and_mul_masked_post_quant,
)
down_input = torch.empty(
(num_experts, max_tokens, intermediate),
dtype=torch.float8_e4m3fn,
device=gate_up.device,
)
down_scale = torch.empty(
(num_experts, max_tokens, groups),
dtype=torch.float32,
device=gate_up.device,
)
silu_and_mul_masked_post_quant(
gate_up,
down_input,
down_scale,
128,
expert_num_tokens,
scale_ue8m0=False,
topk=self.config.top_k,
swiglu_limit=self.swiglu_limit,
)
return down_input.flatten(0, 1), down_scale.flatten(0, 1)
self.apply_activation(
inputs=buffers["gate_up_output"],
outputs=buffers["activation_output"],
)
return HummingMethod.may_quant_input(
layer=self.layer,
inputs=buffers["activation_output"],
quanted_input=buffers.get("quanted_down_input"),
sublayer_name="w2",
)
def run(
self,
runner_input: HummingRunnerInput,
@@ -406,7 +473,9 @@ class HummingRunnerCore(MoeRunnerCore):
self.layer = quant_info.layer
if runner_input.hidden_states.size(0) == 0:
return HummingRunnerOutput(
hidden_states=torch.empty_like(runner_input.hidden_states)
hidden_states=torch.empty_like(
runner_input.hidden_states, dtype=self.config.params_dtype
)
)
# To make it compatible with dynamic shapes in torch.compile,
@@ -420,6 +489,7 @@ class HummingRunnerCore(MoeRunnerCore):
topk_ids=runner_input.topk_ids,
expected_m=runner_input.expected_m,
expert_num_tokens=runner_input.expert_num_tokens,
hidden_states_scale=runner_input.hidden_states_scale,
apply_routed_scaling_factor=runner_input.apply_routed_scaling_factor,
)
@@ -471,9 +541,14 @@ class HummingRunnerCore(MoeRunnerCore):
hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
hidden_states_scale: torch.Tensor | None = None,
apply_routed_scaling_factor: bool = True,
):
hidden_states = hidden_states.view(-1, hidden_states.size(-1))
if hidden_states_scale is not None:
hidden_states_scale = hidden_states_scale.reshape(
-1, hidden_states_scale.size(-1)
).contiguous()
buffers = self.prepare_buffers(
hidden_states=hidden_states,
topk_ids=topk_ids,
@@ -485,6 +560,7 @@ class HummingRunnerCore(MoeRunnerCore):
inputs, input_scale = HummingMethod.may_quant_input(
layer=self.layer,
inputs=hidden_states,
input_scale=hidden_states_scale,
quanted_input=buffers.get("quanted_gate_up_input", None),
sublayer_name="w13",
)
@@ -538,6 +614,7 @@ class HummingRunnerCore(MoeRunnerCore):
hidden_states: torch.Tensor,
topk_weights: torch.Tensor,
topk_ids: torch.Tensor,
hidden_states_scale: torch.Tensor | None = None,
apply_routed_scaling_factor: bool = True,
):
configs = self.get_humming_gemm_configs(HummingGemmType.GROUPED_CONTIGUOUS)
@@ -549,16 +626,36 @@ class HummingRunnerCore(MoeRunnerCore):
gemm_type=HummingGemmType.GROUPED_CONTIGUOUS,
)
hidden_states, src2dst, expert_first_token_offset = moe_permute(
inputs=hidden_states,
topk_ids=topk_ids,
num_experts=self.num_experts,
is_ep=self.num_experts != self.global_num_experts,
)
is_ep = self.num_experts != self.global_num_experts
if hidden_states_scale is None:
hidden_states, src2dst, expert_first_token_offset = moe_permute(
inputs=hidden_states,
topk_ids=topk_ids,
num_experts=self.num_experts,
is_ep=is_ep,
)
else:
hidden_states_scale = hidden_states_scale.reshape(
-1, hidden_states_scale.size(-1)
).contiguous()
(
hidden_states,
hidden_states_scale,
src2dst,
expert_first_token_offset,
) = moe_permute_with_scale(
inputs=hidden_states,
input_scale=hidden_states_scale,
topk_ids=topk_ids,
num_experts=self.num_experts,
is_ep=is_ep,
outputs=buffers["quanted_gate_up_input"],
)
inputs, input_scale = HummingMethod.may_quant_input(
layer=self.layer,
inputs=hidden_states,
input_scale=hidden_states_scale,
quanted_input=buffers.get("quanted_gate_up_input", None),
sublayer_name="w13",
)
@@ -620,10 +717,15 @@ class HummingRunnerCore(MoeRunnerCore):
topk_weights: torch.Tensor,
expert_num_tokens: torch.Tensor,
expected_m: int,
hidden_states_scale: torch.Tensor | None = None,
):
configs = self.get_humming_gemm_configs(HummingGemmType.GROUPED_MASKED)
valid_shape_m = self.estimate_local_valid_shape_m(topk_ids, expected_m)
hidden_states = hidden_states.view(-1, hidden_states.size(-1))
if hidden_states_scale is not None:
hidden_states_scale = hidden_states_scale.reshape(
-1, hidden_states_scale.size(-1)
).contiguous()
buffers = self.prepare_buffers(
hidden_states=hidden_states,
@@ -634,6 +736,7 @@ class HummingRunnerCore(MoeRunnerCore):
inputs, input_scale = HummingMethod.may_quant_input(
layer=self.layer,
inputs=hidden_states,
input_scale=hidden_states_scale,
quanted_input=buffers.get("quanted_gate_up_input", None),
sublayer_name="w13",
)
@@ -650,22 +753,19 @@ class HummingRunnerCore(MoeRunnerCore):
sublayer_name="w13",
)
self.apply_activation(
inputs=buffers["gate_up_output"],
outputs=buffers["activation_output"],
gate_up = buffers["gate_up_output"].view(
expert_num_tokens.shape[0],
-1,
buffers["gate_up_output"].shape[-1],
)
inputs, input_scale = HummingMethod.may_quant_input(
layer=self.layer,
inputs=buffers["activation_output"],
quanted_input=buffers.get("quanted_down_input", None),
sublayer_name="w2",
down_input, down_scale = self._grouped_masked_act_quant(
gate_up, expert_num_tokens, buffers
)
HummingMethod.forward_layer(
layer=self.layer,
inputs=inputs,
input_scale=input_scale,
inputs=down_input,
input_scale=down_scale,
outputs=buffers["down_output"].view(-1, hidden_states.size(-1)),
valid_shape_m=valid_shape_m,
expert_layout=expert_num_tokens,
@@ -703,6 +803,41 @@ def fused_experts_none_to_humming(
return StandardCombineInput(hidden_states=runner_output.hidden_states)
def _validate_deepep_dispatch_input(
hidden_states: torch.Tensor,
hidden_states_scale: torch.Tensor | None,
layer: torch.nn.Module,
) -> None:
expects_fp8 = layer._humming_uses_deepep_fp8_dispatch
is_fp8 = hidden_states.dtype == torch.float8_e4m3fn
if not expects_fp8:
if is_fp8 or hidden_states_scale is not None:
raise ValueError(
"DeepEP returned FP8 input while Humming is configured for BF16 "
"dispatch."
)
return
if not is_fp8 or hidden_states_scale is None:
raise ValueError(
"Humming expected DeepEP FP8 hidden states and group-128 scales."
)
expected_groups, remainder = divmod(hidden_states.size(-1), 128)
if (
remainder != 0
or hidden_states_scale.dtype != torch.float32
or hidden_states_scale.size(-1) != expected_groups
or hidden_states_scale.numel() != hidden_states.numel() // 128
):
raise ValueError(
"Humming requires row-major FP32 DeepEP scales with group size 128."
)
meta = layer.humming_metas["w13"]
if meta.a_dtype != dtypes.float8e4m3 or meta.input_scale_group_size != 128:
raise ValueError("Humming w13 must use FP8 group-128 input metadata.")
@register_pre_permute("deepep_ll", "humming")
def pre_permute_deepep_ll_to_humming(
dispatch_output: DeepEPLLDispatchOutput,
@@ -715,6 +850,9 @@ def pre_permute_deepep_ll_to_humming(
topk_weights = dispatch_output.topk_weights
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
_validate_deepep_dispatch_input(
hidden_states, dispatch_output.hidden_states_scale, quant_info.layer
)
return HummingRunnerInput(
hidden_states=hidden_states,
@@ -723,6 +861,7 @@ def pre_permute_deepep_ll_to_humming(
expert_num_tokens=dispatch_output.masked_m,
expected_m=dispatch_output.expected_m,
gemm_type=HummingGemmType.GROUPED_MASKED,
hidden_states_scale=dispatch_output.hidden_states_scale,
apply_routed_scaling_factor=False,
)
@@ -759,12 +898,16 @@ def pre_permute_deepep_normal_to_humming(
topk_weights = dispatch_output.topk_weights
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
_validate_deepep_dispatch_input(
hidden_states, dispatch_output.hidden_states_scale, quant_info.layer
)
return HummingRunnerInput(
hidden_states=hidden_states,
topk_weights=topk_weights,
topk_ids=topk_ids.int(),
gemm_type=get_standard_humming_moe_gemm_type(),
hidden_states_scale=dispatch_output.hidden_states_scale,
apply_routed_scaling_factor=False,
)
@@ -790,6 +790,12 @@ class HummingMoEMethod(FusedMoEMethodBase):
if getattr(self, "processed", False):
return
self.processed = True
from sglang.srt.layers.quantization.humming_utils import (
configure_humming_deepep_dispatch,
make_humming_deepep_input_schema,
)
use_deepep_fp8_dispatch = configure_humming_deepep_dispatch(layer)
self.weight_schemas = {}
self.input_schemas = {}
for sublayer_name, configs in layer.sublayer_configs.items():
@@ -869,6 +875,12 @@ class HummingMoEMethod(FusedMoEMethodBase):
del tensors
if use_deepep_fp8_dispatch:
input_schema = make_humming_deepep_input_schema(
sublayer_name, configs["shape_k"]
)
self.input_schemas[sublayer_name] = input_schema
# prepare layer config from humming kernel
HummingMethod.prepare_layer_meta(
layer=layer,
@@ -887,9 +899,6 @@ class HummingMoEMethod(FusedMoEMethodBase):
# preprocess weight for inference
HummingMethod.transform_humming_layer(layer, sublayer_name=sublayer_name)
if hasattr(layer, "dispatcher"):
layer.dispatcher.set_quant_config({"dispatcher_output_dtype": "bf16"})
def create_moe_runner(
self,
layer: torch.nn.Module,
@@ -7,7 +7,9 @@ from humming.schema import BaseWeightSchema
from sglang.srt.environ import envs
from sglang.srt.layers.linear import LinearBase
from sglang.srt.layers.moe import get_moe_a2a_backend
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.runtime_context import get_exec
def humming_is_layer_skipped(config: dict[str, Any], prefix: str):
@@ -35,6 +37,50 @@ def humming_is_layer_skipped(config: dict[str, Any], prefix: str):
return False
def _set_humming_dispatcher_output_dtype(
layer: torch.nn.Module, output_dtype: str
) -> None:
dispatcher = getattr(layer, "dispatcher", None)
if dispatcher is None:
return
quant_config = dict(getattr(dispatcher, "quant_config", None) or {})
quant_config["dispatcher_output_dtype"] = output_dtype
dispatcher.set_quant_config(quant_config)
def configure_humming_deepep_dispatch(layer: torch.nn.Module) -> bool:
if not get_moe_a2a_backend().is_deepep():
layer._humming_uses_deepep_fp8_dispatch = False
_set_humming_dispatcher_output_dtype(layer, "bf16")
return False
output_dtype = get_exec().moe.deepep_dispatcher_output_dtype
if output_dtype == "auto":
output_dtype = "bf16" if envs.SGLANG_DEEPEP_BF16_DISPATCH.get() else "fp8"
if output_dtype not in ("bf16", "fp8"):
raise ValueError(
f"Humming does not support DeepEP {output_dtype} dispatch; "
"use --deepep-dispatcher-output-dtype=bf16 or fp8."
)
_set_humming_dispatcher_output_dtype(layer, output_dtype)
use_fp8 = output_dtype == "fp8"
layer._humming_uses_deepep_fp8_dispatch = use_fp8
return use_fp8
def make_humming_deepep_input_schema(
sublayer_name: str, shape_k: int
) -> HummingInputSchema:
if shape_k % 128 != 0:
raise ValueError(
f"Humming FP8 dispatch requires {sublayer_name} K={shape_k} "
"to be divisible by 128."
)
return HummingInputSchema(a_dtype="float8e4m3", input_scale_group_size=128)
def prepare_humming_layer(layer: LinearBase, quant_config: dict):
weight_schema = BaseWeightSchema.from_config(quant_config)
input_schema = HummingInputSchema()
@@ -84,6 +130,8 @@ def prepare_humming_moe_layer(layer: FusedMoE, quant_config: dict):
# TODO: read input_quant_config from quant_config
input_schema = HummingInputSchema.from_config(input_quant_config)
use_deepep_fp8_dispatch = configure_humming_deepep_dispatch(layer)
shape_config = {
"w13": (
layer.intermediate_size_per_partition * 2,
@@ -121,7 +169,10 @@ def prepare_humming_moe_layer(layer: FusedMoE, quant_config: dict):
)
layer.weight_schemas[sublayer_name] = weight_schema_new
layer.input_schemas[sublayer_name] = input_schema
sub_input_schema = input_schema
if use_deepep_fp8_dispatch:
sub_input_schema = make_humming_deepep_input_schema(sublayer_name, shape_k)
layer.input_schemas[sublayer_name] = sub_input_schema
for name, _ in list(layer.named_parameters()):
if not name.startswith(sublayer_name + "_"):
@@ -140,7 +191,7 @@ def prepare_humming_moe_layer(layer: FusedMoE, quant_config: dict):
shape_k=shape_k,
pad_n_to_multiple=256,
pad_k_to_multiple=128,
input_schema=input_schema,
input_schema=sub_input_schema,
weight_schema=weight_schema_new,
has_bias=layer.with_bias,
num_experts=layer.num_local_experts,
@@ -0,0 +1,265 @@
"""SM90 component CI for Humming FP8 DeepEP dispatch and fused SiLU quant."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.kernels.ops.attention.dsv4.moe import (
silu_and_mul_masked_post_quant,
)
from sglang.kernels.ops.moe.ep_moe_kernels import moe_permute_with_scale
from sglang.srt.layers.moe.moe_runner import humming as humming_runner
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
from sglang.srt.layers.quantization import humming_utils
from sglang.srt.utils import get_device_sm
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=90, stage="base-b-kernel-unit", runner_config="1-gpu-large")
def _silu_reference(gate_up: torch.Tensor, swiglu_limit: float | None) -> torch.Tensor:
gate, up = gate_up.chunk(2, dim=-1)
if swiglu_limit is not None:
gate = gate.clamp_max(swiglu_limit)
up = up.clamp(-swiglu_limit, swiglu_limit)
gate = gate.float()
up = up.float()
return gate * torch.sigmoid(gate) * up
class TestHummingFp8Dispatch(CustomTestCase):
@classmethod
def setUpClass(cls):
if not torch.cuda.is_available():
raise unittest.SkipTest("CUDA is not available")
if get_device_sm() != 90:
raise unittest.SkipTest("Humming FP8 dispatch requires SM90")
def test_fp8_dispatch_preserves_group_scales(self):
generator = torch.Generator(device="cuda").manual_seed(0)
hidden_states = torch.randn(
5,
256,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
).to(torch.float8_e4m3fn)
hidden_states_scale = torch.arange(10, device="cuda", dtype=torch.float32).view(
5, 2
)
topk_ids = torch.tensor(
[[2, 0], [1, -1], [0, 2], [1, 0], [-1, 2]],
device="cuda",
dtype=torch.int32,
)
output, output_scale, src2dst, expert_offsets = moe_permute_with_scale(
inputs=hidden_states,
input_scale=hidden_states_scale,
topk_ids=topk_ids,
num_experts=3,
is_ep=True,
)
for token in range(topk_ids.size(0)):
for slot in range(topk_ids.size(1)):
source_index = token * topk_ids.size(1) + slot
destination_index = int(src2dst[source_index])
if topk_ids[token, slot] < 0:
self.assertLess(destination_index, 0)
continue
self.assertTrue(
torch.equal(
output[destination_index].view(torch.uint8),
hidden_states[token].view(torch.uint8),
)
)
self.assertTrue(
torch.equal(
output_scale[destination_index], hidden_states_scale[token]
)
)
self.assertTrue(
torch.equal(
expert_offsets,
torch.tensor([0, 3, 5, 8], device="cuda", dtype=torch.int32),
)
)
dispatcher = SimpleNamespace(
quant_config={"existing_key": "preserved"},
set_quant_config=MagicMock(),
)
layer = SimpleNamespace(
dispatcher=dispatcher,
humming_metas={
"w13": SimpleNamespace(
a_dtype=humming_runner.dtypes.float8e4m3,
input_scale_group_size=128,
)
},
)
with (
patch.object(
humming_utils,
"get_moe_a2a_backend",
return_value=SimpleNamespace(is_deepep=lambda: True),
),
patch.object(
humming_utils,
"get_exec",
return_value=SimpleNamespace(
moe=SimpleNamespace(deepep_dispatcher_output_dtype="auto")
),
),
patch.object(
humming_utils.envs.SGLANG_DEEPEP_BF16_DISPATCH,
"get",
return_value=False,
),
):
self.assertTrue(humming_utils.configure_humming_deepep_dispatch(layer))
self.assertTrue(layer._humming_uses_deepep_fp8_dispatch)
dispatcher.set_quant_config.assert_called_once_with(
{"existing_key": "preserved", "dispatcher_output_dtype": "fp8"}
)
humming_runner._validate_deepep_dispatch_input(output, output_scale, layer)
def test_silu_fused_masked_quant(self):
num_experts = 8
max_tokens = 17
intermediate = 768
num_groups = intermediate // 128
generator = torch.Generator(device="cuda").manual_seed(1)
gate_up = (
torch.randn(
num_experts,
max_tokens,
2 * intermediate,
generator=generator,
device="cuda",
)
* 3.0
).to(torch.bfloat16)
expert_num_tokens = torch.tensor(
[17, 0, 7, 10, 17, 1, 8, 5], device="cuda", dtype=torch.int32
)
for swiglu_limit in (None, 10.0):
with self.subTest(swiglu_limit=swiglu_limit):
runner = humming_runner.HummingRunnerCore(
MoeRunnerConfig(
num_experts=num_experts,
num_local_experts=num_experts,
intermediate_size_per_partition=intermediate,
top_k=8,
activation="silu",
is_gated=True,
swiglu_limit=swiglu_limit,
gate_up_interleaved=False,
)
)
runner.layer = SimpleNamespace(
humming_metas={
"w2": SimpleNamespace(
a_dtype=humming_runner.dtypes.float8e4m3,
input_scale_group_size=128,
)
}
)
activation_output = torch.full(
(num_experts * max_tokens, intermediate),
float("nan"),
device="cuda",
dtype=torch.bfloat16,
)
buffers = {
"gate_up_output": gate_up.flatten(0, 1),
"activation_output": activation_output,
}
def run_poisoned_kernel(*args, **kwargs):
args[1].view(torch.uint8).fill_(0x7F)
args[2].fill_(float("nan"))
return silu_and_mul_masked_post_quant(*args, **kwargs)
with (
patch.object(
runner,
"apply_activation",
side_effect=AssertionError("SiLU fused path was not selected"),
) as fallback_activation,
patch(
"sglang.kernels.ops.attention.dsv4.moe."
"silu_and_mul_masked_post_quant",
side_effect=run_poisoned_kernel,
) as fused_silu,
):
down_input, down_scale = runner._grouped_masked_act_quant(
gate_up, expert_num_tokens, buffers
)
fused_silu.assert_called_once()
fallback_activation.assert_not_called()
call = fused_silu.call_args
self.assertIs(call.args[0], gate_up)
self.assertIs(call.args[4], expert_num_tokens)
self.assertEqual(call.args[3], 128)
self.assertEqual(call.kwargs["topk"], 8)
self.assertEqual(call.kwargs["swiglu_limit"], swiglu_limit)
self.assertFalse(call.kwargs["scale_ue8m0"])
self.assertEqual(call.args[1].data_ptr(), down_input.data_ptr())
self.assertEqual(call.args[2].data_ptr(), down_scale.data_ptr())
self.assertTrue(bool(activation_output.isnan().all()))
actual = down_input.view(num_experts, max_tokens, intermediate)
actual_scale = down_scale.view(num_experts, max_tokens, num_groups)
reference = _silu_reference(gate_up, swiglu_limit)
grouped_reference = reference.view(
num_experts, max_tokens, num_groups, 128
)
expected_scale = (
grouped_reference.abs().amax(dim=-1).clamp_min(1e-10) / 448.0
)
for expert in range(num_experts):
valid = int(expert_num_tokens[expert])
if valid:
torch.testing.assert_close(
actual_scale[expert, :valid],
expected_scale[expert, :valid],
rtol=5e-3,
atol=1e-6,
)
expanded_scale = actual_scale[expert, :valid].repeat_interleave(
128, dim=-1
)
dequantized = actual[expert, :valid].float() * expanded_scale
self.assertTrue(
bool(
(
(dequantized - reference[expert, :valid]).abs()
<= expanded_scale * 17.0 + 1e-4
).all()
)
)
if valid < max_tokens:
self.assertTrue(
bool(
(actual[expert, valid:].view(torch.uint8) == 0x7F).all()
)
)
self.assertTrue(
bool(actual_scale[expert, valid:].isnan().all())
)
if __name__ == "__main__":
unittest.main()