Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
55 commits
Select commit Hold shift + click to select a range
6d08663
MiniMax-M3: MXFP8 dense-only block convert + torch._scaled_mm 1x32 path
zcnrex Aug 25, 2026
be4ed6d
MXFP8: platform-gate the dense block-convert env and the bf16 backend
kevin-mii Aug 31, 2026
950005e
Merge remote-tracking branch 'upstream/main' into cl/36574
Sep 10, 2026
fc52c84
Merge branch 'main' into m3-mxfp8-dense-block-convert
kevin-mii Sep 16, 2026
d4a44f8
Merge remote-tracking branch 'origin/main' into ci/pr-36574
kevin-mii Sep 17, 2026
c169701
Merge commit '126c2f1bc13fc072aa1345a13ac79a62e8d8bfde' into HEAD
kevin-mii Sep 17, 2026
f492380
Reject missing alpha for AITER MXFP8 OAI activation
kevin-mii Sep 17, 2026
ea596e9
Merge upstream main into m3-mxfp8-dense-block-convert to retrigger CI
kevin-mii Sep 18, 2026
af21ed6
Merge upstream main into m3-mxfp8-dense-block-convert to retrigger CI
kevin-mii Sep 18, 2026
990c77d
Trigger CI after transient runner failure
kevin-mii Sep 18, 2026
2481589
Merge upstream main into m3-mxfp8-dense-block-convert to retrigger CI
kevin-mii Sep 18, 2026
19c7e13
Merge upstream main into m3-mxfp8-dense-block-convert to retrigger fa…
kevin-mii Sep 18, 2026
0410fa2
Trigger CI after shared MLX failure
kevin-mii Sep 18, 2026
3dea28a
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
5a0c8e2
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
98d0030
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
85d0cca
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
f7385f5
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
bcb5e6d
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
3d85429
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
1deab8a
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
45e9bb8
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
d361d49
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
4fa0ffd
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
1ffd82b
Merge upstream main to retrigger failed CI
kevin-mii Sep 18, 2026
94dd4cc
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
05ef582
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
873a2af
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
e3a1ea5
Trigger CI after failed runner checks
kevin-mii Sep 18, 2026
0b06297
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
5fbf10f
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
7782ade
Merge upstream main to retrigger failed CI
kevin-mii Sep 19, 2026
08aa2d8
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
37cdad9
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
babb4d0
Merge upstream main to retrigger failed CI
kevin-mii Sep 19, 2026
d361beb
Merge upstream main to retrigger failed CI
kevin-mii Sep 19, 2026
7012d6a
Merge upstream main to retrigger failed CI
kevin-mii Sep 19, 2026
cb8f3d6
Merge upstream main to retrigger failed CI
kevin-mii Sep 19, 2026
62eefb1
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
f14d25a
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
2a0c916
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
7ccc581
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
419d1ba
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
9a75a55
Trigger CI after failed runner checks
kevin-mii Sep 19, 2026
567d8f6
Merge branch 'main' into m3-mxfp8-dense-block-convert
yctseng0211 Sep 21, 2026
fc19612
Merge remote-tracking branch 'origin/main' into claude/pr-36574-revie…
kevin-mii Sep 23, 2026
809f6b0
mxfp8: drop the --fp8-gemm-backend bf16 choice: unmeasured A/B scaffo…
kevin-mii Sep 23, 2026
feabac5
mxfp8: drop the torch._scaled_mm Blockwise-1x32 route: unmeasured def…
kevin-mii Sep 23, 2026
c60e1e5
Revert the behavior-neutral dispatch and minimax_m3 refactors
kevin-mii Sep 23, 2026
f9d9ff3
mxfp8: drop the _fp8_qinput consumer: its producer is not in this PR
kevin-mii Sep 23, 2026
8325593
mxfp8: build the rowwise-fp8 copy only when the aiter block runner co…
kevin-mii Sep 23, 2026
5dab498
Trim the SGLANG_FORCE_MXFP8_BLOCK_CONVERT_DENSE comment to the constr…
kevin-mii Sep 23, 2026
e1f9d41
mxfp8: state where the ptpc M cutoff comes from; restore an unrelated…
kevin-mii Sep 23, 2026
c369443
Merge remote-tracking branch 'origin/main' into claude/pr-36574-revie…
kevin-mii Sep 23, 2026
067e05f
mxfp8: allow --moe-runner-backend aiter for MXFP8 MoE on gfx950
kevin-mii Sep 23, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions python/sglang/srt/arg_groups/overrides.py
Original file line number Diff line number Diff line change
Expand Up @@ -1428,6 +1428,9 @@ def _moe_runner_backend_quant_constraints(view: Any) -> dict:
allowed = list(MXFP8_MOE_RUNNER_BACKEND_CHOICES)
if is_gfx95_mxfp8:
allowed.append("triton")
# the aiter MXFP8 MoE quant info is built only when aiter is enabled
if envs.SGLANG_USE_AITER.get():
allowed.append("aiter")

