feat(humming): FP8 DeepEP dispatch for humming MoE backend (#31429)
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user