[AMD] [GLM-5.3-Flash Day 0] Enable FP8 and Quark MXFP4 MoE on gfx950 (#38546)

Co-authored-by: Raiden-Makoto <Raiden-Makoto@users.noreply.github.com>
Co-authored-by: Thomas Wang <thomawan@amd.com>
Co-authored-by: andyluo7 <andy.luo@amd.com>
Co-authored-by: Kevin Mi <mikevin920@yahoo.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Raiden Makoto
2026-09-21 21:07:37 -07:00
committed by GitHub
co-authored by Raiden-Makoto Thomas Wang andyluo7 Kevin Mi Claude Opus 5
parent 15ba54bd5d
commit b44e248682
8 changed files with 756 additions and 15 deletions
@@ -282,10 +282,13 @@ class AiterRunnerCore(MoeRunnerCore):
# `SGLANG_USE_AITER_MOE_GU_ITLV=0` to switch to SEPARATED, which # `SGLANG_USE_AITER_MOE_GU_ITLV=0` to switch to SEPARATED, which
# matches the layout produced by `Mxfp4MoEMethod` (gpt-oss # matches the layout produced by `Mxfp4MoEMethod` (gpt-oss
# MXFP4) and the gptoss_fp4 tuned FlyDSL kernels. # MXFP4) and the gptoss_fp4 tuned FlyDSL kernels.
extra["gate_mode"] = ( extra.setdefault(
GateMode.INTERLEAVE.value "gate_mode",
if envs.SGLANG_USE_AITER_MOE_GU_ITLV.get() (
else GateMode.SEPARATED.value GateMode.INTERLEAVE.value
if envs.SGLANG_USE_AITER_MOE_GU_ITLV.get()
else GateMode.SEPARATED.value
),
) )
extra["swiglu_limit"] = quant_info.swiglu_limit extra["swiglu_limit"] = quant_info.swiglu_limit
if self.config.no_combine: if self.config.no_combine:
@@ -1886,6 +1886,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
) )
layer.w13_weight.is_shuffled = True layer.w13_weight.is_shuffled = True
layer.w2_weight.is_shuffled = True layer.w2_weight.is_shuffled = True
layer._aiter_gate_up_interleaved = False
return return
elif self.use_mxfp8 and get_moe_a2a_backend().is_flashinfer_megamoe(): elif self.use_mxfp8 and get_moe_a2a_backend().is_flashinfer_megamoe():
from sglang.srt.layers.moe.flashinfer_megamoe import ( from sglang.srt.layers.moe.flashinfer_megamoe import (
@@ -1933,6 +1934,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
) )
layer.w13_weight.is_shuffled = True layer.w13_weight.is_shuffled = True
layer.w2_weight.is_shuffled = True layer.w2_weight.is_shuffled = True
layer._aiter_gate_up_interleaved = False
elif _use_aiter: elif _use_aiter:
# Pre-shuffle weights # Pre-shuffle weights
t = shuffle_weight(layer.w13_weight, (16, 16)) t = shuffle_weight(layer.w13_weight, (16, 16))
@@ -1943,6 +1945,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
del t del t
layer.w13_weight.is_shuffled = True layer.w13_weight.is_shuffled = True
layer.w2_weight.is_shuffled = True layer.w2_weight.is_shuffled = True
layer._aiter_gate_up_interleaved = False
elif _is_cpu: elif _is_cpu:
assert _is_cpu_amx_available, ( assert _is_cpu_amx_available, (
"Fp8MoEMethod on CPU requires that CPU has AMX support" "Fp8MoEMethod on CPU requires that CPU has AMX support"
@@ -2650,6 +2653,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
requires_grad=False, requires_grad=False,
) )
torch.cuda.empty_cache() torch.cuda.empty_cache()
layer._aiter_gate_up_interleaved = False
# ROCm (_use_aiter): using column-wise scaling # ROCm (_use_aiter): using column-wise scaling
layer.w13_weight_scale1 *= layer.w13_weight_scale.unsqueeze(-1) layer.w13_weight_scale1 *= layer.w13_weight_scale.unsqueeze(-1)
@@ -3158,6 +3162,23 @@ class Fp8MoEMethod(FusedMoEMethodBase):
quant_type = AiterQuantType.PER_TOKEN quant_type = AiterQuantType.PER_TOKEN
w13_scale = layer.w13_weight_scale1 w13_scale = layer.w13_weight_scale1
w2_scale = layer.w2_weight_scale1 w2_scale = layer.w2_weight_scale1
fused_moe_kwargs = None
gate_up_interleaved = getattr(layer, "_aiter_gate_up_interleaved", None)
if (
gate_up_interleaved is not None
and (self.moe_runner_config.swiglu_limit or 0.0) > 0
):
from aiter.ops.flydsl.moe_common import GateMode
fused_moe_kwargs = {
"gate_mode": (
GateMode.INTERLEAVE.value
if gate_up_interleaved
else GateMode.SEPARATED.value
)
}
return AiterMoeQuantInfo( return AiterMoeQuantInfo(
w13_weight=w13_weight, w13_weight=w13_weight,
w2_weight=w2_weight, w2_weight=w2_weight,
@@ -3168,6 +3189,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
swiglu_limit=self.moe_runner_config.swiglu_limit or 0.0, swiglu_limit=self.moe_runner_config.swiglu_limit or 0.0,
hidden_pad=getattr(layer, "hidden_pad", 0), hidden_pad=getattr(layer, "hidden_pad", 0),
intermediate_pad=getattr(layer, "intermediate_pad", 0), intermediate_pad=getattr(layer, "intermediate_pad", 0),
fused_moe_kwargs=fused_moe_kwargs,
) )
@@ -15,7 +15,11 @@ from sglang.srt.layers.quantization.base_config import ( # noqa: E501
QuantizationConfig, QuantizationConfig,
QuantizeMethodBase, QuantizeMethodBase,
) )
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8LinearMethod from sglang.srt.layers.quantization.fp8 import (
Fp8Config,
Fp8LinearMethod,
Fp8MoEMethod,
)
from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod from sglang.srt.layers.quantization.kv_cache import BaseKVCacheMethod
from sglang.srt.layers.quantization.quark.schemes import ( from sglang.srt.layers.quantization.quark.schemes import (
QuarkLinearScheme, QuarkLinearScheme,
@@ -375,6 +379,46 @@ class QuarkConfig(QuantizationConfig):
expanded.append(name.removeprefix("language_model.")) expanded.append(name.removeprefix("language_model."))
self.exclude_layers = list(dict.fromkeys(expanded)) self.exclude_layers = list(dict.fromkeys(expanded))
layer_quant_config = self.quant_config.get("layer_quant_config")
if layer_quant_config:
self.quant_config["layer_quant_config"] = hf_to_sglang_mapper.apply_dict(
layer_quant_config
)
if self.kv_cache_group:
self.kv_cache_group = hf_to_sglang_mapper.apply_list(self.kv_cache_group)
@staticmethod
def _get_block_fp8_config(
layer_quant_config: Optional[dict[str, Any]],
packed_modules_mapping: dict[str, list[str]],
) -> Optional[Fp8Config]:
if layer_quant_config is None:
return None
weight_config = layer_quant_config.get("weight") or {}
input_config = layer_quant_config.get("input_tensors") or {}
block_size = weight_config.get("block_size")
if not (
not layer_quant_config.get("output_tensors")
and not layer_quant_config.get("bias")
and weight_config.get("dtype") in {"fp8_e4m3", "fp8_e4m3fn"}
and weight_config.get("qscheme") == "per_block"
and weight_config.get("is_dynamic") is False
and isinstance(block_size, list)
and len(block_size) == 2
and input_config.get("dtype") in {"fp8_e4m3", "fp8_e4m3fn"}
and input_config.get("is_dynamic") is True
):
return None
return Fp8Config(
is_checkpoint_fp8_serialized=True,
activation_scheme="dynamic",
weight_block_size=block_size,
packed_modules_mapping=packed_modules_mapping,
)
def get_quant_method( def get_quant_method(
self, layer: torch.nn.Module, prefix: str self, layer: torch.nn.Module, prefix: str
) -> Optional["QuantizeMethodBase"]: ) -> Optional["QuantizeMethodBase"]:
@@ -396,6 +440,17 @@ class QuarkConfig(QuantizationConfig):
return QuarkKVCacheMethod(self) return QuarkKVCacheMethod(self)
return None return None
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
block_fp8_config = self._get_block_fp8_config(
self._find_matched_config(prefix, layer), self.packed_modules_mapping
)
if block_fp8_config is not None:
if isinstance(layer, LinearBase):
return Fp8LinearMethod(block_fp8_config)
if isinstance(layer, FusedMoE):
return Fp8MoEMethod(block_fp8_config)
if isinstance(layer, LinearBase): if isinstance(layer, LinearBase):
scheme = self.get_linear_scheme(layer=layer, layer_name=prefix) scheme = self.get_linear_scheme(layer=layer, layer_name=prefix)
layer.scheme = scheme layer.scheme = scheme
@@ -406,8 +461,6 @@ class QuarkConfig(QuantizationConfig):
self._online_quantized_layers.add(prefix) self._online_quantized_layers.add(prefix)
return QuarkKVCacheMethod(self) return QuarkKVCacheMethod(self)
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
if isinstance(layer, FusedMoE): if isinstance(layer, FusedMoE):
self._online_quantized_layers.add(prefix) self._online_quantized_layers.add(prefix)
layer.scheme = self.get_moe_scheme(layer, prefix) layer.scheme = self.get_moe_scheme(layer, prefix)
@@ -38,12 +38,12 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
_is_shuffle_moe_mxfp4 = is_gfx95_supported()
__all__ = ["QuarkW4A4MXFp4MoE"] __all__ = ["QuarkW4A4MXFp4MoE"]
_is_hip = is_hip() _is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
_is_gfx95 = is_gfx95_supported()
_is_shuffle_moe_mxfp4 = _use_aiter and _is_gfx95
if _use_aiter: if _use_aiter:
from aiter.ops.shuffle import moe_shuffle_scale, moe_shuffle_weight, shuffle_weight from aiter.ops.shuffle import moe_shuffle_scale, moe_shuffle_weight, shuffle_weight
from aiter.utility.fp4_utils import e8m0_shuffle from aiter.utility.fp4_utils import e8m0_shuffle
@@ -202,6 +202,8 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
is_concat=True, is_concat=True,
is_packed=True, is_packed=True,
) )
layer.hidden_pad = 0
layer.intermediate_pad = w13_up_dim // 2 - intermediate_size_per_partition
# Add the quantization method used (per tensor/grouped/channel) # Add the quantization method used (per tensor/grouped/channel)
# to ensure the weight scales are loaded in properly # to ensure the weight scales are loaded in properly
@@ -819,6 +821,11 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
layer.w2_weight = torch.nn.Parameter(qw2_weight, requires_grad=False) layer.w2_weight = torch.nn.Parameter(qw2_weight, requires_grad=False)
def process_weights_after_loading(self, layer: torch.nn.Module) -> None: def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
if not getattr(self, "_owns_moe_runner", False):
raise RuntimeError(
"Quark MXFP4 weight preshuffling requires an owned AITER runner."
)
if ( if (
not self.is_checkpoint_mxfp4_serialized not self.is_checkpoint_mxfp4_serialized
or self.dequantization_config is not None or self.dequantization_config is not None
@@ -889,15 +896,19 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
) )
self.moe_runner_config = moe_runner_config self.moe_runner_config = moe_runner_config
self._owns_moe_runner = False
moe_runner_backend = get_moe_runner_backend() moe_runner_backend = get_moe_runner_backend()
if moe_runner_backend.is_auto() and get_moe_a2a_backend().supports_aiter(): if moe_runner_backend.is_auto() and get_moe_a2a_backend().supports_aiter():
moe_runner_backend = MoeRunnerBackend.AITER moe_runner_backend = MoeRunnerBackend.AITER
if moe_runner_backend.is_aiter(): if moe_runner_backend.is_aiter():
self.runner = MoeRunner(moe_runner_backend, moe_runner_config) self.runner = MoeRunner(moe_runner_backend, moe_runner_config)
self._owns_moe_runner = True
else: else:
# TODO(cwan): refactor other backends raise NotImplementedError(
pass "Quark MXFP4 MoE currently requires the AITER runner; "
f"got {moe_runner_backend.value!r}."
)
def apply_weights( def apply_weights(
self, self,
@@ -924,6 +935,12 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
from aiter.ops.flydsl.moe_common import GateMode from aiter.ops.flydsl.moe_common import GateMode
_fused_moe_kwargs = {"gate_mode": GateMode.INTERLEAVE.value} _fused_moe_kwargs = {"gate_mode": GateMode.INTERLEAVE.value}
elif _is_gfx95:
from aiter.ops.flydsl.moe_common import GateMode
# Quark checkpoints store gate and up projections as separate
# contiguous row ranges. Keep that ordering for correctness.
_fused_moe_kwargs = {"gate_mode": GateMode.SEPARATED.value}
else: else:
_fused_moe_kwargs = None _fused_moe_kwargs = None
@@ -934,6 +951,9 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
w13_scale=layer.w13_weight_scale, w13_scale=layer.w13_weight_scale,
w2_scale=layer.w2_weight_scale, w2_scale=layer.w2_weight_scale,
expert_mask=layer.dispatcher.expert_mask_gpu, expert_mask=layer.dispatcher.expert_mask_gpu,
hidden_pad=getattr(layer, "hidden_pad", 0),
intermediate_pad=getattr(layer, "intermediate_pad", 0),
swiglu_limit=self.moe_runner_config.swiglu_limit or 0.0,
fused_moe_kwargs=_fused_moe_kwargs, fused_moe_kwargs=_fused_moe_kwargs,
) )
return self.runner.run(dispatch_output, quant_info) return self.runner.run(dispatch_output, quant_info)
+17 -3
View File
@@ -1233,11 +1233,13 @@ class Glm5NextModel(nn.Module):
class Glm5NextForConditionalGeneration(nn.Module): class Glm5NextForConditionalGeneration(nn.Module):
hf_to_sglang_mapper = WeightsMapper( hf_to_sglang_mapper = WeightsMapper(
orig_to_new_substr={ orig_to_new_substr={
"model.language_model.": "model.",
"model.visual": "visual", "model.visual": "visual",
} },
orig_to_new_prefix={
"model.language_model.": "model.",
},
orig_to_new_suffix={".attn.qkv": ".attn.qkv_proj"},
) )
packed_modules_mapping = { packed_modules_mapping = {
"fused_qkv_a_proj_with_mqa": ["q_a_proj", "kv_a_proj_with_mqa"], "fused_qkv_a_proj_with_mqa": ["q_a_proj", "kv_a_proj_with_mqa"],
**Glm5NextLinearAttention._PACKED_MODULES_MAPPING, **Glm5NextLinearAttention._PACKED_MODULES_MAPPING,
@@ -1570,6 +1572,14 @@ class Glm5NextForConditionalGeneration(nn.Module):
fused_cat_dim = 0 fused_cat_dim = 0
params_dict = dict(self.named_parameters()) params_dict = dict(self.named_parameters())
def maybe_map_fp8_block_scale_name(name: str) -> str:
if name.endswith("weight_scale"):
candidate = name.removesuffix("weight_scale") + "weight_scale_inv"
if candidate in params_dict:
return candidate
return name
weight_names = [] weight_names = []
for name, loaded_weight in weights: for name, loaded_weight in weights:
is_visual_weight = "visual" in name is_visual_weight = "visual" in name
@@ -1633,6 +1643,7 @@ class Glm5NextForConditionalGeneration(nn.Module):
if "mlp.experts" in name: if "mlp.experts" in name:
continue continue
candidate = name.replace(weight_name, param_name) candidate = name.replace(weight_name, param_name)
candidate = maybe_map_fp8_block_scale_name(candidate)
if ( if (
param_name param_name
in { in {
@@ -1662,6 +1673,7 @@ class Glm5NextForConditionalGeneration(nn.Module):
continue continue
is_expert_weight = True is_expert_weight = True
name = name.replace(weight_name, param_name) name = name.replace(weight_name, param_name)
name = maybe_map_fp8_block_scale_name(name)
if name not in params_dict: if name not in params_dict:
continue continue
param = params_dict[name] param = params_dict[name]
@@ -1713,6 +1725,7 @@ class Glm5NextForConditionalGeneration(nn.Module):
"fused_qkv_a_proj_with_mqa", "fused_qkv_a_proj_with_mqa",
) )
) )
target = maybe_map_fp8_block_scale_name(target)
if target in params_dict: if target in params_dict:
param = params_dict[target] param = params_dict[target]
weight_loader = getattr( weight_loader = getattr(
@@ -1723,6 +1736,7 @@ class Glm5NextForConditionalGeneration(nn.Module):
cached_a_proj.pop(kv_a_proj_name, None) cached_a_proj.pop(kv_a_proj_name, None)
continue continue
name = maybe_map_fp8_block_scale_name(name)
if name not in params_dict: if name not in params_dict:
continue continue
@@ -0,0 +1,397 @@
"""Isolated gfx950 numerical tests for GLM-5.3-Flash Quark MoE."""
import unittest
from types import SimpleNamespace
import torch
import torch.nn.functional as F
from aiter.ops.flydsl.moe_common import GateMode
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from sglang.srt.layers.moe.moe_runner.aiter import (
AiterMoeQuantInfo,
AiterQuantType,
AiterRunnerCore,
AiterRunnerInput,
)
from sglang.srt.layers.quantization.fp8_utils import dequant_mxfp4
from sglang.srt.utils import is_gfx95_supported, is_hip
from sglang.test.ci.ci_register import register_amd_ci
from sglang.test.test_utils import CustomTestCase
register_amd_ci(est_time=180, suite="stage-b-test-1-gpu-small-amd-mi35x")
@unittest.skipUnless(
torch.cuda.is_available() and is_hip() and is_gfx95_supported(),
"requires one gfx950 GPU",
)
class TestGLM53FlashQuarkMoE(CustomTestCase):
hidden_size = 4096
intermediate_size = 2048
num_experts = 9
swiglu_limit = 10.0
@classmethod
def setUpClass(cls):
super().setUpClass()
torch.manual_seed(7)
cls.weights = cls._make_mxfp4_bank()
cls.weights["w13_deq"] = cls._dequant(
cls.weights["w13_raw"], cls.weights["s13_raw"]
)
cls.weights["w2_deq"] = cls._dequant(
cls.weights["w2_raw"], cls.weights["s2_raw"]
)
cls.runner = AiterRunnerCore(
SimpleNamespace(
no_combine=False,
activation="silu",
gemm1_alpha=None,
gemm1_clamp_limit=None,
)
)
@classmethod
def _make_mxfp4_bank(cls):
gate_weights = []
up_weights = []
down_weights = []
gate_scales = []
up_scales = []
down_scales = []
for expert in range(cls.num_experts):
generator = torch.Generator(device="cuda")
generator.manual_seed(100 + expert)
gate = (
torch.randn(
cls.intermediate_size,
cls.hidden_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.05
)
up = (
torch.randn(
cls.intermediate_size,
cls.hidden_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.05
)
down = (
torch.randn(
cls.hidden_size,
cls.intermediate_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.01
)
gate_q, gate_s = dynamic_mxfp4_quant(gate)
up_q, up_s = dynamic_mxfp4_quant(up)
down_q, down_s = dynamic_mxfp4_quant(down)
gate_weights.append(gate_q)
up_weights.append(up_q)
down_weights.append(down_q)
gate_scales.append(gate_s)
up_scales.append(up_s)
down_scales.append(down_s)
w13 = torch.cat([torch.stack(gate_weights), torch.stack(up_weights)], dim=1)
w2 = torch.stack(down_weights)
s13 = torch.cat([torch.stack(gate_scales), torch.stack(up_scales)], dim=1)
s2 = torch.stack(down_scales)
return {
"w13_raw": w13,
"w2_raw": w2,
"s13_raw": s13,
"s2_raw": s2,
"w13": shuffle_weight(w13.contiguous(), (16, 16)),
"w2": shuffle_weight(w2.contiguous(), (16, 16)),
"s13": e8m0_shuffle(s13.view(-1, s13.shape[-1])).view_as(s13),
"s2": e8m0_shuffle(s2.view(-1, s2.shape[-1])).view_as(s2),
}
@staticmethod
def _quantize_fp8_weight(weight):
rows, width = weight.shape
blocks = (
weight.float().view(rows // 128, 128, width // 128, 128).permute(0, 2, 1, 3)
)
scale = blocks.abs().amax(dim=(2, 3)).clamp(min=1e-12) / 448.0
quantized = (blocks / scale[:, :, None, None]).to(torch.float8_e4m3fn)
return (
quantized.permute(0, 2, 1, 3).reshape(rows, width),
scale,
)
@staticmethod
def _dequantize_fp8_weight(weight, scale):
return weight.float() * scale.repeat_interleave(128, dim=0).repeat_interleave(
128, dim=1
)
@staticmethod
def _quant_dequant_fp8_activation(activation):
tokens, width = activation.shape
groups = activation.float().view(tokens, width // 128, 128)
scale = groups.abs().amax(dim=-1).clamp(min=1e-12) / 448.0
quantized = (groups / scale.unsqueeze(-1)).to(torch.float8_e4m3fn)
return (quantized.float() * scale.unsqueeze(-1)).reshape(tokens, width)
@classmethod
def tearDownClass(cls):
if hasattr(cls, "weights"):
del cls.weights
if hasattr(cls, "runner"):
del cls.runner
torch.cuda.empty_cache()
super().tearDownClass()
@classmethod
def _dequant(cls, weight, scale):
experts, rows, packed = weight.shape
blocks = packed // 16
return dequant_mxfp4(
weight.view(experts, rows, blocks, 16),
scale,
torch.bfloat16,
)
@classmethod
def _quant_dequant_activation(cls, activation):
quantized, scale = dynamic_mxfp4_quant(activation)
tokens, packed = quantized.shape
blocks = packed // 16
return dequant_mxfp4(
quantized.view(1, tokens, blocks, 16),
scale.view(1, tokens, blocks),
torch.bfloat16,
).squeeze(0)
@classmethod
def _torch_oracle(cls, hidden_states, topk_ids, topk_weights):
w13 = cls.weights["w13_deq"]
w2 = cls.weights["w2_deq"]
output = torch.zeros_like(hidden_states)
hidden_qdq = cls._quant_dequant_activation(hidden_states)
for token in range(hidden_states.shape[0]):
for route in range(topk_ids.shape[1]):
expert = int(topk_ids[token, route])
gate = F.linear(
hidden_qdq[token].float(),
w13[expert, : cls.intermediate_size].float(),
)
up = F.linear(
hidden_qdq[token].float(),
w13[expert, cls.intermediate_size :].float(),
)
gate = gate.clamp(max=cls.swiglu_limit)
up = up.clamp(min=-cls.swiglu_limit, max=cls.swiglu_limit)
activated = F.silu(gate) * up
activated = cls._quant_dequant_activation(
activated.unsqueeze(0).bfloat16()
).squeeze(0)
expert_output = F.linear(activated.float(), w2[expert].float())
output[token] += (expert_output * topk_weights[token, route]).to(
output.dtype
)
return output
@classmethod
def _aiter(cls, hidden_states, topk_ids, topk_weights):
w13 = cls.weights["w13"].view(torch.float4_e2m1fn_x2)
w2 = cls.weights["w2"].view(torch.float4_e2m1fn_x2)
w13.is_shuffled = True
w2.is_shuffled = True
quant_info = AiterMoeQuantInfo(
w13_weight=w13,
w2_weight=w2,
quant_type=AiterQuantType.PER_1X32,
w13_scale=cls.weights["s13"],
w2_scale=cls.weights["s2"],
swiglu_limit=cls.swiglu_limit,
fused_moe_kwargs={"gate_mode": GateMode.SEPARATED.value},
)
runner_input = AiterRunnerInput(
hidden_states=hidden_states,
topk_ids=topk_ids.to(torch.int32),
topk_weights=topk_weights.to(torch.float32),
quant_type=AiterQuantType.PER_1X32,
)
return cls.runner.run(runner_input, quant_info, {}).hidden_states
def _assert_numerics(self, actual, expected, max_abs=None):
self.assertTrue(torch.isfinite(actual).all())
actual_float = actual.float()
expected_float = expected.float()
cosine = F.cosine_similarity(
actual_float.flatten().unsqueeze(0),
expected_float.flatten().unsqueeze(0),
).item()
self.assertGreater(cosine, 0.98)
relative_l2 = (
torch.linalg.vector_norm(actual_float - expected_float)
/ torch.linalg.vector_norm(expected_float).clamp(min=1e-12)
).item()
self.assertLess(relative_l2, 0.20)
if max_abs is not None:
self.assertLess(
(actual_float - expected_float).abs().max().item(),
max_abs,
)
def test_top1_and_top8_match_dequantized_oracle(self):
for tokens in (1, 8, 17, 32, 64, 128):
generator = torch.Generator(device="cuda")
generator.manual_seed(tokens)
hidden = (
torch.randn(
tokens,
self.hidden_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.5
)
# topk=9 models eight routed experts plus one fused shared slot.
for topk in (1, 8, 9):
with self.subTest(tokens=tokens, topk=topk):
ids = torch.arange(topk, device="cuda", dtype=torch.int64).repeat(
tokens, 1
)
weights = torch.rand(
tokens,
topk,
generator=generator,
device="cuda",
dtype=torch.float32,
)
if topk > 1:
weights /= weights.sum(dim=-1, keepdim=True)
expected = self._torch_oracle(hidden, ids, weights)
actual = self._aiter(hidden, ids, weights)
repeated = self._aiter(hidden, ids, weights)
self._assert_numerics(actual, expected, max_abs=0.75)
if topk == 1:
torch.testing.assert_close(actual, repeated, atol=0, rtol=0)
else:
# Stage-2 combines top-k routes with atomics; reduction
# order may differ while remaining BF16-equivalent.
torch.testing.assert_close(
actual, repeated, atol=2e-2, rtol=1e-2
)
def test_clamp_boundary(self):
hidden = torch.full(
(1, self.hidden_size),
4.0,
device="cuda",
dtype=torch.bfloat16,
)
ids = torch.tensor([[0]], device="cuda", dtype=torch.int64)
weights = torch.ones((1, 1), device="cuda", dtype=torch.float32)
expected = self._torch_oracle(hidden, ids, weights)
actual = self._aiter(hidden, ids, weights)
self._assert_numerics(actual, expected)
def test_plain_block_fp8_matches_separated_oracle(self):
generator = torch.Generator(device="cuda")
generator.manual_seed(1234)
gate = (
torch.randn(
self.intermediate_size,
self.hidden_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.05
)
up = (
torch.randn(
self.intermediate_size,
self.hidden_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.05
)
down = (
torch.randn(
self.hidden_size,
self.intermediate_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.01
)
gate_q, gate_s = self._quantize_fp8_weight(gate)
up_q, up_s = self._quantize_fp8_weight(up)
down_q, down_s = self._quantize_fp8_weight(down)
w13_raw = torch.cat([gate_q, up_q], dim=0).unsqueeze(0)
w13_scale = torch.cat([gate_s, up_s], dim=0).unsqueeze(0)
w2_raw = down_q.unsqueeze(0)
w2_scale = down_s.unsqueeze(0)
w13 = shuffle_weight(w13_raw.contiguous(), (16, 16))
w2 = shuffle_weight(w2_raw.contiguous(), (16, 16))
quant_info = AiterMoeQuantInfo(
w13_weight=w13,
w2_weight=w2,
quant_type=AiterQuantType.PER_128X128,
w13_scale=w13_scale,
w2_scale=w2_scale,
swiglu_limit=self.swiglu_limit,
fused_moe_kwargs={"gate_mode": GateMode.SEPARATED.value},
)
gate_deq = self._dequantize_fp8_weight(gate_q, gate_s)
up_deq = self._dequantize_fp8_weight(up_q, up_s)
down_deq = self._dequantize_fp8_weight(down_q, down_s)
for tokens in (1, 8, 32):
with self.subTest(tokens=tokens):
hidden = (
torch.randn(
tokens,
self.hidden_size,
generator=generator,
device="cuda",
dtype=torch.bfloat16,
)
* 0.5
)
hidden_qdq = self._quant_dequant_fp8_activation(hidden)
gate_out = F.linear(hidden_qdq, gate_deq).clamp(max=self.swiglu_limit)
up_out = F.linear(hidden_qdq, up_deq).clamp(
-self.swiglu_limit, self.swiglu_limit
)
activated = self._quant_dequant_fp8_activation(
(F.silu(gate_out) * up_out).bfloat16()
)
expected = F.linear(activated, down_deq).bfloat16()
runner_input = AiterRunnerInput(
hidden_states=hidden,
topk_ids=torch.zeros((tokens, 1), device="cuda", dtype=torch.int32),
topk_weights=torch.ones(
(tokens, 1), device="cuda", dtype=torch.float32
),
quant_type=AiterQuantType.PER_128X128,
)
actual = self.runner.run(runner_input, quant_info, {}).hidden_states
self._assert_numerics(actual, expected, max_abs=0.75)
if __name__ == "__main__":
unittest.main()
@@ -7,14 +7,18 @@ the MxFP4 wrapper methods borrow an `Fp8MoEMethod` for weight loading only
and never give it a `moe_runner_config` (issue #36264). and never give it a `moe_runner_config` (issue #36264).
""" """
import sys
import types
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import torch import torch
from sglang.srt.layers.moe.moe_runner.aiter import AiterQuantType
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
from sglang.srt.layers.moe.utils import MoeRunnerBackend from sglang.srt.layers.moe.utils import MoeRunnerBackend
from sglang.srt.layers.quantization import fp8 as fp8_module
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8MoEMethod from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8MoEMethod
from sglang.srt.runtime_context import get_flags from sglang.srt.runtime_context import get_flags
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
@@ -114,5 +118,49 @@ class TestFp8MoERunnerOwnership(CustomTestCase):
self._assert_activation_params_absent(layer) self._assert_activation_params_absent(layer)
class TestFp8MoEAiterQuantInfo(CustomTestCase):
"""maybe_get_hip_aiter_quant_info assembles what the AITER runner consumes.
The gfx950 e2e builds AiterMoeQuantInfo by hand, so dropping the gate/up
layout or the clamp here would leave it passing while served experts read
the gate and up halves swapped.
"""
def test_block_fp8_forwards_separated_layout_and_clamp(self):
method = Fp8MoEMethod(
Fp8Config(is_checkpoint_fp8_serialized=True, weight_block_size=[128, 128])
)
# create_moe_runner is not called: it resolves a global backend and
# builds a MoeRunner, none of which this assembly reads.
method.moe_runner_config = MoeRunnerConfig(swiglu_limit=10.0)
layer = SimpleNamespace(
w13_weight=torch.zeros((1, 4, 4), dtype=torch.float8_e4m3fn),
w2_weight=torch.zeros((1, 4, 2), dtype=torch.float8_e4m3fn),
w13_weight_scale_inv=torch.ones((1, 4, 1), dtype=torch.float32),
w2_weight_scale_inv=torch.ones((1, 4, 1), dtype=torch.float32),
hidden_pad=0,
intermediate_pad=0,
_aiter_gate_up_interleaved=False,
dispatcher=SimpleNamespace(expert_mask_gpu=torch.tensor([True, False])),
)
fake_moe_common = types.ModuleType("aiter.ops.flydsl.moe_common")
fake_moe_common.GateMode = SimpleNamespace(
SEPARATED=SimpleNamespace(value="separated"),
INTERLEAVE=SimpleNamespace(value="interleave"),
)
with (
patch.dict(sys.modules, {"aiter.ops.flydsl.moe_common": fake_moe_common}),
patch.object(fp8_module, "_use_aiter", True),
):
quant_info = method.maybe_get_hip_aiter_quant_info(layer)
self.assertIsNotNone(quant_info)
self.assertEqual(quant_info.quant_type, AiterQuantType.PER_128X128)
self.assertEqual(quant_info.swiglu_limit, 10.0)
self.assertEqual(quant_info.fused_moe_kwargs, {"gate_mode": "separated"})
self.assertIs(quant_info.expert_mask, layer.dispatcher.expert_mask_gpu)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -1,21 +1,35 @@
"""Unit tests for QuarkConfig — CPU-only, no model loading.""" """Unit tests for QuarkConfig and its MoE scheme — CPU-only, no model loading."""
from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=12, suite="base-a-test-cpu") register_cpu_ci(est_time=12, suite="base-a-test-cpu")
import sys
import types
import unittest import unittest
from copy import deepcopy
from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import torch import torch
from sglang.srt.layers.linear import LinearBase
from sglang.srt.layers.moe.moe_runner.aiter import AiterQuantType
from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod
from sglang.srt.layers.quantization.quark.quark import ( from sglang.srt.layers.quantization.quark.quark import (
QuarkConfig, QuarkConfig,
_build_mixed_precision_layer_quant_config, _build_mixed_precision_layer_quant_config,
_mixed_precision_layer_map, _mixed_precision_layer_map,
_parse_nvfp4_excludes, _parse_nvfp4_excludes,
) )
from sglang.srt.layers.quantization.quark.schemes import (
quark_w4a4_mxfp4_moe as quark_moe,
)
from sglang.srt.layers.quantization.quark.schemes.quark_w4a4_mxfp4_moe import (
QuarkW4A4MXFp4MoE,
)
from sglang.srt.layers.quantization.quark.utils import check_equal_or_regex_match from sglang.srt.layers.quantization.quark.utils import check_equal_or_regex_match
from sglang.srt.models.glm5_next import Glm5NextForConditionalGeneration
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
_GET_CAP = "sglang.srt.layers.quantization.quark.quark.get_device_capability" _GET_CAP = "sglang.srt.layers.quantization.quark.quark.get_device_capability"
@@ -207,5 +221,175 @@ class TestParseNvfp4Excludes(CustomTestCase):
) )
class TestQuarkPerLayerBlockFp8(CustomTestCase):
_BLOCK_FP8_CONFIG = {
"weight": {
"dtype": "fp8_e4m3",
"qscheme": "per_block",
"block_size": [128, 128],
"is_dynamic": False,
},
"input_tensors": {
"dtype": "fp8_e4m3",
"qscheme": "per_group",
"group_size": 128,
"is_dynamic": True,
},
"output_tensors": None,
"bias": None,
}
def _build_bare_config(self) -> QuarkConfig:
config = _bare_config()
config.quant_config = {
"layer_quant_config": {
"model.language_model.layers.0.mlp.down_proj": self._BLOCK_FP8_CONFIG
},
"layer_type_quant_config": {},
"global_quant_config": {
"weight": {
"dtype": "fp4",
"qscheme": "per_group",
"group_size": 32,
"is_dynamic": False,
"scale_format": "e8m0",
},
"input_tensors": {
"dtype": "fp4",
"qscheme": "per_group",
"group_size": 32,
"is_dynamic": True,
"scale_format": "e8m0",
},
},
}
config.exclude_layers = []
config.kv_cache_group = []
config.packed_modules_mapping = {}
config.excluded_fp8_config = None
config._online_quantized_layers = set()
return config
def test_model_mapper_rewrites_explicit_layer_config(self):
config = self._build_bare_config()
config.apply_weight_name_mapper(
Glm5NextForConditionalGeneration.hf_to_sglang_mapper
)
self.assertIn(
"model.layers.0.mlp.down_proj",
config.quant_config["layer_quant_config"],
)
def test_model_mapper_rewrites_fused_visual_exclusion(self):
config = self._build_bare_config()
config.exclude_layers = ["model.visual.blocks.0.attn.qkv"]
config.apply_weight_name_mapper(
Glm5NextForConditionalGeneration.hf_to_sglang_mapper
)
self.assertEqual(
config.exclude_layers,
["visual.blocks.0.attn.qkv_proj"],
)
self.assertNotIn(
"model.language_model.layers.0.mlp.down_proj",
config.quant_config["layer_quant_config"],
)
def test_explicit_block_fp8_linear_uses_fp8_method(self):
config = self._build_bare_config()
config.apply_weight_name_mapper(
Glm5NextForConditionalGeneration.hf_to_sglang_mapper
)
layer = LinearBase.__new__(LinearBase)
method = config.get_quant_method(layer, "model.layers.0.mlp.down_proj")
self.assertIsInstance(method, Fp8LinearMethod)
self.assertTrue(method.quant_config.is_checkpoint_fp8_serialized)
self.assertEqual(method.quant_config.weight_block_size, [128, 128])
def test_dynamic_block_fp8_weight_is_not_treated_as_serialized(self):
layer_config = deepcopy(self._BLOCK_FP8_CONFIG)
layer_config["weight"]["is_dynamic"] = True
self.assertIsNone(QuarkConfig._get_block_fp8_config(layer_config, {}))
def test_unmatched_layer_still_uses_global_quark_config(self):
config = self._build_bare_config()
config.apply_weight_name_mapper(
Glm5NextForConditionalGeneration.hf_to_sglang_mapper
)
matched = config._find_matched_config(
"model.layers.4.mlp.down_proj", torch.nn.Module()
)
self.assertEqual(matched["weight"]["dtype"], "fp4")
class _Runner:
"""Records the quant_info apply_weights() hands to the runner."""
def __init__(self):
self.quant_info = None
def run(self, dispatch_output, quant_info):
self.quant_info = quant_info
return dispatch_output
class TestQuarkMxfp4MoEAiterQuantInfo(CustomTestCase):
"""apply_weights assembles what the AITER runner consumes.
The gfx950 e2e builds AiterMoeQuantInfo by hand, so dropping the gate/up
layout, the clamp or the padding here would leave it passing while served
experts read the gate and up halves swapped.
"""
def test_apply_forwards_clamp_separated_layout_and_padding(self):
scheme = object.__new__(QuarkW4A4MXFp4MoE)
scheme.moe_runner_config = SimpleNamespace(swiglu_limit=10.0)
scheme.runner = _Runner()
layer = SimpleNamespace(
w13_weight=torch.zeros((1, 4, 2), dtype=torch.uint8),
w2_weight=torch.zeros((1, 2, 2), dtype=torch.uint8),
w13_weight_scale=torch.ones((1, 4, 1), dtype=torch.uint8),
w2_weight_scale=torch.ones((1, 2, 1), dtype=torch.uint8),
hidden_pad=0,
intermediate_pad=128,
dispatcher=SimpleNamespace(expert_mask_gpu=torch.tensor([True, False])),
)
layer.w13_weight.is_shuffled = True
fake_moe_common = types.ModuleType("aiter.ops.flydsl.moe_common")
fake_moe_common.GateMode = SimpleNamespace(
SEPARATED=SimpleNamespace(value="separated"),
INTERLEAVE=SimpleNamespace(value="interleave"),
)
with (
patch.dict(sys.modules, {"aiter.ops.flydsl.moe_common": fake_moe_common}),
patch.object(quark_moe, "_is_gfx95", True),
patch.object(quark_moe, "_is_gfx1250", False),
):
marker = object()
result = scheme.apply_weights(layer, marker)
self.assertIs(result, marker)
quant_info = scheme.runner.quant_info
self.assertEqual(quant_info.quant_type, AiterQuantType.PER_1X32)
self.assertEqual(quant_info.swiglu_limit, 10.0)
self.assertEqual(quant_info.hidden_pad, 0)
self.assertEqual(quant_info.intermediate_pad, 128)
self.assertEqual(quant_info.fused_moe_kwargs, {"gate_mode": "separated"})
self.assertIs(quant_info.expert_mask, layer.dispatcher.expert_mask_gpu)
self.assertTrue(quant_info.w13_weight.is_shuffled)
self.assertTrue(quant_info.w2_weight.is_shuffled)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()