if view.moe_a2a_backend == "flashinfer_megamoe":
mxfp8_default = "flashinfer_megamoe"
Expand Down
3 changes: 3 additions & 0 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -1000,6 +1000,9 @@ class Envs:
# on load. Unrelated to the NVFP4 block-FP8 NextN path above.
SGLANG_GLM_NEXTN_MOE_PTPC = EnvBool(False)
SGLANG_QUANT_ALLOW_DOWNCASTING = EnvBool(False)
# HIP: convert only MXFP8 dense linears to block-fp8; fused MoE stays MX 1x32,
# since the block-scale MoE kernel lacks SwiGLU-OAI
SGLANG_FORCE_MXFP8_BLOCK_CONVERT_DENSE = EnvBool(False)
SGLANG_FP8_IGNORED_LAYERS = EnvStr("")
SGLANG_FP4_IGNORED_LAYERS = EnvStr("")
# On by default; set SGLANG_ENABLE_FP8_GEMM_CONFIG_TUNE=0 as a kill switch.
Expand Down
81 changes: 79 additions & 2 deletions python/sglang/srt/layers/quantization/fp8.py
Original file line number Diff line number Diff line change
Expand Up @@ -510,7 +510,10 @@ def __init__(self, quant_config: Union[Fp8Config, W4AFp8Config]):
self.block_quant = (
self.use_mxfp8 or self.quant_config.weight_block_size is not None
)
self.convert_mxfp8_to_block = self.use_mxfp8 and _mxfp8_to_block_fp8_required
self.convert_mxfp8_to_block = self.use_mxfp8 and (
_mxfp8_to_block_fp8_required
or (_is_hip and envs.SGLANG_FORCE_MXFP8_BLOCK_CONVERT_DENSE.get())
)
self.weight_block_size = self.quant_config.weight_block_size
self.w8a8_block_fp8_linear = None
self.w8a8_mxfp8_linear = None
Expand Down Expand Up @@ -722,12 +725,34 @@ def process_weights_after_loading_block_quant(self, layer: Module) -> None:
if self.convert_mxfp8_to_block:
from sglang.srt.layers.quantization.mxfp8_block_convert import (
convert_mxfp8_weight_to_block_fp8,
dequant_mxfp8_2d_to_bf16,
)

mx_weight, mx_scale = layer.weight.data, layer.weight_scale_inv.data
qweight, scale = convert_mxfp8_weight_to_block_fp8(
layer.weight.data, layer.weight_scale_inv.data, block=128
mx_weight, mx_scale, block=128
)
layer.weight = Parameter(qweight, requires_grad=False)
if (
_use_aiter
and _is_gfx95_supported
and self.w8a8_block_fp8_linear is aiter_w8a8_block_fp8_linear
):
# rowwise-fp8 copy for the small-M path of aiter_w8a8_block_fp8_linear;
# the later bpreshuffle is an in-place copy_, so these attrs survive
weight_fp32 = dequant_mxfp8_2d_to_bf16(mx_weight, mx_scale).float()
fp8_max = torch.finfo(torch.float8_e4m3fn).max
row_scale = (
weight_fp32.abs().amax(dim=1, keepdim=True).clamp(min=1e-12)
/ fp8_max
)
layer.weight._ptpc_weight = shuffle_weight(
(weight_fp32 / row_scale)
.clamp(-fp8_max, fp8_max)
.to(torch.float8_e4m3fn),
(16, 16),
)
layer.weight._ptpc_scale = row_scale
layer.weight_scale_inv = Parameter(scale, requires_grad=False)
self.use_mxfp8 = False
self.convert_mxfp8_to_block = False
Expand Down Expand Up @@ -2347,6 +2372,30 @@ def _copy_or_rebind(param: Parameter, new_value: torch.Tensor) -> None:

align_mxfp8_moe_weights_for_flashinfer_trtllm(layer)

if _is_hip and _is_gfx95_supported and get_moe_runner_backend().is_aiter():
from aiter.ops.shuffle import shuffle_scale_a16w4, shuffle_weight_a16w4
from aiter.utility import fp4_utils

num_experts = layer.w13_weight.shape[0]
layer.w13_weight.data = shuffle_weight_a16w4(
layer.w13_weight.data.contiguous(), 16, True
)
w13_s3d = layer.w13_weight_scale_inv.data
layer.w13_weight_scale_inv.data = shuffle_scale_a16w4(
w13_s3d.reshape(-1, w13_s3d.shape[-1]).contiguous(),
num_experts,
True,
)
layer.w2_weight.data = shuffle_weight_a16w4(
layer.w2_weight.data.contiguous(), 16, False
)
w2_s3d = layer.w2_weight_scale_inv.data
layer.w2_weight_scale_inv.data = fp4_utils.e8m0_shuffle(
w2_s3d.reshape(-1, w2_s3d.shape[-1]).contiguous()
)
layer.w13_weight.is_shuffled = True
layer.w2_weight.is_shuffled = True

