diff --git a/python/pyproject.toml b/python/pyproject.toml index 3340ecfdc..5bfd23211 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -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", diff --git a/python/sglang/kernels/ops/moe/ep_moe_kernels.py b/python/sglang/kernels/ops/moe/ep_moe_kernels.py index bc8aaba8d..938b82951 100644 --- a/python/sglang/kernels/ops/moe/ep_moe_kernels.py +++ b/python/sglang/kernels/ops/moe/ep_moe_kernels.py @@ -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, diff --git a/python/sglang/srt/layers/moe/moe_runner/humming.py b/python/sglang/srt/layers/moe/moe_runner/humming.py index 845bf4e40..7063565b2 100644 --- a/python/sglang/srt/layers/moe/moe_runner/humming.py +++ b/python/sglang/srt/layers/moe/moe_runner/humming.py @@ -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, ) diff --git a/python/sglang/srt/layers/quantization/humming.py b/python/sglang/srt/layers/quantization/humming.py index 4c58edf9a..3f976b1eb 100644 --- a/python/sglang/srt/layers/quantization/humming.py +++ b/python/sglang/srt/layers/quantization/humming.py @@ -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, diff --git a/python/sglang/srt/layers/quantization/humming_utils.py b/python/sglang/srt/layers/quantization/humming_utils.py index 6a5eefd45..753b3a22c 100644 --- a/python/sglang/srt/layers/quantization/humming_utils.py +++ b/python/sglang/srt/layers/quantization/humming_utils.py @@ -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, diff --git a/test/registered/kernels/ops/moe/test_humming_fp8_dispatch.py b/test/registered/kernels/ops/moe/test_humming_fp8_dispatch.py new file mode 100644 index 000000000..7c4d5400b --- /dev/null +++ b/test/registered/kernels/ops/moe/test_humming_fp8_dispatch.py @@ -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()