MoE Refactor: Refactor modelopt_quant.py -> flashinfer_trllm.py (#16685)

Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
This commit is contained in:
b8zhong
2026-02-02 20:45:14 -08:00
committed by GitHub
co-authored by Cheng Wan
parent eedd472025
commit 78bf13db44
3 changed files with 341 additions and 190 deletions
@@ -1467,6 +1467,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
global_num_experts=global_num_experts,
local_expert_offset=moe_ep_rank * num_local_experts,
local_num_experts=num_local_experts,
intermediate_size=layer.w2_weight.shape[2],
routing_method_type=int(
getattr(layer, "routing_method_type", RoutingMethodType.DeepSeekV3)
),
@@ -44,7 +44,6 @@ from sglang.srt.layers.quantization.utils import (
convert_to_channelwise,
is_layer_skipped,
per_tensor_dequantize,
prepare_static_weights_for_trtllm_fp4_moe,
requantize_with_max_scale,
swizzle_blockscale,
)
@@ -600,81 +599,12 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase):
# Align FP8 weights to FlashInfer per-tensor kernel layout if enabled
if get_moe_runner_backend().is_flashinfer_trtllm():
from flashinfer import reorder_rows_for_gated_act_gemm, shuffle_matrix_a
# 1) Swap W13 halves: [Up, Gate] -> [Gate, Up] expected by FI
num_experts, two_n, hidden = layer.w13_weight.shape
inter = two_n // 2
w13_swapped = (
layer.w13_weight.reshape(num_experts, 2, inter, hidden)
.flip(dims=[1])
.reshape(num_experts, two_n, hidden)
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
align_fp8_moe_weights_for_flashinfer_trtllm,
)
# 2) Reorder rows for fused gated activation (W13)
w13_interleaved = [
reorder_rows_for_gated_act_gemm(w13_swapped[i])
for i in range(num_experts)
]
w13_interleaved = torch.stack(w13_interleaved).reshape(
num_experts, two_n, hidden
)
# 3) Shuffle weights for transposed MMA output (both W13, W2)
epilogue_tile_m = 128
w13_shuffled = [
shuffle_matrix_a(w13_interleaved[i].view(torch.uint8), epilogue_tile_m)
for i in range(num_experts)
]
w2_shuffled = [
shuffle_matrix_a(layer.w2_weight[i].view(torch.uint8), epilogue_tile_m)
for i in range(num_experts)
]
layer.w13_weight = Parameter(
torch.stack(w13_shuffled).view(torch.float8_e4m3fn),
requires_grad=False,
)
layer.w2_weight = Parameter(
torch.stack(w2_shuffled).view(torch.float8_e4m3fn),
requires_grad=False,
)
# Precompute and register per-expert output scaling factors for FI MoE
if get_moe_runner_backend().is_flashinfer_trtllm():
# Note: w13_input_scale and w2_input_scale are scalar Parameters post-reduction
assert (
hasattr(layer, "w13_input_scale") and layer.w13_input_scale is not None
)
assert hasattr(layer, "w2_input_scale") and layer.w2_input_scale is not None
assert (
hasattr(layer, "w13_weight_scale")
and layer.w13_weight_scale is not None
)
assert (
hasattr(layer, "w2_weight_scale") and layer.w2_weight_scale is not None
)
input_scale = layer.w13_input_scale.to(torch.float32)
activation_scale = layer.w2_input_scale.to(torch.float32)
w13_weight_scale = layer.w13_weight_scale.to(torch.float32)
w2_weight_scale = layer.w2_weight_scale.to(torch.float32)
output1_scales_scalar = (
w13_weight_scale * input_scale * (1.0 / activation_scale)
)
output1_scales_gate_scalar = w13_weight_scale * input_scale
output2_scales_scalar = activation_scale * w2_weight_scale
layer.output1_scales_scalar = Parameter(
output1_scales_scalar, requires_grad=False
)
layer.output1_scales_gate_scalar = Parameter(
output1_scales_gate_scalar, requires_grad=False
)
layer.output2_scales_scalar = Parameter(
output2_scales_scalar, requires_grad=False
)
# ModelOpt FP8 stores weights in [Up, Gate] order, so we need to swap
align_fp8_moe_weights_for_flashinfer_trtllm(layer, swap_w13_halves=True)
elif get_moe_runner_backend().is_flashinfer_cutlass():
assert (
hasattr(layer, "w13_input_scale") and layer.w13_input_scale is not None
@@ -725,28 +655,19 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase):
get_moe_runner_backend().is_flashinfer_trtllm()
and TopKOutputChecker.format_is_bypassed(topk_output)
):
router_logits = topk_output.router_logits
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
FlashInferTrtllmFp8MoeQuantInfo,
fused_experts_none_to_flashinfer_trtllm_fp8,
)
from sglang.srt.layers.moe.utils import RoutingMethodType
topk_config = topk_output.topk_config
# Constraints
# Constraints for ModelOpt FP8 MoE
assert (
self.moe_runner_config.activation == "silu"
), "Only silu is supported for flashinfer fp8 moe"
from flashinfer import RoutingMethodType
from flashinfer.fused_moe import trtllm_fp8_per_tensor_scale_moe
correction_bias = (
None
if topk_config.correction_bias is None
else topk_config.correction_bias
)
# Pre-quantize activations to FP8 per-tensor using provided input scale
x_fp8, _ = scaled_fp8_quant(x, layer.w13_input_scale)
use_routing_scales_on_input = True
routed_scaling_factor = self.moe_runner_config.routed_scaling_factor
# Enforce Llama4 routing for ModelOpt FP8 MoE for now.
# TODO(brayden): support other routing methods
assert topk_config.top_k == 1, "ModelOpt FP8 MoE requires top_k==1"
@@ -756,50 +677,26 @@ class ModelOptFp8MoEMethod(FusedMoEMethodBase):
assert (
not topk_config.topk_group
), "ModelOpt FP8 MoE does not support grouped top-k"
routing_method_type = RoutingMethodType.Llama4
# FlashInfer TRTLLM requires routing_logits (and bias) to be bfloat16
routing_logits_cast = router_logits.to(torch.bfloat16)
routing_bias_cast = (
None if correction_bias is None else correction_bias.to(torch.bfloat16)
quant_info = FlashInferTrtllmFp8MoeQuantInfo(
w13_weight=layer.w13_weight,
w2_weight=layer.w2_weight,
global_num_experts=layer.num_experts,
local_expert_offset=layer.moe_ep_rank * layer.num_local_experts,
local_num_experts=layer.num_local_experts,
intermediate_size=layer.w2_weight.shape[2],
routing_method_type=RoutingMethodType.Llama4,
block_quant=False,
w13_input_scale=layer.w13_input_scale,
output1_scales_scalar=layer.output1_scales_scalar,
output1_scales_gate_scalar=layer.output1_scales_gate_scalar,
output2_scales_scalar=layer.output2_scales_scalar,
use_routing_scales_on_input=True,
)
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
# FIXME: there is a bug in the trtllm_fp8_block_scale_moe.
# It ignored the `output`` argument. https://github.com/flashinfer-ai/flashinfer/blob/da01b1bd8f9f22aec8c0eea189ad54860b034947/flashinfer/fused_moe/core.py#L1323-L1325
# so we put the whole function under the ``use_symmetric_memory`` context manager.
# If the bug is fixed, we can only put the output tensor allocation under the context manager.
output = trtllm_fp8_per_tensor_scale_moe(
routing_logits=routing_logits_cast,
routing_bias=routing_bias_cast,
hidden_states=x_fp8,
gemm1_weights=layer.w13_weight,
output1_scales_scalar=layer.output1_scales_scalar,
output1_scales_gate_scalar=layer.output1_scales_gate_scalar,
gemm2_weights=layer.w2_weight,
output2_scales_scalar=layer.output2_scales_scalar,
num_experts=layer.num_experts,
top_k=topk_config.top_k,
n_group=0,
topk_group=0,
intermediate_size=layer.w2_weight.shape[2],
local_expert_offset=layer.moe_ep_rank * layer.num_local_experts,
local_num_experts=layer.num_local_experts,
routed_scaling_factor=(
routed_scaling_factor
if routed_scaling_factor is not None
else 1.0
),
use_routing_scales_on_input=use_routing_scales_on_input,
routing_method_type=routing_method_type,
tune_max_num_tokens=next_power_of_2(x.shape[0]),
)
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
return StandardCombineInput(hidden_states=output)
return fused_experts_none_to_flashinfer_trtllm_fp8(
dispatch_output, quant_info, self.moe_runner_config
)
if get_moe_runner_backend().is_flashinfer_cutlass():
activation = ACT_STR_TO_TYPE_MAP[self.moe_runner_config.activation]
@@ -1532,49 +1429,12 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
and reorder_rows_for_gated_act_gemm is not None
and shuffle_matrix_sf_a is not None
):
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
align_fp4_moe_weights_for_flashinfer_trtllm,
)
# FlashInfer TRTLLM processing - handles both w13 and w2
(
gemm1_weights_fp4_shuffled,
gemm1_scales_fp4_shuffled,
gemm2_weights_fp4_shuffled,
gemm2_scales_fp4_shuffled,
) = prepare_static_weights_for_trtllm_fp4_moe(
layer.w13_weight,
layer.w2_weight,
layer.w13_weight_scale,
layer.w2_weight_scale,
layer.w2_weight.size(-2), # hidden_size
layer.w13_weight.size(-2) // 2, # intermediate_size
layer.w13_weight.size(0), # num_experts
)
# Set flashinfer parameters
layer.gemm1_weights_fp4_shuffled = Parameter(
gemm1_weights_fp4_shuffled, requires_grad=False
)
layer.gemm2_weights_fp4_shuffled = Parameter(
gemm2_weights_fp4_shuffled, requires_grad=False
)
layer.gemm1_scales_fp4_shuffled = Parameter(
gemm1_scales_fp4_shuffled, requires_grad=False
)
layer.gemm2_scales_fp4_shuffled = Parameter(
gemm2_scales_fp4_shuffled, requires_grad=False
)
# Additional parameter needed for TRT-LLM
layer.g1_scale_c = Parameter(
(layer.w2_input_scale_quant * layer.g1_alphas).to(torch.float32),
requires_grad=False,
)
# Clean up weights that won't be used by TRT-LLM
del (
layer.w2_weight,
layer.w2_weight_scale,
layer.w13_weight,
layer.w13_weight_scale,
)
align_fp4_moe_weights_for_flashinfer_trtllm(layer)
else:
# CUTLASS processing - handle w13 and w2 separately
@@ -1644,6 +1504,10 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
):
self.moe_runner_config = moe_runner_config
if get_moe_runner_backend().is_flashinfer_trtllm():
self.runner = MoeRunner(
MoeRunnerBackend.FLASHINFER_TRTLLM, moe_runner_config
)
def apply(
self,
@@ -1662,12 +1526,36 @@ class ModelOptNvFp4FusedMoEMethod(FusedMoEMethodBase):
), f"{activation=} missing from {ACT_STR_TO_TYPE_MAP.keys()=}"
moe_runner_config = self.moe_runner_config
# Check if this is a FlashInferFP4MoE layer that should handle its own forward
# FlashInfer TRTLLM FP4 path - layer has shuffled weights only when
# backend is flashinfer_trtllm
if hasattr(layer, "gemm1_weights_fp4_shuffled"):
# This layer was processed with flashinfer TRTLLM - delegate to its own forward
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
FlashInferTrtllmFp4MoeQuantInfo,
)
from sglang.srt.layers.moe.utils import RoutingMethodType
return StandardCombineInput(hidden_states=layer.forward(x, topk_output))
# Determine routing method type based on layer configuration
routing_method_type = getattr(
layer, "routing_method_type", RoutingMethodType.Default
)
quant_info = FlashInferTrtllmFp4MoeQuantInfo(
gemm1_weights_fp4_shuffled=layer.gemm1_weights_fp4_shuffled.data,
gemm2_weights_fp4_shuffled=layer.gemm2_weights_fp4_shuffled.data,
gemm1_scales_fp4_shuffled=layer.gemm1_scales_fp4_shuffled.data,
gemm2_scales_fp4_shuffled=layer.gemm2_scales_fp4_shuffled.data,
g1_scale_c=layer.g1_scale_c.data,
g1_alphas=layer.g1_alphas.data,
g2_alphas=layer.g2_alphas.data,
w13_input_scale_quant=layer.w13_input_scale_quant,
global_num_experts=layer.num_experts,
local_expert_offset=layer.moe_ep_rank * layer.num_local_experts,
local_num_experts=layer.num_local_experts,
intermediate_size_per_partition=layer.intermediate_size_per_partition,
routing_method_type=routing_method_type,
)
return self.runner.run(dispatch_output, quant_info)
if self.enable_flashinfer_cutlass_moe:
from sglang.srt.layers.moe.token_dispatcher import DispatchOutputChecker