[GLM5][MoE] perf: Write FlashInfer TRT-LLM MoE output directly (#28416)

This commit is contained in:
Kaixi
2026-07-20 13:45:23 -07:00
committed by GitHub
parent e856eae921
commit 91b210f7b0
2 changed files with 34 additions and 24 deletions
@@ -5,7 +5,7 @@ import torch
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
def _fake_fp8_block_scale_moe( def _fake_fp8_block_scale_moe_out(
routing_logits: torch.Tensor, routing_logits: torch.Tensor,
routing_bias: Optional[torch.Tensor], routing_bias: Optional[torch.Tensor],
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -14,6 +14,7 @@ def _fake_fp8_block_scale_moe(
gemm1_weights_scale: torch.Tensor, gemm1_weights_scale: torch.Tensor,
gemm2_weights: torch.Tensor, gemm2_weights: torch.Tensor,
gemm2_weights_scale: torch.Tensor, gemm2_weights_scale: torch.Tensor,
output: torch.Tensor,
num_experts: int, num_experts: int,
top_k: int, top_k: int,
n_group: Optional[int], n_group: Optional[int],
@@ -29,14 +30,15 @@ def _fake_fp8_block_scale_moe(
tune_max_num_tokens: int = 8192, tune_max_num_tokens: int = 8192,
fp8_quantization_type: Optional[int] = None, fp8_quantization_type: Optional[int] = None,
activation_type: Optional[int] = None, activation_type: Optional[int] = None,
) -> torch.Tensor: ) -> None:
return torch.empty( return None
hidden_states.shape, dtype=torch.bfloat16, device=hidden_states.device
)
@register_custom_op(fake_impl=_fake_fp8_block_scale_moe) @register_custom_op(
def trtllm_fp8_block_scale_moe_wrapper( fake_impl=_fake_fp8_block_scale_moe_out,
mutates_args=["output"],
)
def trtllm_fp8_block_scale_moe_out_wrapper(
routing_logits: torch.Tensor, routing_logits: torch.Tensor,
routing_bias: Optional[torch.Tensor], routing_bias: Optional[torch.Tensor],
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -45,6 +47,7 @@ def trtllm_fp8_block_scale_moe_wrapper(
gemm1_weights_scale: torch.Tensor, gemm1_weights_scale: torch.Tensor,
gemm2_weights: torch.Tensor, gemm2_weights: torch.Tensor,
gemm2_weights_scale: torch.Tensor, gemm2_weights_scale: torch.Tensor,
output: torch.Tensor,
num_experts: int, num_experts: int,
top_k: int, top_k: int,
n_group: Optional[int], n_group: Optional[int],
@@ -60,7 +63,7 @@ def trtllm_fp8_block_scale_moe_wrapper(
tune_max_num_tokens: int = 8192, tune_max_num_tokens: int = 8192,
fp8_quantization_type: Optional[int] = None, fp8_quantization_type: Optional[int] = None,
activation_type: Optional[int] = None, activation_type: Optional[int] = None,
) -> torch.Tensor: ) -> None:
try: try:
from flashinfer.fused_moe import trtllm_fp8_block_scale_moe from flashinfer.fused_moe import trtllm_fp8_block_scale_moe
except ImportError as e: except ImportError as e:
@@ -68,6 +71,7 @@ def trtllm_fp8_block_scale_moe_wrapper(
"Can't import trtllm_fp8_block_scale_moe from flashinfer. " "Can't import trtllm_fp8_block_scale_moe from flashinfer. "
"Please check flashinfer version." "Please check flashinfer version."
) from e ) from e
kwargs = { kwargs = {
"routing_logits": routing_logits, "routing_logits": routing_logits,
"routing_bias": routing_bias, "routing_bias": routing_bias,
@@ -77,6 +81,7 @@ def trtllm_fp8_block_scale_moe_wrapper(
"gemm1_weights_scale": gemm1_weights_scale, "gemm1_weights_scale": gemm1_weights_scale,
"gemm2_weights": gemm2_weights, "gemm2_weights": gemm2_weights,
"gemm2_weights_scale": gemm2_weights_scale, "gemm2_weights_scale": gemm2_weights_scale,
"output": output,
"num_experts": num_experts, "num_experts": num_experts,
"top_k": top_k, "top_k": top_k,
"n_group": n_group, "n_group": n_group,
@@ -101,10 +106,10 @@ def trtllm_fp8_block_scale_moe_wrapper(
kwargs["activation_type"] = ActivationType(activation_type) kwargs["activation_type"] = ActivationType(activation_type)
return trtllm_fp8_block_scale_moe(**kwargs) trtllm_fp8_block_scale_moe(**kwargs)
def _fake_fp8_block_scale_routed_moe( def _fake_fp8_block_scale_routed_moe_out(
topk_ids: torch.Tensor, topk_ids: torch.Tensor,
routing_bias: Optional[torch.Tensor], routing_bias: Optional[torch.Tensor],
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -121,6 +126,7 @@ def _fake_fp8_block_scale_routed_moe(
local_expert_offset: int, local_expert_offset: int,
local_num_experts: int, local_num_experts: int,
routed_scaling_factor: Optional[float], routed_scaling_factor: Optional[float],
output: torch.Tensor,
routing_method_type: int = 0, routing_method_type: int = 0,
use_shuffled_weight: bool = False, use_shuffled_weight: bool = False,
weight_layout: int = 0, weight_layout: int = 0,
@@ -128,14 +134,15 @@ def _fake_fp8_block_scale_routed_moe(
tune_max_num_tokens: int = 8192, tune_max_num_tokens: int = 8192,
fp8_quantization_type: Optional[int] = None, fp8_quantization_type: Optional[int] = None,
activation_type: Optional[int] = None, activation_type: Optional[int] = None,
) -> torch.Tensor: ) -> None:
return torch.empty( return None
hidden_states.shape, dtype=torch.bfloat16, device=hidden_states.device
)
@register_custom_op(fake_impl=_fake_fp8_block_scale_routed_moe) @register_custom_op(
def trtllm_fp8_block_scale_routed_moe_wrapper( fake_impl=_fake_fp8_block_scale_routed_moe_out,
mutates_args=["output"],
)
def trtllm_fp8_block_scale_routed_moe_out_wrapper(
topk_ids: torch.Tensor, topk_ids: torch.Tensor,
routing_bias: Optional[torch.Tensor], routing_bias: Optional[torch.Tensor],
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -152,6 +159,7 @@ def trtllm_fp8_block_scale_routed_moe_wrapper(
local_expert_offset: int, local_expert_offset: int,
local_num_experts: int, local_num_experts: int,
routed_scaling_factor: Optional[float], routed_scaling_factor: Optional[float],
output: torch.Tensor,
routing_method_type: int = 0, routing_method_type: int = 0,
use_shuffled_weight: bool = False, use_shuffled_weight: bool = False,
weight_layout: int = 0, weight_layout: int = 0,
@@ -159,7 +167,7 @@ def trtllm_fp8_block_scale_routed_moe_wrapper(
tune_max_num_tokens: int = 8192, tune_max_num_tokens: int = 8192,
fp8_quantization_type: Optional[int] = None, fp8_quantization_type: Optional[int] = None,
activation_type: Optional[int] = None, activation_type: Optional[int] = None,
) -> torch.Tensor: ) -> None:
try: try:
from flashinfer.fused_moe import trtllm_fp8_block_scale_routed_moe from flashinfer.fused_moe import trtllm_fp8_block_scale_routed_moe
except ImportError as e: except ImportError as e:
@@ -167,6 +175,7 @@ def trtllm_fp8_block_scale_routed_moe_wrapper(
"Can't import trtllm_fp8_block_scale_routed_moe from flashinfer. " "Can't import trtllm_fp8_block_scale_routed_moe from flashinfer. "
"Please check flashinfer version." "Please check flashinfer version."
) from e ) from e
kwargs = { kwargs = {
"topk_ids": topk_ids, "topk_ids": topk_ids,
"routing_bias": routing_bias, "routing_bias": routing_bias,
@@ -176,6 +185,7 @@ def trtllm_fp8_block_scale_routed_moe_wrapper(
"gemm1_weights_scale": gemm1_weights_scale, "gemm1_weights_scale": gemm1_weights_scale,
"gemm2_weights": gemm2_weights, "gemm2_weights": gemm2_weights,
"gemm2_weights_scale": gemm2_weights_scale, "gemm2_weights_scale": gemm2_weights_scale,
"output": output,
"num_experts": num_experts, "num_experts": num_experts,
"top_k": top_k, "top_k": top_k,
"n_group": n_group, "n_group": n_group,
@@ -200,7 +210,7 @@ def trtllm_fp8_block_scale_routed_moe_wrapper(
kwargs["activation_type"] = ActivationType(activation_type) kwargs["activation_type"] = ActivationType(activation_type)
return trtllm_fp8_block_scale_routed_moe(**kwargs) trtllm_fp8_block_scale_routed_moe(**kwargs)
def _fake_fp8_per_tensor_scale_moe( def _fake_fp8_per_tensor_scale_moe(
@@ -25,8 +25,8 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.dp_attention import is_allocation_symmetric from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe.flashinfer_trtllm_moe import ( from sglang.srt.layers.moe.flashinfer_trtllm_moe import (
trtllm_fp8_block_scale_moe_wrapper, trtllm_fp8_block_scale_moe_out_wrapper,
trtllm_fp8_block_scale_routed_moe_wrapper, trtllm_fp8_block_scale_routed_moe_out_wrapper,
trtllm_fp8_per_tensor_scale_moe_wrapper, trtllm_fp8_per_tensor_scale_moe_wrapper,
) )
from sglang.srt.layers.moe.moe_runner.base import ( from sglang.srt.layers.moe.moe_runner.base import (
@@ -735,7 +735,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
), "runner_config.top_k is required for flashinfer_trtllm_routed." ), "runner_config.top_k is required for flashinfer_trtllm_routed."
packed_topk_ids = _get_packed_topk_ids_for_flashinfer_routed(topk_output) packed_topk_ids = _get_packed_topk_ids_for_flashinfer_routed(topk_output)
output = trtllm_fp8_block_scale_routed_moe_wrapper( trtllm_fp8_block_scale_routed_moe_out_wrapper(
topk_ids=packed_topk_ids, topk_ids=packed_topk_ids,
routing_bias=None, routing_bias=None,
hidden_states=a_q, hidden_states=a_q,
@@ -762,6 +762,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
else routing_method_type else routing_method_type
), ),
use_shuffled_weight=use_shuffled_weight, use_shuffled_weight=use_shuffled_weight,
output=symm_output,
tune_max_num_tokens=next_power_of_2(a_q.shape[0]), tune_max_num_tokens=next_power_of_2(a_q.shape[0]),
fp8_quantization_type=int(fp8_quantization_type), fp8_quantization_type=int(fp8_quantization_type),
activation_type=quant_info.activation_type, activation_type=quant_info.activation_type,
@@ -769,7 +770,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
else: else:
assert TopKOutputChecker.format_is_bypassed(topk_output) assert TopKOutputChecker.format_is_bypassed(topk_output)
output = trtllm_fp8_block_scale_moe_wrapper( trtllm_fp8_block_scale_moe_out_wrapper(
routing_logits=router_logits, routing_logits=router_logits,
routing_bias=correction_bias, routing_bias=correction_bias,
hidden_states=a_q, hidden_states=a_q,
@@ -778,6 +779,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
gemm1_weights_scale=quant_info.w13_weight_scale_inv, gemm1_weights_scale=quant_info.w13_weight_scale_inv,
gemm2_weights=quant_info.w2_weight, gemm2_weights=quant_info.w2_weight,
gemm2_weights_scale=quant_info.w2_weight_scale_inv, gemm2_weights_scale=quant_info.w2_weight_scale_inv,
output=symm_output,
num_experts=quant_info.global_num_experts, num_experts=quant_info.global_num_experts,
top_k=topk_config.top_k, top_k=topk_config.top_k,
n_group=topk_config.num_expert_group, n_group=topk_config.num_expert_group,
@@ -796,8 +798,6 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
fp8_quantization_type=int(fp8_quantization_type), fp8_quantization_type=int(fp8_quantization_type),
activation_type=quant_info.activation_type, activation_type=quant_info.activation_type,
) )
# TODO: Once https://github.com/flashinfer-ai/flashinfer/issues/2703 is fixed, pass output to moe kernel and remove this copy.
symm_output.copy_(output)
output = symm_output output = symm_output
else: else:
assert TopKOutputChecker.format_is_bypassed(topk_output) assert TopKOutputChecker.format_is_bypassed(topk_output)