Add aiter bias moe support in gpt-oss mxfp4 model (#17735)

Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
kk
2026-01-29 17:50:11 +08:00
committed by GitHub
parent 319f6886fe
commit ef1c512754
2 changed files with 115 additions and 1 deletions

View File

@@ -51,6 +51,7 @@ from sglang.srt.utils import (
round_up,
set_weight_attrs,
)
from sglang.srt.utils.common import get_bool_env_var
from sglang.srt.utils.custom_op import register_custom_op
_is_sm100_supported = is_cuda() and is_sm100_supported()
@@ -75,6 +76,7 @@ if TYPE_CHECKING:
)
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
_is_shuffle_moe_mxfp4 = is_gfx95_supported()
if _is_hip:
@@ -82,7 +84,11 @@ if _is_hip:
try:
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.shuffle import (
shuffle_scale_a16w4,
shuffle_weight,
shuffle_weight_a16w4,
)
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
except ImportError as err:
@@ -292,6 +298,18 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
intermediate_size_per_partition_after_pad = round_up(
intermediate_size_per_partition, 64
)
elif _use_aiter:
intermediate_size_per_partition_after_pad = round_up(
intermediate_size_per_partition, 256
)
hidden_size = round_up(hidden_size, 256)
self.hidden_pad = hidden_size - layer.hidden_size
self.intermediate_pad = (
intermediate_size_per_partition_after_pad
- layer.intermediate_size_per_partition
)
elif has_triton_kernels:
# TODO: this is a hack to make
# intermediate_size_per_partition_after_pad the same as the
@@ -530,6 +548,58 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
requires_grad=False,
)
return
if _use_aiter:
if layer.w13_weight_bias is not None:
layer.w13_weight_bias.data = layer.w13_weight_bias.data.to(
torch.float32
)
if layer.w2_weight_bias is not None:
layer.w2_weight_bias.data = layer.w2_weight_bias.data.to(torch.float32)
e, n, k = layer.w13_weight.shape
layer.w13_weight.view(torch.uint8).copy_(
layer.w13_weight.data.view(torch.uint8)
.view(e, n // 2, 2, k)
.permute(0, 2, 1, 3)
.contiguous()
.view(e, n, k)
)
layer.w13_weight_scale.data = (
layer.w13_weight_scale.data.view(e, n // 2, 2, -1)
.permute(0, 2, 1, 3)
.contiguous()
.view(e, n, -1)
)
layer.w13_weight.data = shuffle_weight_a16w4(layer.w13_weight, 16, True)
shuffled_w13_scale = shuffle_scale_a16w4(
layer.w13_weight_scale.view(-1, layer.w13_weight_scale.shape[-1]),
self.num_experts,
True,
)
layer.w2_weight.data = shuffle_weight_a16w4(layer.w2_weight, 16, False)
shuffled_w2_scale = shuffle_scale_a16w4(
layer.w2_weight_scale.view(-1, layer.w2_weight_scale.shape[-1]),
self.num_experts,
False,
)
layer.w13_weight_bias.data = (
layer.w13_weight_bias.data.view(-1, n // 2, 2)
.permute(0, 2, 1)
.contiguous()
.view(-1, n)
)
layer.w13_weight_scale = torch.nn.Parameter(
shuffled_w13_scale, requires_grad=False
)
layer.w2_weight_scale = torch.nn.Parameter(
shuffled_w2_scale, requires_grad=False
)
return
if self.use_triton_kernels:
@@ -680,6 +750,43 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
output=symm_output,
)[0]
return StandardCombineInput(hidden_states=trtllm_gen_output)
if _use_aiter:
topk_weights, topk_ids, _ = topk_output
if hasattr(torch, "float4_e2m1fn_x2"):
w13_weight = layer.w13_weight.view(torch.float4_e2m1fn_x2)
w2_weight = layer.w2_weight.view(torch.float4_e2m1fn_x2)
else:
w13_weight = layer.w13_weight
w2_weight = layer.w2_weight
origi_hidden_size = self.hidden_size - self.hidden_pad
x = torch.nn.functional.pad(
x,
(0, self.hidden_pad),
mode="constant",
value=0.0,
)
output = fused_moe(
x,
w13_weight,
w2_weight,
topk_weights,
topk_ids,
expert_mask=layer.expert_mask_gpu,
activation=ActivationType.Swiglu,
quant_type=QuantType.per_1x32,
w1_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale,
doweight_stage1=self.moe_runner_config.apply_router_weight_on_input,
hidden_pad=self.hidden_pad,
intermediate_pad=self.intermediate_pad,
bias1=layer.w13_weight_bias,
bias2=layer.w2_weight_bias,
)
return StandardCombineInput(hidden_states=output)
backend = self.runner.runner_backend
if backend.is_triton_kernels():

View File

@@ -1358,6 +1358,13 @@ class ServerArgs:
logger.warning(
"Detected SM100 and MXFP4 quantization format for GPT-OSS model, enabling FlashInfer MXFP4 MOE kernel."
)
elif (
is_hip() and get_bool_env_var("SGLANG_USE_AITER")
) and is_mxfp4_quant_format:
self.moe_runner_backend = "auto"
logger.warning(
"Detected ROCm and MXFP4 quantization format for GPT-OSS model, enabling aiter MXFP4 MOE kernel."
)
elif self.ep_size == 1 and is_triton_kernels_available():
self.moe_runner_backend = "triton_kernel"
logger.warning(