[AMD] Implement QuarkW4A8MXFp4MoE to support amd/gpt-oss-120b-w-mxfp4-a-fp8 (#27204)
Signed-off-by: Stanley Winata <stanley.winata@amd.com>
This commit is contained in:
@@ -20,6 +20,7 @@ from sglang.srt.layers.quantization.quark.schemes import (
|
||||
QuarkMoEScheme,
|
||||
QuarkW4A4MXFP4,
|
||||
QuarkW4A4MXFp4MoE,
|
||||
QuarkW4A8MXFp4MoE,
|
||||
QuarkW8A8Fp8,
|
||||
QuarkW8A8FP8MoE,
|
||||
)
|
||||
@@ -385,6 +386,28 @@ class QuarkConfig(QuantizationConfig):
|
||||
|
||||
return True
|
||||
|
||||
def _is_mx_w4a8(
|
||||
self,
|
||||
weight_quant: Optional[dict[str, Any]],
|
||||
input_quant: Optional[dict[str, Any]],
|
||||
) -> bool:
|
||||
if weight_quant is None or input_quant is None:
|
||||
return False
|
||||
|
||||
is_mx_fp4_weight = (
|
||||
weight_quant.get("dtype") == "fp4"
|
||||
and weight_quant.get("qscheme") == "per_group"
|
||||
and weight_quant.get("group_size") == 32
|
||||
and not weight_quant.get("is_dynamic")
|
||||
and weight_quant.get("scale_format") == "e8m0"
|
||||
)
|
||||
is_static_fp8_activation = (
|
||||
input_quant.get("dtype") in ("fp8_e4m3", "fp8_e4m3fn")
|
||||
and input_quant.get("qscheme") == "per_tensor"
|
||||
and not input_quant.get("is_dynamic")
|
||||
)
|
||||
return is_mx_fp4_weight and is_static_fp8_activation
|
||||
|
||||
def _find_matched_config(
|
||||
self, layer_name: str, module: torch.nn.Module
|
||||
) -> dict[str, Any]:
|
||||
@@ -496,6 +519,9 @@ class QuarkConfig(QuantizationConfig):
|
||||
input_config,
|
||||
is_checkpoint_mxfp4_serialized=self.is_prequantized,
|
||||
)
|
||||
elif self._is_mx_w4a8(weight_config, input_config):
|
||||
logger.info_once("Using Quark MXFP4-W/FP8-A MoE scheme")
|
||||
return QuarkW4A8MXFp4MoE(weight_config, input_config)
|
||||
elif self._is_fp8_w8a8(weight_config, input_config):
|
||||
return QuarkW8A8FP8MoE(weight_config, input_config)
|
||||
else:
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
from .quark_scheme import QuarkLinearScheme, QuarkMoEScheme
|
||||
from .quark_w4a4_mxfp4 import QuarkW4A4MXFP4
|
||||
from .quark_w4a4_mxfp4_moe import QuarkW4A4MXFp4MoE
|
||||
from .quark_w4a8_mxfp4_moe import QuarkW4A8MXFp4MoE
|
||||
from .quark_w8a8_fp8 import QuarkW8A8Fp8
|
||||
from .quark_w8a8_fp8_moe import QuarkW8A8FP8MoE
|
||||
|
||||
@@ -12,5 +13,6 @@ __all__ = [
|
||||
"QuarkW4A4MXFP4",
|
||||
"QuarkW8A8Fp8",
|
||||
"QuarkW4A4MXFp4MoE",
|
||||
"QuarkW4A8MXFp4MoE",
|
||||
"QuarkW8A8FP8MoE",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,407 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import replace
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.utils import get_moe_weight_sizes
|
||||
from sglang.srt.layers.quantization.quark.schemes import QuarkMoEScheme
|
||||
from sglang.srt.layers.quantization.utils import all_close_1d
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
is_gfx95_supported,
|
||||
is_hip,
|
||||
round_up,
|
||||
set_weight_attrs,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.token_dispatcher import (
|
||||
CombineInput,
|
||||
StandardDispatchOutput,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_is_shuffle_moe_mxfp4 = is_gfx95_supported()
|
||||
|
||||
__all__ = ["QuarkW4A8MXFp4MoE"]
|
||||
|
||||
_is_hip = is_hip()
|
||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
if _use_aiter:
|
||||
from aiter.ops.shuffle import (
|
||||
shuffle_scale,
|
||||
shuffle_scale_a16w4,
|
||||
shuffle_weight,
|
||||
shuffle_weight_a16w4,
|
||||
)
|
||||
|
||||
OCP_MX_BLOCK_SIZE = 32
|
||||
|
||||
|
||||
class QuarkW4A8MXFp4MoE(QuarkMoEScheme):
|
||||
"""Quark MoE scheme for MXFP4 weights with static FP8 activations."""
|
||||
|
||||
def __init__(self, weight_config: dict[str, Any], input_config: dict[str, Any]):
|
||||
self.weight_quant = weight_config
|
||||
self.input_quant = input_config
|
||||
|
||||
weight_qscheme = self.weight_quant.get("qscheme")
|
||||
input_qscheme = self.input_quant.get("qscheme")
|
||||
weight_dtype = self.weight_quant.get("dtype")
|
||||
input_dtype = self.input_quant.get("dtype")
|
||||
|
||||
if not (
|
||||
weight_dtype == "fp4"
|
||||
and weight_qscheme == "per_group"
|
||||
and self.weight_quant.get("group_size") == OCP_MX_BLOCK_SIZE
|
||||
and not self.weight_quant.get("is_dynamic")
|
||||
and self.weight_quant.get("scale_format") == "e8m0"
|
||||
):
|
||||
raise ValueError(
|
||||
"For W4A8 MXFP4-FP8 Fused MoE layers, weights must be "
|
||||
"static per-group FP4 with group_size=32 and e8m0 scales. "
|
||||
f"Found {self.weight_quant}."
|
||||
)
|
||||
|
||||
if not (
|
||||
input_dtype in ("fp8_e4m3", "fp8_e4m3fn")
|
||||
and input_qscheme == "per_tensor"
|
||||
and not self.input_quant.get("is_dynamic")
|
||||
):
|
||||
raise ValueError(
|
||||
"For W4A8 MXFP4-FP8 Fused MoE layers, activations must be "
|
||||
"static per-tensor fp8_e4m3/fp8_e4m3fn. "
|
||||
f"Found {self.input_quant}."
|
||||
)
|
||||
|
||||
self.with_bias = False
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 70
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
num_experts: int,
|
||||
hidden_size: int,
|
||||
intermediate_size_per_partition: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
|
||||
|
||||
self.num_experts = num_experts
|
||||
self.with_bias = extra_weight_attrs.get("with_bias", False)
|
||||
if _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
|
||||
)
|
||||
else:
|
||||
intermediate_size_per_partition_after_pad = intermediate_size_per_partition
|
||||
self.hidden_pad = 0
|
||||
self.intermediate_pad = 0
|
||||
|
||||
w13_up_dim, w2_down_dim, weight_padded = get_moe_weight_sizes(
|
||||
intermediate_size_per_partition_after_pad,
|
||||
is_aiter_moe=_use_aiter,
|
||||
is_concat=True,
|
||||
is_packed=True,
|
||||
)
|
||||
self.intermediate_size_per_partition = intermediate_size_per_partition_after_pad
|
||||
self.hidden_size = hidden_size
|
||||
|
||||
# Add the quantization method used (per tensor/grouped/channel)
|
||||
# to ensure the weight scales are loaded in properly.
|
||||
extra_weight_attrs.update(
|
||||
{
|
||||
"quant_method": FusedMoeWeightScaleSupported.BLOCK.value,
|
||||
"weight_padded": weight_padded,
|
||||
},
|
||||
)
|
||||
|
||||
weight_dtype = torch.uint8
|
||||
|
||||
# WEIGHTS
|
||||
# MXFP4 weights are stored as uint8, with two FP4 values packed per
|
||||
# byte. The AITER path later views these buffers as float4_e2m1fn_x2.
|
||||
# Use ``zeros`` (not ``empty``) so the alignment padding (hidden
|
||||
# 2880->3072, intermediate 2880->3072 for GPT-OSS) dequantizes to
|
||||
# 0.0 if it ever reaches the matmul. The current AITER kernel
|
||||
# skips the padded tail via ``n_pad_zeros`` / ``k_pad_zeros`` so
|
||||
# this is defensive, but it matches ``Mxfp4MoEMethod``'s
|
||||
# convention for the same kernel.
|
||||
w13_weight = torch.nn.Parameter(
|
||||
torch.zeros(
|
||||
num_experts,
|
||||
w13_up_dim,
|
||||
hidden_size // 2,
|
||||
dtype=weight_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight", w13_weight)
|
||||
set_weight_attrs(w13_weight, extra_weight_attrs)
|
||||
|
||||
w2_weight = torch.nn.Parameter(
|
||||
torch.zeros(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
w2_down_dim,
|
||||
dtype=weight_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight", w2_weight)
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
w13_weight_bias = torch.nn.Parameter(
|
||||
torch.zeros(
|
||||
num_experts,
|
||||
w13_up_dim,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_bias", w13_weight_bias)
|
||||
set_weight_attrs(w13_weight_bias, extra_weight_attrs)
|
||||
|
||||
w2_weight_bias = torch.nn.Parameter(
|
||||
torch.zeros(num_experts, hidden_size, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_bias", w2_weight_bias)
|
||||
set_weight_attrs(w2_weight_bias, extra_weight_attrs)
|
||||
|
||||
# WEIGHT_SCALES
|
||||
# MXFP4 uses one e8m0 scale per 32-value block. These scales are
|
||||
# loaded as uint8 and shuffled after loading for the kernel layout.
|
||||
w13_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
w13_up_dim,
|
||||
hidden_size // OCP_MX_BLOCK_SIZE,
|
||||
dtype=weight_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
# 1. w2 scale is floor division of inter_dim by blockscale.
|
||||
# 2. w2 scale needs to scale up just as w2.
|
||||
# We combine 1. and 2. to keep the integer precision.
|
||||
w2_weight_scale = torch.nn.Parameter(
|
||||
torch.ones(
|
||||
num_experts,
|
||||
hidden_size,
|
||||
(w2_down_dim * 2) // OCP_MX_BLOCK_SIZE,
|
||||
dtype=weight_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(w2_weight_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w13_weight_scale, extra_weight_attrs)
|
||||
|
||||
layer.register_parameter("w13_weight_scale", w13_weight_scale)
|
||||
layer.register_parameter("w2_weight_scale", w2_weight_scale)
|
||||
|
||||
# Add the quantization method used (per tensor/grouped/channel)
|
||||
# to ensure the activation scales are loaded in properly.
|
||||
extra_weight_attrs.update(
|
||||
{"quant_method": FusedMoeWeightScaleSupported.TENSOR.value}
|
||||
)
|
||||
|
||||
# INPUT_SCALES
|
||||
# W4A8 checkpoints carry static per-tensor FP8 activation scales for
|
||||
# gate_up_proj and down_proj. These are separate from the MXFP4 weight
|
||||
# block scales above.
|
||||
w13_input_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
w2_input_scale = torch.nn.Parameter(
|
||||
torch.ones(num_experts, dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_input_scale", w13_input_scale)
|
||||
layer.register_parameter("w2_input_scale", w2_input_scale)
|
||||
set_weight_attrs(w13_input_scale, extra_weight_attrs)
|
||||
set_weight_attrs(w2_input_scale, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
# Mirror native MXFP4 post-load shuffling. The default
|
||||
# `SGLANG_USE_AITER_MOE_GU_ITLV=1` path uses the gate-up-aware
|
||||
# a16w4 layout; the `=0` fallback keeps the separated gate/up layout.
|
||||
# The Quark loader (`_load_quark_experts_weights` in
|
||||
# `python/sglang/srt/models/gpt_oss.py`) already writes the
|
||||
# SEPARATED-layout `[g0..g_{N-1}, u0..u_{N-1}]` buffer per expert,
|
||||
# which is exactly the starting state the native path is in after
|
||||
# its post-load `.view(e, n//2, 2, k).permute(0, 2, 1, 3)` step.
|
||||
if envs.SGLANG_USE_AITER_MOE_GU_ITLV.get():
|
||||
if _is_shuffle_moe_mxfp4:
|
||||
layer.w13_weight.data = shuffle_weight_a16w4(
|
||||
layer.w13_weight.contiguous(), 16, True
|
||||
)
|
||||
layer.w2_weight.data = shuffle_weight_a16w4(
|
||||
layer.w2_weight.contiguous(), 16, False
|
||||
)
|
||||
layer.w13_weight.is_shuffled = True
|
||||
layer.w2_weight.is_shuffled = True
|
||||
shuffled_w13_scale = shuffle_scale_a16w4(
|
||||
layer.w13_weight_scale.view(-1, layer.w13_weight_scale.shape[-1]),
|
||||
self.num_experts,
|
||||
True,
|
||||
)
|
||||
shuffled_w2_scale = shuffle_scale_a16w4(
|
||||
layer.w2_weight_scale.view(-1, layer.w2_weight_scale.shape[-1]),
|
||||
self.num_experts,
|
||||
False,
|
||||
)
|
||||
else:
|
||||
if _is_shuffle_moe_mxfp4:
|
||||
layer.w13_weight.data = shuffle_weight(
|
||||
layer.w13_weight.contiguous(),
|
||||
is_guinterleave=False,
|
||||
gate_up=True,
|
||||
)
|
||||
layer.w2_weight.data = shuffle_weight(
|
||||
layer.w2_weight.contiguous(),
|
||||
is_guinterleave=False,
|
||||
gate_up=False,
|
||||
)
|
||||
layer.w13_weight.is_shuffled = True
|
||||
layer.w2_weight.is_shuffled = True
|
||||
shuffled_w13_scale = shuffle_scale(
|
||||
layer.w13_weight_scale.view(-1, layer.w13_weight_scale.shape[-1]),
|
||||
experts_cnt=self.num_experts,
|
||||
is_guinterleave=False,
|
||||
gate_up=True,
|
||||
)
|
||||
shuffled_w2_scale = shuffle_scale(
|
||||
layer.w2_weight_scale.view(-1, layer.w2_weight_scale.shape[-1]),
|
||||
experts_cnt=self.num_experts,
|
||||
is_guinterleave=False,
|
||||
gate_up=False,
|
||||
)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
# Static FP8 MoE kernels consume a single activation scale. Use the
|
||||
# maximum if expert-local checkpoint scales differ.
|
||||
if layer.w13_input_scale is None or layer.w2_input_scale is None:
|
||||
raise ValueError("W4A8 MXFP4-FP8 MoE requires static input scales.")
|
||||
if not all_close_1d(layer.w13_input_scale) or not all_close_1d(
|
||||
layer.w2_input_scale
|
||||
):
|
||||
logger.warning(
|
||||
"Found input_scales that are not equal for W4A8 MXFP4-FP8 "
|
||||
"MoE layer. Using the maximum across experts for each layer."
|
||||
)
|
||||
layer.w13_input_scale = torch.nn.Parameter(
|
||||
layer.w13_input_scale.max().to(torch.float32), requires_grad=False
|
||||
)
|
||||
layer.w2_input_scale = torch.nn.Parameter(
|
||||
layer.w2_input_scale.max().to(torch.float32), requires_grad=False
|
||||
)
|
||||
|
||||
if hasattr(layer, "dispatcher"):
|
||||
# Weights are stored as torch.uint8 but semantically MXFP4
|
||||
layer.dispatcher.set_quant_config({"weight_dtype": torch.float4_e2m1fn_x2})
|
||||
|
||||
def create_moe_runner(
|
||||
self, layer: torch.nn.Module, moe_runner_config: MoeRunnerConfig
|
||||
):
|
||||
from sglang.srt.layers.moe.utils import (
|
||||
get_moe_a2a_backend,
|
||||
get_moe_runner_backend,
|
||||
)
|
||||
|
||||
self.moe_runner_config = moe_runner_config
|
||||
moe_runner_backend = get_moe_runner_backend()
|
||||
if _use_aiter and get_moe_a2a_backend().supports_aiter():
|
||||
moe_runner_backend = MoeRunnerBackend.AITER
|
||||
|
||||
if moe_runner_backend.is_aiter():
|
||||
# MXFP4 hard-codes Swiglu in the AITER kernel path.
|
||||
self.runner = MoeRunner(
|
||||
moe_runner_backend, replace(moe_runner_config, activation="swiglu")
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"QuarkW4A8MXFp4MoE is currently only supported with AITER."
|
||||
)
|
||||
|
||||
def apply_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
dispatch_output: StandardDispatchOutput,
|
||||
) -> CombineInput:
|
||||
from sglang.srt.layers.moe.moe_runner.aiter import (
|
||||
AiterMoeQuantInfo,
|
||||
AiterQuantType,
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
if hasattr(layer.w13_weight, "is_shuffled"):
|
||||
w13_weight.is_shuffled = True
|
||||
w2_weight.is_shuffled = True
|
||||
|
||||
x_padded = torch.nn.functional.pad(
|
||||
dispatch_output.hidden_states,
|
||||
(0, self.hidden_pad),
|
||||
mode="constant",
|
||||
value=0.0,
|
||||
)
|
||||
quant_info = AiterMoeQuantInfo(
|
||||
w13_weight=w13_weight,
|
||||
w2_weight=w2_weight,
|
||||
quant_type=AiterQuantType.PER_1X32,
|
||||
w13_scale=layer.w13_weight_scale,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
a13_scale=layer.w13_input_scale,
|
||||
a2_scale=layer.w2_input_scale,
|
||||
b13=layer.w13_weight_bias,
|
||||
b2=layer.w2_weight_bias,
|
||||
expert_mask=layer.dispatcher.expert_mask_gpu,
|
||||
doweight_stage1=self.moe_runner_config.apply_router_weight_on_input,
|
||||
hidden_pad=self.hidden_pad,
|
||||
intermediate_pad=self.intermediate_pad,
|
||||
# gpt-oss populates `gemm1_clamp_limit` (renamed in
|
||||
# `models/gpt_oss.py` from `config.swiglu_limit`); DSv4 populates
|
||||
# `swiglu_limit` directly. Accept either so the AITER `gate_mode`
|
||||
# + `swiglu_limit` dispatch block in `moe_runner/aiter.py` (gated
|
||||
# on `quant_info.swiglu_limit > 0`) is actually entered for both
|
||||
# families. Mirrors the same fix PR #27201 applied to the native
|
||||
# `Mxfp4MoEMethod.apply` path.
|
||||
swiglu_limit=(
|
||||
self.moe_runner_config.gemm1_clamp_limit
|
||||
or self.moe_runner_config.swiglu_limit
|
||||
or 0.0
|
||||
),
|
||||
)
|
||||
return self.runner.run(
|
||||
dispatch_output._replace(hidden_states=x_padded), quant_info
|
||||
)
|
||||
@@ -0,0 +1,248 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
import re
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_moe_expert_parallel_rank,
|
||||
get_moe_expert_parallel_world_size,
|
||||
get_moe_tensor_parallel_rank,
|
||||
get_moe_tensor_parallel_world_size,
|
||||
)
|
||||
from sglang.srt.utils import is_cuda
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
|
||||
|
||||
def load_gptoss_weight_quark(
|
||||
model,
|
||||
weights,
|
||||
*,
|
||||
is_nextn: bool,
|
||||
weight_name_mapping,
|
||||
) -> None:
|
||||
# Regex matching `model.layers.{L}.mlp.experts.{N}.{gate_up_proj|down_proj}.{suffix}`
|
||||
# used by the AMD Quark GPT-OSS per-expert checkpoint layout.
|
||||
quark_expert_pat = re.compile(
|
||||
r"^(.*\.mlp\.experts)\.(\d+)\.(gate_up_proj|down_proj)\."
|
||||
r"(weight|weight_scale|input_scale|bias)$"
|
||||
)
|
||||
quark_experts_weights = []
|
||||
normal_weights = []
|
||||
|
||||
for name, weight in weights:
|
||||
if quark_expert_pat.match(name) is not None:
|
||||
quark_experts_weights.append((name, weight))
|
||||
else:
|
||||
normal_weights.append((name, weight))
|
||||
|
||||
quark_loaded = _load_gptoss_quark_expert_weights(
|
||||
model, quark_experts_weights, quark_expert_pat
|
||||
)
|
||||
model._load_normal_weights(
|
||||
normal_weights,
|
||||
is_nextn=is_nextn,
|
||||
weight_name_mapping=weight_name_mapping,
|
||||
other_loaded_param_names=quark_loaded,
|
||||
)
|
||||
|
||||
|
||||
def _load_gptoss_quark_expert_weights(model, weights, quark_expert_pat):
|
||||
"""GPT-OSS per-expert style loader for Quark MoE tensors into padded fused buffers.
|
||||
|
||||
Quark stores each expert separately:
|
||||
experts.{N}.gate_up_proj.{weight,weight_scale,input_scale,bias}
|
||||
experts.{N}.down_proj.{weight,weight_scale,input_scale,bias}
|
||||
|
||||
We mirror the static MXFP4 expert loader: slice the checkpoint along
|
||||
the TP-sharded dimension (intermediate axis) and copy into a window
|
||||
of the padded ``w13_*`` / ``w2_*`` parameters allocated by
|
||||
:class:`QuarkW4A8MXFp4MoE`. Down-proj bias is loaded only on
|
||||
``moe_tp_rank == 0`` to avoid double-counting after all-reduce.
|
||||
"""
|
||||
params_dict = dict(model.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
mxfp4_block = 32
|
||||
|
||||
moe_tp_rank = get_moe_tensor_parallel_rank()
|
||||
moe_tp_size = get_moe_tensor_parallel_world_size()
|
||||
moe_ep_rank = get_moe_expert_parallel_rank()
|
||||
moe_ep_size = get_moe_expert_parallel_world_size()
|
||||
|
||||
intermediate_size = model.config.intermediate_size
|
||||
assert (
|
||||
intermediate_size % mxfp4_block == 0
|
||||
), f"{intermediate_size=} must be divisible by {mxfp4_block=}"
|
||||
intermediate_size_block = intermediate_size // mxfp4_block
|
||||
|
||||
per_rank_intermediate_size_block = math.ceil(intermediate_size_block / moe_tp_size)
|
||||
|
||||
per_rank_intermediate_size = per_rank_intermediate_size_block * mxfp4_block
|
||||
|
||||
# Calculate common slicing bounds for current rank
|
||||
assert model.config.num_local_experts % moe_ep_size == 0
|
||||
moe_num_local_experts = model.config.num_local_experts // moe_ep_size
|
||||
|
||||
moe_tp_rank_start = moe_tp_rank * per_rank_intermediate_size
|
||||
moe_tp_rank_end = min(
|
||||
(moe_tp_rank + 1) * per_rank_intermediate_size, intermediate_size
|
||||
)
|
||||
|
||||
moe_ep_rank_start = moe_ep_rank * moe_num_local_experts
|
||||
moe_ep_rank_end = (moe_ep_rank + 1) * moe_num_local_experts
|
||||
|
||||
for name, weight in weights:
|
||||
# Quark stores experts separately as
|
||||
# `experts.{N}.{gate_up_proj|down_proj}.{suffix}`; pull the
|
||||
# expert id out of the name (mxfp4 has it as axis 0 instead).
|
||||
m = quark_expert_pat.match(name)
|
||||
if m is None:
|
||||
continue
|
||||
prefix, expert_str, proj, suffix = m.groups()
|
||||
global_expert_id = int(expert_str)
|
||||
if global_expert_id < moe_ep_rank_start or global_expert_id >= moe_ep_rank_end:
|
||||
continue
|
||||
local_expert_id = global_expert_id - moe_ep_rank_start
|
||||
|
||||
if _is_cuda:
|
||||
weight = weight.cuda()
|
||||
|
||||
dispatch_key = f"{proj}.{suffix}"
|
||||
|
||||
if dispatch_key == "gate_up_proj.weight":
|
||||
# Handle MLP gate and up projection weights
|
||||
new_name = f"{prefix}.w13_weight"
|
||||
|
||||
# De-interleave gate/up rows ([g0,u0,g1,u1,...] -> [g..., u...])
|
||||
# then slice the TP window. Each half is written into its own
|
||||
# slot of the padded fused buffer; the gap between halves is
|
||||
# pre-zeroed by `create_weights` and must not be overwritten.
|
||||
narrow_gate = weight[0::2][moe_tp_rank_start:moe_tp_rank_end].contiguous()
|
||||
narrow_up = weight[1::2][moe_tp_rank_start:moe_tp_rank_end].contiguous()
|
||||
|
||||
param = params_dict[new_name]
|
||||
intermediate_pad = param.data.shape[1] // 2
|
||||
g0, g1 = narrow_gate.shape
|
||||
u0, u1 = narrow_up.shape
|
||||
param.data[local_expert_id, :g0, :g1].copy_(
|
||||
narrow_gate.to(param.data.dtype)
|
||||
)
|
||||
param.data[
|
||||
local_expert_id,
|
||||
intermediate_pad : intermediate_pad + u0,
|
||||
:u1,
|
||||
].copy_(narrow_up.to(param.data.dtype))
|
||||
loaded_params.add(new_name)
|
||||
|
||||
elif dispatch_key == "down_proj.weight":
|
||||
# Handle MLP down projection weights
|
||||
# packed FP4 -> halve the TP bound on the contracting K dim
|
||||
new_name = f"{prefix}.w2_weight"
|
||||
|
||||
narrow_weight = weight[
|
||||
...,
|
||||
moe_tp_rank_start // 2 : moe_tp_rank_end // 2,
|
||||
]
|
||||
|
||||
param = params_dict[new_name]
|
||||
d0, d1 = narrow_weight.shape
|
||||
param.data[local_expert_id, :d0, :d1].copy_(
|
||||
narrow_weight.to(param.data.dtype)
|
||||
)
|
||||
loaded_params.add(new_name)
|
||||
|
||||
elif dispatch_key == "gate_up_proj.weight_scale":
|
||||
# Handle MLP gate and up projection weight scales
|
||||
new_name = f"{prefix}.w13_weight_scale"
|
||||
|
||||
narrow_gate = weight[0::2][moe_tp_rank_start:moe_tp_rank_end].contiguous()
|
||||
narrow_up = weight[1::2][moe_tp_rank_start:moe_tp_rank_end].contiguous()
|
||||
|
||||
param = params_dict[new_name]
|
||||
intermediate_pad = param.data.shape[1] // 2
|
||||
g0, g1 = narrow_gate.shape
|
||||
u0, u1 = narrow_up.shape
|
||||
param.data[local_expert_id, :g0, :g1].copy_(
|
||||
narrow_gate.to(param.data.dtype)
|
||||
)
|
||||
param.data[
|
||||
local_expert_id,
|
||||
intermediate_pad : intermediate_pad + u0,
|
||||
:u1,
|
||||
].copy_(narrow_up.to(param.data.dtype))
|
||||
loaded_params.add(new_name)
|
||||
|
||||
elif dispatch_key == "down_proj.weight_scale":
|
||||
# Handle MLP down projection weight scales
|
||||
# 32 fp4 values per block -> slice by mxfp4_block
|
||||
new_name = f"{prefix}.w2_weight_scale"
|
||||
|
||||
narrow_weight = weight[
|
||||
...,
|
||||
moe_tp_rank_start // mxfp4_block : moe_tp_rank_end // mxfp4_block,
|
||||
]
|
||||
|
||||
param = params_dict[new_name]
|
||||
d0, d1 = narrow_weight.shape
|
||||
param.data[local_expert_id, :d0, :d1].copy_(
|
||||
narrow_weight.to(param.data.dtype)
|
||||
)
|
||||
loaded_params.add(new_name)
|
||||
|
||||
elif dispatch_key == "gate_up_proj.bias":
|
||||
# Handle MLP gate and up projection biases
|
||||
new_name = f"{prefix}.w13_weight_bias"
|
||||
|
||||
narrow_gate = weight[0::2][moe_tp_rank_start:moe_tp_rank_end].contiguous()
|
||||
narrow_up = weight[1::2][moe_tp_rank_start:moe_tp_rank_end].contiguous()
|
||||
|
||||
param = params_dict[new_name]
|
||||
intermediate_pad = param.data.shape[1] // 2
|
||||
param.data[local_expert_id, : narrow_gate.shape[0]].copy_(
|
||||
narrow_gate.to(param.data.dtype)
|
||||
)
|
||||
param.data[
|
||||
local_expert_id,
|
||||
intermediate_pad : intermediate_pad + narrow_up.shape[0],
|
||||
].copy_(narrow_up.to(param.data.dtype))
|
||||
loaded_params.add(new_name)
|
||||
|
||||
elif dispatch_key == "down_proj.bias":
|
||||
# Handle MLP down projection bias
|
||||
# Only TP rank 0 owns the bias; others zero out so the
|
||||
# post-MoE all-reduce sums to the correct value once.
|
||||
narrow_weight = weight
|
||||
if moe_tp_rank != 0:
|
||||
narrow_weight = torch.zeros_like(narrow_weight)
|
||||
|
||||
new_name = f"{prefix}.w2_weight_bias"
|
||||
param = params_dict[new_name]
|
||||
d0 = narrow_weight.shape[0]
|
||||
param.data[local_expert_id, :d0].copy_(narrow_weight.to(param.data.dtype))
|
||||
loaded_params.add(new_name)
|
||||
|
||||
elif dispatch_key == "gate_up_proj.input_scale":
|
||||
# Handle MLP gate/up FP8 activation scale (per-tensor scalar)
|
||||
new_name = f"{prefix}.w13_input_scale"
|
||||
if new_name not in params_dict:
|
||||
# Scheme didn't allocate the parameter (e.g. W4A16); skip.
|
||||
continue
|
||||
|
||||
param = params_dict[new_name]
|
||||
param.data[local_expert_id].copy_(weight.to(param.data.dtype).reshape(()))
|
||||
loaded_params.add(new_name)
|
||||
|
||||
elif dispatch_key == "down_proj.input_scale":
|
||||
# Handle MLP down FP8 activation scale (per-tensor scalar)
|
||||
new_name = f"{prefix}.w2_input_scale"
|
||||
if new_name not in params_dict:
|
||||
# Scheme didn't allocate the parameter (e.g. W4A16); skip.
|
||||
continue
|
||||
|
||||
param = params_dict[new_name]
|
||||
param.data[local_expert_id].copy_(weight.to(param.data.dtype).reshape(()))
|
||||
loaded_params.add(new_name)
|
||||
|
||||
return loaded_params
|
||||
@@ -879,12 +879,23 @@ class GptOssForCausalLM(nn.Module):
|
||||
quant_config_name = (
|
||||
self.quant_config.get_name() if self.quant_config is not None else None
|
||||
)
|
||||
if quant_config_name != "mxfp4":
|
||||
self._load_normal_weights(
|
||||
if quant_config_name == "mxfp4":
|
||||
self._load_weights_mxfp4(
|
||||
weights, is_nextn=is_nextn, weight_name_mapping=weight_name_mapping
|
||||
)
|
||||
elif quant_config_name == "quark":
|
||||
from sglang.srt.layers.quantization.quark.weights import (
|
||||
load_gptoss_weight_quark,
|
||||
)
|
||||
|
||||
load_gptoss_weight_quark(
|
||||
self,
|
||||
weights,
|
||||
is_nextn=is_nextn,
|
||||
weight_name_mapping=weight_name_mapping,
|
||||
)
|
||||
else:
|
||||
self._load_weights_mxfp4(
|
||||
self._load_normal_weights(
|
||||
weights, is_nextn=is_nextn, weight_name_mapping=weight_name_mapping
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
"""MI35x GPT-OSS W4A8 MXFP4-FP8 GSM8K Completion Evaluation Test (8-GPU)
|
||||
|
||||
Tests the AMD Quark `gpt-oss-120b-w-mxfp4-a-fp8` checkpoint (MXFP4
|
||||
weights + static per-tensor FP8 activations) using few-shot completion
|
||||
benchmark on MI35x.
|
||||
|
||||
Registry: nightly-amd-8-gpu-mi35x suite
|
||||
"""
|
||||
|
||||
import ast
|
||||
import os
|
||||
|
||||
# Set HF cache for MI35x
|
||||
os.environ.setdefault("HF_HOME", "/data2/models/huggingface")
|
||||
os.environ.setdefault("HF_HUB_CACHE", "/data2/models/huggingface/hub")
|
||||
|
||||
import re
|
||||
import time
|
||||
import unittest
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.test.ci.ci_register import register_amd_ci
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
is_in_ci,
|
||||
popen_launch_server,
|
||||
write_github_step_summary,
|
||||
)
|
||||
from sglang.utils import download_and_cache_file, read_jsonl
|
||||
|
||||
# Register for AMD CI - MI35x GPT-OSS W4A8 MXFP4-FP8 accuracy test (~15 min)
|
||||
register_amd_ci(est_time=900, suite="nightly-amd-8-gpu-mi35x", nightly=True)
|
||||
|
||||
INVALID = -9999999
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelConfig:
|
||||
"""Configuration for a model to test."""
|
||||
|
||||
model_path: str
|
||||
tp_size: int = 8
|
||||
accuracy_threshold: float = 0.50
|
||||
other_args: Optional[List[str]] = None
|
||||
env_vars: Optional[dict] = None
|
||||
timeout: Optional[int] = None
|
||||
|
||||
def __post_init__(self):
|
||||
if self.other_args is None:
|
||||
self.other_args = []
|
||||
if self.env_vars is None:
|
||||
self.env_vars = {}
|
||||
|
||||
|
||||
# GPT-OSS Quark W4A8 MXFP4-FP8 models for MI35x
|
||||
MI35X_GPT_OSS_W4A8_MXFP4_MODELS = [
|
||||
ModelConfig(
|
||||
model_path="amd/gpt-oss-120b-w-mxfp4-a-fp8",
|
||||
tp_size=8,
|
||||
# Match the peer `openai/gpt-oss-120b` GSM8K threshold in
|
||||
# `test_gpt_oss_eval_mi35x.py` (0.79). The AMD Quark W4A8 MXFP4-FP8
|
||||
# measures ~0.855 on this host, comfortably above 0.79.
|
||||
accuracy_threshold=0.79,
|
||||
timeout=900,
|
||||
other_args=[
|
||||
"--chunked-prefill-size",
|
||||
"130172",
|
||||
"--max-running-requests",
|
||||
"128",
|
||||
"--mem-fraction-static",
|
||||
"0.85",
|
||||
"--attention-backend",
|
||||
"triton",
|
||||
"--trust-remote-code",
|
||||
],
|
||||
# AITER MXFP4 fused-MoE for gpt-oss uses the SEPARATED gate/up tile
|
||||
# layout (matches `gptoss_fp4_tuned_fmoe.csv` flydsl entries, the
|
||||
# `QuarkW4A8MXFp4MoE.process_weights_after_loading` shuffle and the
|
||||
# Mxfp4MoEMethod weight shuffle, both uses `shuffle_weight_a16w4`
|
||||
# and + `shuffle_scale_a16w4`) Other AITER MXFP4 callers default
|
||||
# to INTERLEAVE, so opt out explicitly here.
|
||||
env_vars={
|
||||
"SGLANG_USE_AITER": "1",
|
||||
"SGLANG_USE_AITER_MOE_GU_ITLV": "1",
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def get_one_example(lines, i, include_answer):
|
||||
"""Format a single GSM8K example."""
|
||||
ret = "Question: " + lines[i]["question"] + "\nAnswer:"
|
||||
if include_answer:
|
||||
ret += " " + lines[i]["answer"]
|
||||
return ret
|
||||
|
||||
|
||||
def get_few_shot_examples(lines, k):
|
||||
"""Get k few-shot examples for prompting."""
|
||||
ret = ""
|
||||
for i in range(k):
|
||||
ret += get_one_example(lines, i, True) + "\n\n"
|
||||
return ret
|
||||
|
||||
|
||||
def get_answer_value(answer_str):
|
||||
"""Extract numerical answer from response."""
|
||||
answer_str = answer_str.replace(",", "")
|
||||
numbers = re.findall(r"\d+", answer_str)
|
||||
if len(numbers) < 1:
|
||||
return INVALID
|
||||
try:
|
||||
return ast.literal_eval(numbers[-1])
|
||||
except SyntaxError:
|
||||
return INVALID
|
||||
|
||||
|
||||
def run_gsm8k_benchmark(
|
||||
base_url: str,
|
||||
num_questions: int = 200,
|
||||
num_shots: int = 5,
|
||||
parallel: int = 64,
|
||||
) -> Tuple[float, float, float]:
|
||||
"""Run GSM8K few-shot completion benchmark."""
|
||||
import sglang as sgl
|
||||
from sglang.lang.backend.runtime_endpoint import RuntimeEndpoint
|
||||
|
||||
url = "https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl"
|
||||
data_path = download_and_cache_file(url)
|
||||
lines = list(read_jsonl(data_path))
|
||||
|
||||
few_shot_examples = get_few_shot_examples(lines, num_shots)
|
||||
|
||||
questions = []
|
||||
labels = []
|
||||
for i in range(len(lines[:num_questions])):
|
||||
questions.append(get_one_example(lines, i, False))
|
||||
labels.append(get_answer_value(lines[i]["answer"]))
|
||||
assert all(l != INVALID for l in labels)
|
||||
arguments = [{"question": q} for q in questions]
|
||||
|
||||
@sgl.function
|
||||
def few_shot_gsm8k(s, question):
|
||||
s += few_shot_examples + question
|
||||
s += sgl.gen(
|
||||
"answer", max_tokens=512, stop=["Question", "Assistant:", "<|separator|>"]
|
||||
)
|
||||
|
||||
backend = RuntimeEndpoint(base_url)
|
||||
sgl.set_default_backend(backend)
|
||||
|
||||
tic = time.perf_counter()
|
||||
states = few_shot_gsm8k.run_batch(
|
||||
arguments, temperature=0, num_threads=parallel, progress_bar=True
|
||||
)
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
preds = [get_answer_value(states[i]["answer"]) for i in range(len(states))]
|
||||
acc = np.mean(np.array(preds) == np.array(labels))
|
||||
invalid = np.mean(np.array(preds) == INVALID)
|
||||
|
||||
return float(acc), float(invalid), float(latency)
|
||||
|
||||
|
||||
class TestGptOssW4A8Mxfp4EvalMI35x(unittest.TestCase):
|
||||
"""GPT-OSS W4A8 MXFP4-FP8 GSM8K Completion Evaluation Test for AMD MI35x."""
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.models = MI35X_GPT_OSS_W4A8_MXFP4_MODELS
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
cls.num_questions = int(os.environ.get("GSM8K_NUM_QUESTIONS", "200"))
|
||||
|
||||
def test_gpt_oss_accuracy(self):
|
||||
"""Test GPT-OSS W4A8 MXFP4-FP8 models with GSM8K completion benchmark."""
|
||||
all_results = []
|
||||
summary = "### GPT-OSS W4A8 MXFP4-FP8 Models (MI35x)\n\n"
|
||||
summary += "| Model | TP | Accuracy | Threshold | Status |\n"
|
||||
summary += "| ----- | -- | -------- | --------- | ------ |\n"
|
||||
|
||||
for config in self.models:
|
||||
with self.subTest(model=config.model_path):
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Testing: {config.model_path}")
|
||||
print(f"{'='*60}")
|
||||
|
||||
env = os.environ.copy()
|
||||
for key, value in config.env_vars.items():
|
||||
env[key] = value
|
||||
|
||||
other_args = list(config.other_args)
|
||||
other_args.extend(["--tp", str(config.tp_size)])
|
||||
timeout = config.timeout or DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH
|
||||
|
||||
try:
|
||||
process = popen_launch_server(
|
||||
model=config.model_path,
|
||||
base_url=self.base_url,
|
||||
timeout=timeout,
|
||||
other_args=other_args,
|
||||
env=env,
|
||||
)
|
||||
|
||||
try:
|
||||
acc, invalid, latency = run_gsm8k_benchmark(
|
||||
self.base_url, num_questions=self.num_questions
|
||||
)
|
||||
passed = acc >= config.accuracy_threshold
|
||||
status = "✅ PASS" if passed else "❌ FAIL"
|
||||
print(
|
||||
f" accuracy={acc:.3f} threshold={config.accuracy_threshold} {status}"
|
||||
)
|
||||
|
||||
all_results.append(
|
||||
{
|
||||
"model": config.model_path,
|
||||
"accuracy": acc,
|
||||
"passed": passed,
|
||||
}
|
||||
)
|
||||
summary += f"| {config.model_path} | {config.tp_size} | {acc:.3f} | {config.accuracy_threshold} | {status} |\n"
|
||||
|
||||
finally:
|
||||
kill_process_tree(process.pid)
|
||||
|
||||
except Exception as e:
|
||||
summary += f"| {config.model_path} | {config.tp_size} | N/A | {config.accuracy_threshold} | ❌ ERROR |\n"
|
||||
all_results.append(
|
||||
{
|
||||
"model": config.model_path,
|
||||
"accuracy": None,
|
||||
"passed": False,
|
||||
"error": str(e),
|
||||
}
|
||||
)
|
||||
|
||||
if is_in_ci():
|
||||
write_github_step_summary(summary)
|
||||
|
||||
failed = [r for r in all_results if not r["passed"]]
|
||||
if failed:
|
||||
raise AssertionError(f"Failed models: {[r['model'] for r in failed]}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user