def process_weights_after_loading(self, layer: Module) -> None:
if _is_hip and _use_hip_int4:
self.process_weights_hip_int4(layer)
Expand Down Expand Up @@ -3142,6 +3191,34 @@ def maybe_get_hip_aiter_quant_info(
w13_weight = layer.w13_weight
w2_weight = layer.w2_weight

if self.use_mxfp8:
gemm1_alpha = self.moe_runner_config.gemm1_alpha
if gemm1_alpha != 1.702:
raise NotImplementedError(
f"AITER MXFP8 MoE only supports swiglu-oai "
f"alpha=1.702, got {gemm1_alpha=}."
)
Comment thread
kevin-mii marked this conversation as resolved.
from aiter import ActivationType
from aiter.ops.flydsl.moe_common import GateMode

return AiterMoeQuantInfo(
w13_weight=w13_weight,
w2_weight=w2_weight,
quant_type=AiterQuantType.PER_1X32,
w13_scale=layer.w13_weight_scale_inv,
w2_scale=layer.w2_weight_scale_inv,
expert_mask=layer.dispatcher.expert_mask_gpu if _use_aiter else None,
swiglu_limit=self.moe_runner_config.swiglu_limit
or self.moe_runner_config.gemm1_clamp_limit
or 0.0,
hidden_pad=getattr(layer, "hidden_pad", 0),
intermediate_pad=getattr(layer, "intermediate_pad", 0),
fused_moe_kwargs={
"activation": ActivationType.Swiglu,
"gate_mode": GateMode.INTERLEAVE.value,
},
)

if self.block_quant:
quant_type = (
AiterQuantType.PER_1X32
Expand Down
19 changes: 19 additions & 0 deletions python/sglang/srt/layers/quantization/fp8_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,9 @@
get_bool_env_var("SGLANG_USE_AITER") and _is_hip and not _is_gfx1250_supported
)
_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported
# Conservative, not a tuned crossover: ptpc beat block-fp8 up to M=512 on
# MiniMax-M3 TP4 gfx950 shapes, but lost past M=64 on narrow ones (K=768).
MXFP8_DENSE_PTPC_DECODE_MAX_M = 128
# ROCm 7.0 hipcc miscompiles gemm_a8w8_blockscale_bpreshuffle on gfx95 (#23319).
_use_aiter_bpreshuffle_gfx95 = _use_aiter_gfx95 and get_hip_version() >= (7, 2, 0)
# gfx95 + ROCm < 7.2: bpreshuffle CK is disabled (above), and the non-bpreshuffle
Expand Down Expand Up @@ -1334,6 +1337,22 @@ def aiter_w8a8_block_fp8_linear(
input_2d = input.view(-1, input.shape[-1])
output_shape = [*input.shape[:-1], weight.shape[0]]

# dense linears converted from MXFP8 carry a rowwise-fp8 copy; its ptpc GEMM
# beats the block-fp8 GEMMs at decode-sized M
if input_scale is None:
ptpc_weight = getattr(weight, "_ptpc_weight", None)
if (
ptpc_weight is not None
and input_2d.shape[0] <= MXFP8_DENSE_PTPC_DECODE_MAX_M
):
out = apply_fp8_ptpc_linear(
input=input_2d,
weight=ptpc_weight,
weight_scale=weight._ptpc_scale,
bias=bias,
)
return out.to(input.dtype).view(*output_shape)

n, k = weight.shape

if _use_aiter_bpreshuffle_gfx95:
Expand Down
18 changes: 18 additions & 0 deletions test/registered/unit/test_model_overrides.py
Original file line number Diff line number Diff line change
Expand Up @@ -2661,6 +2661,24 @@ def _view(**kw):
_moe_runner_backend_quant_constraints(_view(quantization="mxfp8")),
{"moe_runner_backend": "flashinfer_trtllm"},
)
# gfx950 accepts --moe-runner-backend aiter for MXFP8 only with aiter enabled
with (
override_platform(is_hip=True),
patch(
"sglang.srt.arg_groups.overrides.is_gfx95_supported",
return_value=True,
),
):
aiter_view = dict(quantization="mxfp8", moe_runner_backend="aiter")
with envs.SGLANG_USE_AITER.override(True):
self.assertEqual(
_moe_runner_backend_quant_constraints(_view(**aiter_view)), {}
)
with envs.SGLANG_USE_AITER.override(False):
self.assertEqual(
_moe_runner_backend_quant_constraints(_view(**aiter_view)),
{"moe_runner_backend": "triton"},
)
with override_platform(is_sm120=True):
self.assertEqual(
_moe_runner_backend_quant_constraints(
Expand Down
Loading