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,
|
||||
|
||||
Reference in New Issue
Block a user