[AMD][Quantization] Online MXFP4 quantization 4/N - NVFP4 to MXFP4 Online Requantization on AMD GPUs (#29328)

This commit is contained in:
Colin Z
2026-08-14 21:59:39 -07:00
committed by GitHub
parent 5afdb1caea
commit bc7e3ba66c
14 changed files with 1218 additions and 204 deletions
+19 -2
View File
@@ -865,7 +865,7 @@ Other layers (e.g. projections in the attention layers) have their weights quant
### `quark_mxfp4` online quantization method
SGLang running on AMD GPUs with hardware FP4 support (CDNA4 architecture, e.g. MI355x) supports the quantization method `--quantization quark_mxfp4`, that will quantize BF16 model weights to MXFP4 at load time, use dynamic MXFP4 quantization for activations and MXFP4 GEMMs instead of BF16 GEMMs.
SGLang running on AMD GPUs with hardware FP4 support (CDNA4 architecture, e.g. MI355x) supports the quantization method `--quantization quark_mxfp4`, that will quantize BF16 or NVFP4 model weights to MXFP4 at load time, use dynamic MXFP4 quantization for activations and MXFP4 GEMMs instead of BF16 GEMMs.
Example (BF16 to MXFP4 requantization):
@@ -875,7 +875,24 @@ sglang serve --model-path Qwen/Qwen3-30B-A3B \
--quantization quark_mxfp4
```
The option `--quantization quark_mxfp4` also supports converting FP8 dense and MOE models to MXFP4, following this logic:
#### Online NVFP4 to MXFP4 Requantization
The option `--quantization quark_mxfp4` supports converting NVFP4 checkpoints (e.g. `nvidia/Kimi-K2.6-NVFP4`) to MXFP4 at load time to allow efficient inference using supported AMD hardware (gfx95x+):
- The quantization metadata of the source NVFP4 checkpoint is read from either `config.json` (`quantization_config`) or a standalone `hf_quant_config.json`;
- Producer-declared excluded modules will remain in higher precision;
- NVFP4 checkpoints with mixed precision (`"quant_algo": "MIXED_PRECISION"`, e.g. `nvidia/Qwen3.5-397B-A17B-NVFP4-V2`) are also supported.
Example (NVFP4 to MXFP4 requantization):
```bash
sglang serve --model-path nvidia/Kimi-K2.6-NVFP4 \
--tensor-parallel-size 4 \
--quantization quark_mxfp4 \
```
#### Online FP8 to MXFP4 Requantization
The option `--quantization quark_mxfp4` supports converting FP8 dense and MOE models to MXFP4, following this logic:
1. Load an FP8 weight tensor,
2. Dequantize it to BF16,
+15 -10
View File
@@ -1507,16 +1507,21 @@ class ModelConfig:
and self.quantization == "nvfp4_online"
and quant_method == "modelopt_fp4"
)
# Detect which checkpoint is it
if not preserve_online_draft_quantization:
for _, method in QUANTIZATION_METHODS.items():
quantization_override = method.override_quantization_method(
quant_cfg, self.quantization
)
if quantization_override:
quant_method = quantization_override
self.quantization = quantization_override
break
# An explicit online-requantization request (e.g. quark_mxfp4 on top
# of an NVFP4/mixed checkpoint) must not be overridden back to the
# source format
if self.quantization not in REQUANTIZATION_METHODS:
# Detect which checkpoint is it
if not preserve_online_draft_quantization:
for _, method in QUANTIZATION_METHODS.items():
quantization_override = method.override_quantization_method(
quant_cfg, self.quantization
)
if quantization_override:
quant_method = quantization_override
self.quantization = quantization_override
break
# Verify quantization configurations.
if self.quantization is None:
+18 -2
View File
@@ -155,6 +155,10 @@ class LinearBase(torch.nn.Module):
quant_config: Quantization configure.
"""
# Set by quant methods that attach a per-layer scheme (e.g. Quark) inside
# get_quant_method(), which runs before create_weights() picks the loader.
scheme = None
def __init__(
self,
input_size: int,
@@ -366,7 +370,13 @@ class ColumnParallelLinear(LinearBase):
skip_block_quant_check=skip_block_quant_check,
weight_loader=(
self.weight_loader_v2
if self.quant_method.__class__.__name__ in WEIGHT_LOADER_V2_SUPPORTED
if (
self.quant_method.__class__.__name__ in WEIGHT_LOADER_V2_SUPPORTED
or (
self.scheme is not None
and self.scheme.requires_weight_loader_v2
)
)
else self.weight_loader
),
)
@@ -1462,7 +1472,13 @@ class RowParallelLinear(LinearBase):
params_dtype=self.params_dtype,
weight_loader=(
self.weight_loader_v2
if self.quant_method.__class__.__name__ in WEIGHT_LOADER_V2_SUPPORTED
if (
self.quant_method.__class__.__name__ in WEIGHT_LOADER_V2_SUPPORTED
or (
self.scheme is not None
and self.scheme.requires_weight_loader_v2
)
)
else self.weight_loader
),
)
@@ -185,7 +185,12 @@ class QuantizationConfig(ABC):
if hf_quant_config is None:
return None
if user_quant == "nvfp4_online":
# If the user explicitly requested an online requantization (e.g.
# quark_mxfp4 on top of an NVFP4 checkpoint), do not override it back
# to the source format.
from sglang.srt.configs.model_config import REQUANTIZATION_METHODS
if user_quant == "nvfp4_online" or user_quant in REQUANTIZATION_METHODS:
return None
# Check if this is a ModelOpt config
@@ -14,6 +14,12 @@ class BaseLinearScheme(ABC):
of different quantization schemes.
"""
# Schemes whose parameters only implement the v2 loader API
# (load_{column,row,merged_column,qkv}_weight) set this so LinearBase
# routes them through weight_loader_v2 without flipping the loader for
# every scheme that shares the same LinearMethod class.
requires_weight_loader_v2: bool = False
@abstractmethod
def create_weights(self, *args, **kwargs):
"""
@@ -2,6 +2,8 @@
Utilities to manage the dequantization of weights.
"""
from typing import Optional
import torch
from sglang.srt.layers.quantization.fp8_utils import (
@@ -10,6 +12,29 @@ from sglang.srt.layers.quantization.fp8_utils import (
)
from sglang.srt.utils import set_weight_attrs
NVFP4_BLOCK_SIZE = 16
_FP4_E2M1_LUT = torch.tensor(
[
0.0,
0.5,
1.0,
1.5,
2.0,
3.0,
4.0,
6.0,
-0.0,
-0.5,
-1.0,
-1.5,
-2.0,
-3.0,
-4.0,
-6.0,
],
dtype=torch.float32,
)
def copy_missing_attrs(old: torch.Tensor, new: torch.Tensor) -> None:
"""Copies any attrs present in `old` but not in `new` to `new`"""
@@ -42,3 +67,31 @@ def dequantize_fp8(
)
return w_dequant
def dequantize_nvfp4(
w_q: torch.Tensor,
w_s: torch.Tensor,
w_s2: Optional[torch.Tensor],
out_dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
"""NVFP4 -> ``out_dtype``. ``w_q``: uint8 [..., out, in/2] packed e2m1
(low nibble = even idx). ``w_s``: fp8 e4m3 [..., out, in/16] per-block.
``w_s2``: optional fp32 per-tensor scalar that multiplies the per-block
scale (ModelOpt / AMD Quark NVFP4)."""
device = w_q.device
*batch, out_dim, half_in = w_q.shape
in_dim = half_in * 2
low = (w_q & 0xF).to(torch.int64)
high = (w_q >> 4).to(torch.int64)
lut = _FP4_E2M1_LUT.to(device=device, dtype=torch.float32)
deq = torch.empty(*batch, out_dim, in_dim, dtype=torch.float32, device=device)
deq[..., 0::2] = lut[low]
deq[..., 1::2] = lut[high]
scale = w_s.to(torch.float32)
if w_s2 is not None:
scale = scale * w_s2.to(torch.float32)
scale = scale.repeat_interleave(NVFP4_BLOCK_SIZE, dim=-1)
return (deq * scale).to(out_dtype)
@@ -2,7 +2,8 @@
import fnmatch
import logging
from typing import TYPE_CHECKING, Any, List, Optional, cast
import re
from typing import TYPE_CHECKING, Any, Dict, List, Optional, cast
import torch
@@ -25,7 +26,11 @@ from sglang.srt.layers.quantization.quark.schemes import (
QuarkW8A8Fp8,
QuarkW8A8FP8MoE,
)
from sglang.srt.layers.quantization.quark.utils import deep_compare, should_ignore_layer
from sglang.srt.layers.quantization.quark.utils import (
Nvfp4SourceConfig,
deep_compare,
should_ignore_layer,
)
from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.utils import get_device_capability
@@ -37,6 +42,235 @@ if TYPE_CHECKING:
__all__ = ["QuarkLinearMethod", "QuarkFusedMoEMethod"]
def _parse_nvfp4_excludes(hf_quant_config: Dict[str, Any]) -> List[str]:
"""Extract NVFP4 producer-declared excludes as `re:` patterns.
Reads the producer-specific key:
- `ignore` - ModelOpt (config.json)
- `exclude_modules` - ModelOpt hf_quant_config.json
- `exclude` - AMD Quark export
Entries are usually fnmatch-style (literal strings work too), but ModelOpt
`ignore` lists may already carry `re:`-prefixed regexes (e.g.
`re:.*linear_attn\\.in_proj_a$`); those are passed through untouched.
Wrapping an already-`re:` entry with another `re:` + `fnmatch.translate`
yields a pattern that never matches, silently un-excluding the layer.
Returns [] if no key present.
"""
pats = (
hf_quant_config.get("ignore")
or hf_quant_config.get("exclude_modules")
or hf_quant_config.get("exclude")
or []
)
return [p if p.startswith("re:") else "re:" + fnmatch.translate(p) for p in pats]
def _detect_nvfp4_source(config: Dict[str, Any]) -> Optional["Nvfp4SourceConfig"]:
"""Return an Nvfp4SourceConfig if `config` (the checkpoint's
quantization_config dict) describes a supported NVFP4 source, else None.
Handles two producers:
- ModelOpt: quant_method in {modelopt, modelopt_fp4, nvfp4}
with quant_algo NVFP4/FP4 (or unspecified).
- AMD Quark: quant_method == "quark". global_quant_config.weight is a
2-element list [fp4_per_group_gs16, fp8_e4m3_per_tensor].
compressed-tensors NVFP4 is not supported at this time.
"""
from sglang.srt.layers.quantization.quark.utils import Nvfp4SourceConfig
quant_method = config.get("quant_method", "")
quant_algo = (config.get("quant_algo") or "").upper()
if quant_method in ("modelopt", "modelopt_fp4", "nvfp4") and quant_algo in (
"",
"NVFP4",
"FP4",
):
return Nvfp4SourceConfig()
if quant_method == "quark":
gqc = config.get("global_quant_config", {})
weight = gqc.get("weight")
if not (isinstance(weight, list) and len(weight) == 2):
return None
w0, w1 = weight
is_nvfp4_weight = (
isinstance(w0, dict)
and w0.get("dtype") == "fp4"
and w0.get("qscheme") == "per_group"
and w0.get("group_size") == 16
and not w0.get("is_dynamic")
)
is_nvfp4_scale_2 = (
isinstance(w1, dict)
and w1.get("dtype") == "fp8_e4m3"
and w1.get("qscheme") == "per_tensor"
and not w1.get("is_dynamic")
)
if is_nvfp4_weight and is_nvfp4_scale_2:
return Nvfp4SourceConfig()
return None
if quant_method in ("compressed-tensors", "compressed_tensors"):
raise NotImplementedError(
"Online MXFP4 requantization from compressed-tensors NVFP4 "
"checkpoints is not supported at this time."
)
return None
# Target quant specs used when synthesizing a per-layer config for a
# MIXED_PRECISION source. The MXFP4 spec is the online-requant target shape
# recognized by `_is_mx_fp4`; the FP8 spec is the per-tensor W8A8 shape
# recognized by `_is_fp8_w8a8` (no requantization).
_MXFP4_TARGET_SPEC: Dict[str, Any] = {
"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",
},
"output_tensors": None,
"bias": None,
}
def _fp8_per_tensor_spec(is_dynamic_input: bool) -> Dict[str, Any]:
return {
"weight": {
"dtype": "fp8_e4m3",
"qscheme": "per_tensor",
"is_dynamic": False,
},
"input_tensors": {
"dtype": "fp8_e4m3",
"qscheme": "per_tensor",
"is_dynamic": is_dynamic_input,
},
"output_tensors": None,
"bias": None,
}
def _fp8_is_dynamic_from_config_groups(
config_groups: Any,
) -> bool:
"""Return whether FP8 activation quantization is dynamic, from config_groups.
Reads the `input_activations.dynamic` field of the first config_group whose
`num_bits` is 8, and falls back to True (dynamic) when none exists or the
format is not a recognised dict-of-dicts.
"""
if not isinstance(config_groups, dict):
return True
for group in config_groups.values():
if not isinstance(group, dict):
continue
input_act = group.get("input_activations") or {}
if input_act.get("num_bits") == 8:
return bool(input_act.get("dynamic", True))
return True
def _mixed_precision_layer_map(config: Dict[str, Any]) -> Optional[Dict[str, str]]:
"""Return {layer_name: quant_algo} for a MIXED_PRECISION source, else None.
Reads ModelOpt's per-layer `quantized_layers` map (from
hf_quant_config.json or config.json's quantization_config). Only the
quant_algo string per layer is needed;
"""
if (config.get("quant_algo") or "").upper() != "MIXED_PRECISION":
return None
quantized_layers = config.get("quantized_layers")
if not isinstance(quantized_layers, dict) or not quantized_layers:
return None
layer_map: Dict[str, str] = {}
for name, info in quantized_layers.items():
if isinstance(info, dict):
layer_map[name] = str(info.get("quant_algo", "")).upper()
return layer_map
def _build_mixed_precision_layer_quant_config(
layer_map: Dict[str, str],
config_groups: Optional[Dict[str, Any]] = None,
) -> tuple[Dict[str, Any], bool]:
"""Collapse a per-layer {name: quant_algo} map into a compact
`layer_quant_config` keyed by fnmatch glob patterns.
"""
# suffix tail -> set of algos seen (to detect inconsistency)
tail_algos: Dict[str, set] = {}
for name, algo in layer_map.items():
# Suffix after the last `.layers.<idx>.` (or the whole name if
# unindexed); this is the part shared across all layer indices.
tail = re.split(r"\.layers\.\d+\.", name, maxsplit=1)[-1]
tail_algos.setdefault(tail, set()).add(algo)
fp8_is_dynamic = _fp8_is_dynamic_from_config_groups(config_groups or {})
fp8_spec = _fp8_per_tensor_spec(is_dynamic_input=fp8_is_dynamic)
layer_quant_config: Dict[str, Any] = {}
has_nvfp4 = False
for tail, algos in tail_algos.items():
if len(algos) != 1:
raise NotImplementedError(
f"MIXED_PRECISION layer group {tail!r} has inconsistent "
f"quant algos across layers: {sorted(algos)}. SGLang requires "
"all layers in a group to share one algo."
)
algo = next(iter(algos))
pattern = "*" + tail
if algo in ("NVFP4", "W4A16_NVFP4"):
layer_quant_config[pattern] = _MXFP4_TARGET_SPEC
has_nvfp4 = True
elif algo == "FP8":
layer_quant_config[pattern] = fp8_spec
else:
raise NotImplementedError(
f"MIXED_PRECISION layer group {tail!r} uses unsupported "
f"quant algo {algo!r}; online requantization supports NVFP4 "
"(-> MXFP4) and FP8 (kept as-is) only."
)
return layer_quant_config, has_nvfp4
def _build_excluded_fp8_config(config: Dict[str, Any]) -> Optional["Fp8Config"]:
"""Build a load-as-is `Fp8Config` for the excluded layers of a
mixed-precision NVFP4 source, or None if excluded layers are bf16.
Two producer conventions are handled:
- FP8-serialized base (``quant_method == "fp8"``, e.g.
DeepSeek-V4-Pro-NVFP4): the routed experts are NVFP4 (requantized to
MXFP4) while attn / shared_experts stay FP8 and are listed in the
excludes. Those FP8 layers load through `Fp8LinearMethod`;
``weight_block_size`` selects block (e.g. ``[128, 128]``) vs per-tensor
(``None``), so a single config covers either granularity - and a
checkpoint carrying only per-tensor or only block layers is handled
without any per-layer probing.
- ModelOpt mixed base (``quant_method`` in {modelopt, modelopt_mixed},
e.g. Qwen3.5-397B-A17B-NVFP4-V2): FP8 layers are enumerated in the
per-layer ``quantized_layers`` map (loaded via `QuarkW8A8Fp8`), not in
the excludes, so the excludes are genuinely bf16 -> None.
"""
if config.get("quant_method") != "fp8":
return None
# Fp8Config.from_config reads quant_method/activation_scheme/
# weight_block_size/packed_modules_mapping straight off the checkpoint's
# quantization_config dict, which is exactly what `config` carries here.
return Fp8Config.from_config(config)
logger = logging.getLogger(__name__)
_MOE_SHARED_EXPERT_QUANT_LAYER0_BASES: tuple[str, ...] = (
@@ -64,6 +298,7 @@ class QuarkConfig(QuantizationConfig):
is_prequantized: bool = False,
online_scheme: Optional[str] = None,
dequantization_config: Optional[QuantizationConfig] = None,
excluded_fp8_config: Optional[Fp8Config] = None,
):
super().__init__()
if kv_cache_group is None:
@@ -89,31 +324,49 @@ class QuarkConfig(QuantizationConfig):
self.exclude_layers = cast(list[str], self.quant_config.get("exclude", []))
self.is_prequantized = is_prequantized
self.dequantization_config = dequantization_config
# Load-as-is FP8 config for excluded layers of a mixed-precision source
# (e.g. attn / shared_experts kept in FP8 while routed experts are
# requantized NVFP4 -> MXFP4). Distinct from `dequantization_config`,
# which describes the requantization *source*. `weight_block_size`
# selects block vs per-tensor FP8 within `Fp8LinearMethod`.
self.excluded_fp8_config = excluded_fp8_config
self.packed_modules_mapping = self.quant_config["packed_modules_mapping"]
self._online_quantized_layers = set()
if isinstance(self.dequantization_config, Fp8Config):
self.weight_block_size = self.dequantization_config.weight_block_size
def log_online_quantization(self) -> None:
"""
Log which layers are using online quantization, as well as a count for each layer type.
"""
# Count layers per type (last two parts after ".")
type_counts: dict[str, int] = {}
for name in self._online_quantized_layers:
parts = name.split(".")
layer_type = ".".join(parts[-2:]) if len(parts) >= 2 else parts[-1]
type_counts[layer_type] = type_counts.get(layer_type, 0) + 1
self._maybe_disable_shared_experts_fusion()
type_counts = dict(sorted(type_counts.items()))
count = len(self._online_quantized_layers)
def _maybe_disable_shared_experts_fusion(self) -> None:
"""Turn off shared-expert fusion when the producer keeps shared experts
in a higher precision than the routed experts.
"""
if self.can_fuse_shared_expert():
return
type_summary = ", ".join(f"{t}: {c}" for t, c in type_counts.items())
logger.info_once(
f"Online {self.online_scheme} quantization: "
f"quantized {count} layers in total ({type_summary})."
from sglang.srt.arg_groups.overrides import declare_load_time_override
declare_load_time_override(
"QuarkConfig._maybe_disable_shared_experts_fusion",
{"disable_shared_experts_fusion": True},
)
logger.info(
"Quark: shared experts are excluded from quantization (kept in "
"a higher precision) while routed experts are quantized; "
"disabling shared experts fusion to avoid loading "
"higher-precision shared experts through the quantized "
"routed-expert path."
)
@property
def quantized_layers(self) -> tuple[list[str], int]:
# Consumed by `report_online_quantization` in model_runner. Returns the
# unique layer types (last part after ".") and the total layer count.
layer_types = sorted(
set(name.split(".")[-1] for name in self._online_quantized_layers)
)
return layer_types, len(self._online_quantized_layers)
def get_linear_method(self) -> "QuarkLinearMethod":
return QuarkLinearMethod(self)
@@ -148,12 +401,13 @@ class QuarkConfig(QuantizationConfig):
fused_mapping=self.packed_modules_mapping,
):
if isinstance(layer, LinearBase):
if self.dequantization_config is not None:
# In case of online requantization, "exclude" means keeping the original precision.
# NOTE: Only FP8 supported for now.
return Fp8LinearMethod(quant_config=self.dequantization_config)
else:
return UnquantizedLinearMethod()
# "exclude" means keep the layer in its original precision.
# Mixed-precision sources may keep excluded layers in FP8
# (block or per-tensor, selected by excluded_fp8_config's
# weight_block_size); pure-NVFP4/BF16 sources keep them bf16.
if self.excluded_fp8_config is not None:
return Fp8LinearMethod(quant_config=self.excluded_fp8_config)
return UnquantizedLinearMethod()
elif isinstance(layer, RadixAttention):
return QuarkKVCacheMethod(self)
return None
@@ -179,32 +433,90 @@ class QuarkConfig(QuantizationConfig):
@classmethod
def from_config(cls, config: dict[str, Any]) -> "QuarkConfig":
if config["quant_method"] != "quark":
assert "requantization_method" in config
# Requantization dispatch is gated on requantization_method, NOT on
# quant_method. Quark-exported NVFP4 carries quant_method="quark" too
if config.get("requantization_method") == "quark_mxfp4":
hf_config = config["hf_config"]
# Mixed-precision source: only the NVFP4 layers are requantized to
# MXFP4; layers in other precisions (e.g. FP8) load through their
# own scheme
layer_map = _mixed_precision_layer_map(config)
if layer_map is not None:
config_groups = config.get("config_groups")
layer_quant_config, has_nvfp4 = (
_build_mixed_precision_layer_quant_config(layer_map, config_groups)
)
if not has_nvfp4:
raise NotImplementedError(
"MIXED_PRECISION checkpoint has no NVFP4 layers to "
"requantize; load it with its native quantization "
"method instead of --quantization quark_mxfp4."
)
source_excludes = _parse_nvfp4_excludes(config)
quant_config = QuarkConfig._create_online_mxfp4_config(
model_type=hf_config.model_type,
source_excludes=source_excludes,
layer_quant_config=layer_quant_config,
packed_modules_mapping=config.get("packed_modules_mapping"),
)
# Excluded layers are kept as-is. When the base checkpoint is
# FP8-serialized (e.g. DeepSeek-V4-Pro-NVFP4: FP8 attn/
# shared_experts, NVFP4 routed experts) they load through FP8;
# `weight_block_size` selects block vs per-tensor. Pure
# NVFP4/ModelOpt-mixed sources keep excluded layers in bf16, and
# their FP8 layers (if any) are enumerated in the layer map.
excluded_fp8_config = _build_excluded_fp8_config(config)
return cls(
quant_config=quant_config,
hf_config=hf_config,
is_prequantized=False,
dequantization_config=Nvfp4SourceConfig(),
excluded_fp8_config=excluded_fp8_config,
)
nvfp4_src = _detect_nvfp4_source(config)
if nvfp4_src is not None:
source_excludes = _parse_nvfp4_excludes(config)
quant_config = QuarkConfig._create_online_mxfp4_config(
model_type=hf_config.model_type,
source_excludes=source_excludes,
)
return cls(
quant_config=quant_config,
hf_config=hf_config,
is_prequantized=False,
dequantization_config=nvfp4_src,
)
# Pure FP8 source: every layer is requantized FP8 -> MXFP4.
if (
config["quant_method"] == "fp8"
and config["requantization_method"] == "quark_mxfp4"
and config["activation_scheme"] == "dynamic"
config.get("quant_method") == "fp8"
and config.get("activation_scheme") == "dynamic"
):
hf_config = config["hf_config"]
quant_config = QuarkConfig._create_online_mxfp4_config(
model_type=hf_config.model_type
)
dequantization_config = Fp8Config.from_config(config)
quark_config = cls(
return cls(
quant_config=quant_config,
hf_config=hf_config,
is_prequantized=False,
dequantization_config=dequantization_config,
online_scheme=config["requantization_method"],
)
else:
raise NotImplementedError(
f"Requantization into {config['requantization_method']} is not supported, from the original quant_method={config['quant_method']} and activation_scheme={config['activation_scheme']}. "
)
return quark_config
raise NotImplementedError(
f"Requantization into {config['requantization_method']} is not supported, "
f"from the original quant_method={config['quant_method']} "
f"and activation_scheme={config.get('activation_scheme')}."
)
if config["quant_method"] != "quark":
raise ValueError(
f"QuarkConfig.from_config invoked with non-quark quant_method "
f"{config['quant_method']!r} but no requantization_method set."
)
export_config = config.get("export")
if export_config is None:
@@ -277,9 +589,19 @@ class QuarkConfig(QuantizationConfig):
return []
@staticmethod
def _create_online_mxfp4_config(model_type: str) -> dict[str, Any]:
def _create_online_mxfp4_config(
model_type: str,
source_excludes: Optional[list[str]] = None,
layer_quant_config: Optional[dict[str, Any]] = None,
packed_modules_mapping: Optional[dict[str, list[str]]] = None,
) -> dict[str, Any]:
"""
Create a synthetic quant_config for online MXFP4 quantization.
When `layer_quant_config` is provided (mixed-precision source), the
per-layer map is authoritative about which layers are quantized and in
what precision, so the model_type-specific default excludes
are skipped: non-NVFP4 layers must load through their own scheme
"""
# MOE gate/router is typically implemented as a ReplicatedLinear, and skipped for quantization for accuracy reasons.
# lm_head/embed_tokens is also skipped for accuracy reasons, normally not handled by `QuarkConfig` in any case, but adding them here for safety.
@@ -290,34 +612,37 @@ class QuarkConfig(QuantizationConfig):
"re:.*embed_tokens",
]
# Exclusion for accuracy adapted from
# https://huggingface.co/amd/DeepSeek-V3.2-mxfp4/blob/main/config.json
if model_type in ["deepseek_v3", "deepseek_v32"]:
exclude.extend(
[
"re:.*model.layers.61.*",
"re:.*self_attn.*",
"re:.*mlp.gate$",
]
)
elif model_type == "qwen3_5_moe":
if source_excludes:
exclude.extend(source_excludes)
elif layer_quant_config is None:
# Exclusion for accuracy adapted from
# https://huggingface.co/amd/Qwen3.5-397B-A17B-MXFP4/blob/main/config.json
exclude.extend(
[
"re:.*n_proj_a",
"re:.*in_proj_b",
"re:.*in_proj_qkv",
"re:.*in_proj_z",
"re:.*o_proj",
"re:.*out_proj",
"re:.*qkv_proj",
"re:.*shared_expert",
]
)
# https://huggingface.co/amd/DeepSeek-V3.2-mxfp4/blob/main/config.json
if model_type in ("deepseek_v3", "deepseek_v32", "deepseek_v4"):
exclude.extend(
[
"re:.*model.layers.61.*",
"re:.*self_attn.*",
"re:.*mlp.gate$",
]
)
elif model_type == "qwen3_5_moe":
# Exclusion for accuracy adapted from
# https://huggingface.co/amd/Qwen3.5-397B-A17B-MXFP4/blob/main/config.json
exclude.extend(
[
"re:.*n_proj_a",
"re:.*in_proj_b",
"re:.*in_proj_qkv",
"re:.*in_proj_z",
"re:.*o_proj",
"re:.*out_proj",
"re:.*qkv_proj",
"re:.*shared_expert",
]
)
return {
"packed_modules_mapping": {},
"packed_modules_mapping": packed_modules_mapping or {},
"exclude": exclude,
"global_quant_config": {
"weight": {
@@ -337,7 +662,7 @@ class QuarkConfig(QuantizationConfig):
"output_tensors": None,
"bias": None,
},
"layer_quant_config": {},
"layer_quant_config": layer_quant_config or {},
"layer_type_quant_config": {},
"export": {
"kv_cache_group": [],
@@ -635,9 +960,6 @@ class QuarkLinearMethod(LinearMethodBase):
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.scheme.process_weights_after_loading(layer)
if self.quantization_config.online_scheme is not None:
self.quantization_config.log_online_quantization()
def create_weights(
self,
layer: torch.nn.Module,
@@ -690,9 +1012,6 @@ class QuarkFusedMoEMethod(FusedMoEMethodBase):
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.scheme.process_weights_after_loading(layer)
if self.quantization_config.online_scheme is not None:
self.quantization_config.log_online_quantization()
def create_weights(
self,
layer: torch.nn.Module,
@@ -6,18 +6,27 @@ from typing import Any, Callable, Optional
import torch
from sglang.srt.layers.parameter import GroupQuantScaleParameter, PackedvLLMParameter
from sglang.srt.layers.parameter import (
GroupQuantScaleParameter,
ModelWeightParameter,
PackedvLLMParameter,
PerTensorScaleParameter,
)
from sglang.srt.layers.quantization import QuantizationConfig
from sglang.srt.layers.quantization.dequantization import (
copy_missing_attrs,
dequantize_fp8,
dequantize_nvfp4,
)
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8LinearMethod
from sglang.srt.layers.quantization.online_quantization import CopyNumelCounter
from sglang.srt.layers.quantization.quark.schemes import QuarkLinearScheme
from sglang.srt.layers.quantization.quark.utils import Nvfp4SourceConfig
from sglang.srt.utils import is_hip
from sglang.srt.utils.common import direct_register_custom_op, is_gfx95_supported
NVFP4_BLOCK_SIZE = 16
_is_hip = is_hip()
if _is_hip:
from aiter.ops.triton.gemm.fused.fused_gemm_afp4wfp4_split_cat import (
@@ -165,6 +174,10 @@ OCP_MX_BLOCK_SIZE = 32
class QuarkW4A4MXFP4(QuarkLinearScheme):
# PackedvLLMParameter / ModelWeightParameter (online and NVFP4->MXFP4
# paths) only implement the v2 loader API.
requires_weight_loader_v2 = True
def __init__(
self,
weight_quant_spec: dict[str, Any],
@@ -214,61 +227,74 @@ class QuarkW4A4MXFP4(QuarkLinearScheme):
layer.logical_widths = output_partition_sizes
# If dequantization_config is provided, we need to create FP8 weights first
# for dequantization from FP8 checkpoint to MXFP4
# If dequantization_config is provided, we dequantize the source
# checkpoint and re-quantize to MXFP4 at load time. The source may be
# NVFP4 (ModelOpt/Quark) or FP8 (block-quantized); each has its own
# weight-creation and loader path.
if self.dequantization_config is not None:
if not isinstance(self.dequantization_config, Fp8Config):
raise NotImplementedError(
f"Requantization in QuarkW4A4MXFP4 from {self.dequantization_config.__class__.__name__} is not supported, only Fp8Config is supported."
)
# Create FP8 weights for re-quantization from FP8 checkpoint
# Extract necessary parameters from dequantization_config
self.weight_block_size = self.dequantization_config.weight_block_size
if self.dequantization_config.use_mxfp8:
raise NotImplementedError(
"use_mxfp8=True is not supported in Quark MXFP4 requantization."
)
block_quant = self.weight_block_size is not None
if not block_quant:
raise NotImplementedError(
"Only block_quant=True is supported in Quark MXFP4 requantization, got block_quant=False."
)
layer._fp8_weight_loaded_numel = 0
layer._load_device = torch.get_default_device()
layer._fp8_weight_loading_lock = threading.Lock()
layer._fp8_weight_materialized = False
# Wrap the weight loader to handle FP8->MXFP4 conversion
fp8_to_mxfp4_weight_loader = self.get_online_fp8_to_mxfp4_weight_loader(
layer, weight_loader
)
# Create FP8 MoE weight parameters on meta device to avoid device memory overhead during weight loading, as the resulting model uses MXFP4 using less device memory.
# The weight loader handles progressive FP8 weight materialization on device.
with torch.device("meta"):
Fp8LinearMethod.create_fp8_weight_(
if isinstance(self.dequantization_config, Nvfp4SourceConfig):
self._create_weights_from_nvfp4(
layer=layer,
block_quant=block_quant,
quant_config=self.dequantization_config,
use_mxfp8=False,
output_size_per_partition=output_size_per_partition,
input_size_per_partition=input_size_per_partition,
output_partition_sizes=output_partition_sizes,
weight_loader=fp8_to_mxfp4_weight_loader,
is_checkpoint_fp8_serialized=True,
params_dtype=params_dtype,
skip_block_quant_check=False,
input_size=kwargs.get("input_size", input_size_per_partition),
output_size=kwargs.get("output_size", output_size_per_partition),
weight_loader=weight_loader,
)
elif isinstance(self.dequantization_config, Fp8Config):
# Create FP8 weights for re-quantization from FP8 checkpoint.
# Extract necessary parameters from dequantization_config.
self.weight_block_size = self.dequantization_config.weight_block_size
if self.dequantization_config.use_mxfp8:
raise NotImplementedError(
"use_mxfp8=True is not supported in Quark MXFP4 requantization."
)
block_quant = self.weight_block_size is not None
if not block_quant:
raise NotImplementedError(
"Only block_quant=True is supported in Quark MXFP4 requantization, got block_quant=False."
)
layer._fp8_weight_loaded_numel = 0
layer._load_device = torch.get_default_device()
layer._fp8_weight_loading_lock = threading.Lock()
layer._fp8_weight_materialized = False
# Wrap the weight loader to handle FP8->MXFP4 conversion
fp8_to_mxfp4_weight_loader = self.get_online_fp8_to_mxfp4_weight_loader(
layer, weight_loader
)
# NOTE: ideally, weight_loader should be refactored to be aware of `param_name`.
layer.weight._param_name = "weight"
layer.weight_scale_inv._param_name = "weight_scale_inv"
# Create FP8 weight parameters on meta device to avoid device memory overhead during weight loading, as the resulting model uses MXFP4 using less device memory.
# The weight loader handles progressive FP8 weight materialization on device.
with torch.device("meta"):
Fp8LinearMethod.create_fp8_weight_(
layer=layer,
block_quant=block_quant,
quant_config=self.dequantization_config,
use_mxfp8=False,
output_size_per_partition=output_size_per_partition,
input_size_per_partition=input_size_per_partition,
output_partition_sizes=output_partition_sizes,
weight_loader=fp8_to_mxfp4_weight_loader,
is_checkpoint_fp8_serialized=True,
params_dtype=params_dtype,
skip_block_quant_check=False,
input_size=kwargs.get("input_size", input_size_per_partition),
output_size=kwargs.get(
"output_size", output_size_per_partition
),
)
# NOTE: ideally, weight_loader should be refactored to be aware of `param_name`.
layer.weight._param_name = "weight"
layer.weight_scale_inv._param_name = "weight_scale_inv"
else:
raise NotImplementedError(
f"Requantization in QuarkW4A4MXFP4 from {self.dequantization_config.__class__.__name__} is not supported."
)
else:
original_weight_loader = weight_loader
if not self.is_checkpoint_mxfp4_serialized:
@@ -305,6 +331,143 @@ class QuarkW4A4MXFP4(QuarkLinearScheme):
)
layer.register_parameter("weight_scale", weight_scale)
def _create_weights_from_nvfp4(
self,
layer,
output_size_per_partition,
input_size_per_partition,
output_partition_sizes,
weight_loader,
):
layer._nvfp4_loaded_numel = 0
# torch.get_default_device() may return `cuda` (no index), which breaks
# the `current_device() == idx` assert in the loader
layer._load_device = torch.device(f"cuda:{torch.cuda.current_device()}")
layer._nvfp4_loading_lock = threading.Lock()
nvfp4_loader = self.get_online_nvfp4_to_mxfp4_weight_loader(
layer, weight_loader
)
layer.register_parameter(
"weight",
ModelWeightParameter(
data=torch.empty(
output_size_per_partition,
input_size_per_partition // 2,
dtype=torch.uint8,
device=layer._load_device,
),
input_dim=1,
output_dim=0,
weight_loader=nvfp4_loader,
),
)
layer.register_parameter(
"weight_scale",
ModelWeightParameter(
data=torch.empty(
output_size_per_partition,
input_size_per_partition // NVFP4_BLOCK_SIZE,
dtype=torch.float8_e4m3fn,
device=layer._load_device,
),
input_dim=1,
output_dim=0,
weight_loader=nvfp4_loader,
),
)
layer.register_parameter(
"weight_scale_2",
PerTensorScaleParameter(
data=torch.empty(
len(output_partition_sizes),
dtype=torch.float32,
device=layer._load_device,
),
weight_loader=nvfp4_loader,
),
)
# NVFP4 checkpoints carry per-tensor `input_scale` (activation scale).
# MXFP4 uses dynamic activation quantization, so we discard it, but
# we still register the param so upstream model loaders that rename
# `gate_proj.input_scale` -> `gate_up_proj.input_scale` find a slot
# to write into
def _discard_loader(param, loaded_weight, shard_id=None):
pass
layer.register_parameter(
"input_scale",
PerTensorScaleParameter(
data=torch.empty(
len(output_partition_sizes),
dtype=torch.float32,
device=layer._load_device,
),
weight_loader=_discard_loader,
),
)
layer.weight._param_name = "weight"
layer.weight_scale._param_name = "weight_scale"
layer.weight_scale_2._param_name = "weight_scale_2"
def get_online_nvfp4_to_mxfp4_weight_loader(
self,
layer,
original_weight_loader: Callable,
) -> Callable:
"""NVFP4 -> MXFP4 loader: dequantize+requantize once all source bytes
are in place."""
def loader(param, loaded_weight, shard_id=None):
param_name = getattr(param, "_param_name", None)
assert torch.cuda.current_device() == layer._load_device.index
with layer._nvfp4_loading_lock:
param = getattr(layer, param_name)
kwargs = {"loaded_shard_id": shard_id} if shard_id is not None else {}
counter = CopyNumelCounter()
with counter:
original_weight_loader(param, loaded_weight, **kwargs)
with layer._nvfp4_loading_lock:
layer._nvfp4_loaded_numel += counter.copied_numel
target = (
layer.weight.numel()
+ layer.weight_scale.numel()
+ layer.weight_scale_2.numel()
)
if layer._nvfp4_loaded_numel == target:
# weight_scale_2 is one fp32 per output partition (e.g. 2
# for gate_up_proj, 3 for qkv_proj). Expand to a per-row
# scalar matching layer.weight's output dim so it
# broadcasts against the per-block scale.
per_row_scale_2 = layer.weight_scale_2.repeat_interleave(
torch.tensor(
layer.logical_widths, device=layer.weight_scale_2.device
)
).view(-1, 1)
# Dequantize to fp32: the intermediate feeds straight into the
# MXFP4 requant
dequantized_weight = dequantize_nvfp4(
layer.weight,
layer.weight_scale,
per_row_scale_2,
out_dtype=torch.float32,
)
mxfp4_weight, mxfp4_scale = dynamic_mxfp4_quant(dequantized_weight)
layer.weight = torch.nn.Parameter(mxfp4_weight, requires_grad=False)
layer.weight_scale = torch.nn.Parameter(
mxfp4_scale, requires_grad=False
)
del layer.weight_scale_2
del layer._load_device
return loader
def get_online_mxfp4_weight_loader(
self,
layer,
@@ -14,10 +14,12 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.quantization.dequantization import (
copy_missing_attrs,
dequantize_fp8,
dequantize_nvfp4,
)
from sglang.srt.layers.quantization.fp8 import Fp8Config, Fp8MoEMethod
from sglang.srt.layers.quantization.online_quantization import CopyNumelCounter
from sglang.srt.layers.quantization.quark.schemes import QuarkMoEScheme
from sglang.srt.layers.quantization.quark.utils import Nvfp4SourceConfig
from sglang.srt.utils import (
get_bool_env_var,
is_gfx95_supported,
@@ -26,6 +28,8 @@ from sglang.srt.utils import (
)
from sglang.srt.utils.common import is_gfx95_supported
NVFP4_BLOCK_SIZE = 16
if TYPE_CHECKING:
from sglang.srt.layers.moe.token_dispatcher import (
CombineInput,
@@ -107,57 +111,68 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
original_weight_loader = extra_weight_attrs.get("weight_loader")
with_bias = extra_weight_attrs.pop("with_bias", False)
self.with_bias = with_bias
# Handle FP8 to MXFP4 requantization
# Handle source-checkpoint -> MXFP4 requantization at load time. The
# source may be NVFP4 (ModelOpt/Quark) or FP8 (block-quantized).
if self.dequantization_config is not None:
if not isinstance(self.dequantization_config, Fp8Config):
raise NotImplementedError(
f"Requantization in QuarkW4A4MXFp4MoEMethod from {self.dequantization_config.__class__.__name__} is not supported, only Fp8Config is supported."
)
if self.dequantization_config.use_mxfp8:
raise NotImplementedError(
"use_mxfp8=True is not supported in Quark MXFP4 requantization."
)
block_quant = self.dequantization_config.weight_block_size is not None
if not block_quant:
raise NotImplementedError(
"Only block_quant=True is supported in Quark MXFP4 requantization, got block_quant=False."
)
# `_fp8_loaded_numel` is used to trigger FP8 -> MXFP4 requantization once all weights are loaded.
# `_fp8_materialized` is used to ensure only one thread materializes weights from meta device.
layer._fp8_loaded_numel = 0
layer._fp8_materialized = False
layer._load_device = torch.get_default_device()
layer._fp8_loading_lock = threading.Lock()
# Custom weight loader handling FP8->MXFP4 conversion.
fp8_to_mxfp4_weight_loader = self.get_online_fp8_to_mxfp4_weight_loader(
layer, original_weight_loader
)
extra_weight_attrs["weight_loader"] = fp8_to_mxfp4_weight_loader
# Create FP8 MoE weight parameters on meta device to avoid device memory overhead during weight loading, as the resulting model uses MXFP4 using less device memory.
# The weight loader handles progressive FP8 weight materialization on device.
with torch.device("meta"):
Fp8MoEMethod.create_fp8_moe_weight_(
if isinstance(self.dequantization_config, Nvfp4SourceConfig):
self._create_weights_from_nvfp4_moe(
layer=layer,
num_experts=num_experts,
hidden_size=hidden_size,
intermediate_size_per_partition=intermediate_size_per_partition,
block_quant=block_quant,
quant_config=self.dequantization_config,
use_mxfp8=False,
is_checkpoint_fp8_serialized=True,
is_fp4_expert=False,
params_dtype=params_dtype,
with_bias=with_bias,
**extra_weight_attrs,
original_weight_loader=original_weight_loader,
extra_weight_attrs=extra_weight_attrs,
)
elif isinstance(self.dequantization_config, Fp8Config):
with_bias = extra_weight_attrs.pop("with_bias", False)
self.with_bias = with_bias
if self.dequantization_config.use_mxfp8:
raise NotImplementedError(
"use_mxfp8=True is not supported in Quark MXFP4 requantization."
)
block_quant = self.dequantization_config.weight_block_size is not None
if not block_quant:
raise NotImplementedError(
"Only block_quant=True is supported in Quark MXFP4 requantization, got block_quant=False."
)
# `_fp8_loaded_numel` is used to trigger FP8 -> MXFP4 requantization once all weights are loaded.
# `_fp8_materialized` is used to ensure only one thread materializes weights from meta device.
layer._fp8_loaded_numel = 0
layer._fp8_materialized = False
layer._load_device = torch.get_default_device()
layer._fp8_loading_lock = threading.Lock()
# Custom weight loader handling FP8->MXFP4 conversion.
fp8_to_mxfp4_weight_loader = self.get_online_fp8_to_mxfp4_weight_loader(
layer, original_weight_loader
)
extra_weight_attrs["weight_loader"] = fp8_to_mxfp4_weight_loader
# Create FP8 MoE weight parameters on meta device to avoid device memory overhead during weight loading, as the resulting model uses MXFP4 using less device memory.
# The weight loader handles progressive FP8 weight materialization on device.
with torch.device("meta"):
Fp8MoEMethod.create_fp8_moe_weight_(
layer=layer,
num_experts=num_experts,
hidden_size=hidden_size,
intermediate_size_per_partition=intermediate_size_per_partition,
block_quant=block_quant,
quant_config=self.dequantization_config,
use_mxfp8=False,
is_checkpoint_fp8_serialized=True,
is_fp4_expert=False,
params_dtype=params_dtype,
with_bias=with_bias,
**extra_weight_attrs,
)
else:
raise NotImplementedError(
f"Requantization in QuarkW4A4MXFp4MoE from {self.dequantization_config.__class__.__name__} is not supported."
)
return
@@ -273,6 +288,230 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
layer.register_parameter("w13_weight_scale", w13_weight_scale)
layer.register_parameter("w2_weight_scale", w2_weight_scale)
def _create_weights_from_nvfp4_moe(
self,
*,
layer,
num_experts,
hidden_size,
intermediate_size_per_partition,
original_weight_loader,
extra_weight_attrs,
):
layer._nvfp4_loaded_numel = 0
layer._load_device = torch.device(f"cuda:{torch.cuda.current_device()}")
layer._nvfp4_loading_lock = threading.Lock()
nvfp4_loader = self.get_online_nvfp4_to_mxfp4_weight_loader(
layer, original_weight_loader
)
extra_weight_attrs["weight_loader"] = nvfp4_loader
def _param(shape, dtype):
return torch.nn.Parameter(
torch.empty(*shape, dtype=dtype, device=layer._load_device),
requires_grad=False,
)
params = {
"w13_weight": _param(
(num_experts, 2 * intermediate_size_per_partition, hidden_size // 2),
torch.uint8,
),
"w2_weight": _param(
(num_experts, hidden_size, intermediate_size_per_partition // 2),
torch.uint8,
),
"w13_weight_scale": _param(
(
num_experts,
2 * intermediate_size_per_partition,
hidden_size // NVFP4_BLOCK_SIZE,
),
torch.float8_e4m3fn,
),
"w2_weight_scale": _param(
(
num_experts,
hidden_size,
intermediate_size_per_partition // NVFP4_BLOCK_SIZE,
),
torch.float8_e4m3fn,
),
}
# w13 fuses gate(w1)+up(w3): FusedMoE stores a per-tensor scale for
# each at param[expert][0|1], so shape is [E, 2]. w2 (down) is single.
params["w13_weight_scale_2"] = _param((num_experts, 2), torch.float32)
params["w2_weight_scale_2"] = _param((num_experts,), torch.float32)
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
# FusedMoE's scale loader dispatches on param.quant_method. NVFP4
# per-block weight_scale -> GROUP; per-tensor weight_scale_2 -> TENSOR.
# (The packed weight tensors skip that branch, name has no "scale".)
for name, param in params.items():
layer.register_parameter(name, param)
attrs = dict(extra_weight_attrs)
if name.endswith("weight_scale_2"):
attrs["quant_method"] = FusedMoeWeightScaleSupported.TENSOR.value
elif name.endswith("weight_scale"):
attrs["quant_method"] = FusedMoeWeightScaleSupported.GROUP.value
set_weight_attrs(param, attrs)
# NVFP4 checkpoints carry per-expert `input_scale` (activation scale)
# per projection. MXFP4 uses dynamic activation quant; discard them but
# register slots so upstream MoE loaders that route w1/w3.input_scale ->
# w13_input_scale find a target. No-op loader absorbs any call shape.
def _discard_loader(param, loaded_weight, weight_name, shard_id, expert_id):
pass
w13_input_scale = torch.nn.Parameter(
torch.empty(num_experts, dtype=torch.float32, device=layer._load_device),
requires_grad=False,
)
w2_input_scale = torch.nn.Parameter(
torch.empty(num_experts, dtype=torch.float32, device=layer._load_device),
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, "weight_loader": _discard_loader}
)
set_weight_attrs(
w2_input_scale, {**extra_weight_attrs, "weight_loader": _discard_loader}
)
def get_online_nvfp4_to_mxfp4_weight_loader(self, layer, original_weight_loader):
"""NVFP4 MoE loader: expert-wise dequant+requant once all source bytes
are in place."""
bulk_names = ["w13_weight", "w2_weight", "w13_weight_scale", "w2_weight_scale"]
scale2_names = ["w13_weight_scale_2", "w2_weight_scale_2"]
def loader(param, loaded_weight, weight_name, shard_id, expert_id):
is_scale_2 = "weight_scale_2" in weight_name
is_scale = ("weight_scale" in weight_name) and not is_scale_2
is_w13 = "w13" in weight_name
assert torch.cuda.current_device() == layer._load_device.index
with layer._nvfp4_loading_lock:
if is_scale_2:
name = "w13_weight_scale_2" if is_w13 else "w2_weight_scale_2"
elif is_scale:
name = "w13_weight_scale" if is_w13 else "w2_weight_scale"
else:
name = "w13_weight" if is_w13 else "w2_weight"
param = getattr(layer, name)
counter = CopyNumelCounter()
with counter:
original_weight_loader(
param, loaded_weight, weight_name, shard_id, expert_id
)
with layer._nvfp4_loading_lock:
layer._nvfp4_loaded_numel += counter.copied_numel
total = sum(
getattr(layer, name).numel() for name in bulk_names + scale2_names
)
if layer._nvfp4_loaded_numel == total:
self._requantize_nvfp4_to_mxfp4(layer, "w13")
self._requantize_nvfp4_to_mxfp4(layer, "w2")
for name in scale2_names:
delattr(layer, name)
del layer._load_device
return loader
def _requantize_nvfp4_to_mxfp4(self, layer, prefix):
# dynamic_mxfp4_quant is 2-D only; loop over experts.
packed_weight = getattr(layer, f"{prefix}_weight")
weight_scale = getattr(layer, f"{prefix}_weight_scale")
weight_scale_2 = getattr(layer, f"{prefix}_weight_scale_2")
# Zero-pad the intermediate dim up to the AITER MoE alignment before the
# MXFP4 requant. (process_weights_after_loading's e8m0_shuffle pads column
# count up to a multiple of 8 which could cause weight K-blocks to be
# miscalculated, leading to scale misalignment and garbage output
inter_pad = 0
if _use_aiter:
if prefix == "w2": # [E, hidden, inter // 2]
real_inter = packed_weight.shape[-1] * 2
else: # w13
real_inter = packed_weight.shape[1] // 2
_, w2_down_dim, _ = get_moe_weight_sizes(
real_inter, is_concat=True, is_packed=True, is_aiter_moe=True
)
inter_pad = max(0, w2_down_dim * 2 - real_inter)
num_experts = packed_weight.shape[0]
# Write each expert's MXFP4 result into a preallocated destination
mxfp4_weight = None
mxfp4_scale = None
for expert_idx in range(num_experts):
if prefix == "w13":
# weight_scale_2[expert_idx] = [gate_scale, up_scale]; the fused
# weight is [gate_rows; up_rows] so expand each scalar over its
# half as a per-row [2I, 1] multiplier.
half = packed_weight[expert_idx].shape[0] // 2
expert_scale_2 = torch.cat(
[
weight_scale_2[expert_idx, 0].repeat(half),
weight_scale_2[expert_idx, 1].repeat(half),
]
).view(-1, 1)
else: # w2: single per-expert per-tensor scalar
expert_scale_2 = weight_scale_2[expert_idx]
dequantized_weight = dequantize_nvfp4(
packed_weight[expert_idx],
weight_scale[expert_idx],
expert_scale_2,
out_dtype=torch.float32,
)
if inter_pad:
if prefix == "w2":
# Pad the trailing K dim with zeros.
dequantized_weight = torch.nn.functional.pad(
dequantized_weight, (0, inter_pad)
)
else:
# w13: pad each of the gate/up halves' rows so the [gate; up]
# split properly
half_rows = dequantized_weight.shape[0] // 2
gate = torch.nn.functional.pad(
dequantized_weight[:half_rows], (0, 0, 0, inter_pad)
)
up = torch.nn.functional.pad(
dequantized_weight[half_rows:], (0, 0, 0, inter_pad)
)
dequantized_weight = torch.cat([gate, up], dim=0)
requantized_weight, requantized_scale = dynamic_mxfp4_quant(
dequantized_weight
)
if mxfp4_weight is None:
mxfp4_weight = torch.empty(
(num_experts, *requantized_weight.shape),
dtype=requantized_weight.dtype,
device=requantized_weight.device,
)
mxfp4_scale = torch.empty(
(num_experts, *requantized_scale.shape),
dtype=requantized_scale.dtype,
device=requantized_scale.device,
)
mxfp4_weight[expert_idx] = requantized_weight
mxfp4_scale[expert_idx] = requantized_scale
setattr(
layer,
f"{prefix}_weight",
torch.nn.Parameter(mxfp4_weight, requires_grad=False),
)
setattr(
layer,
f"{prefix}_weight_scale",
torch.nn.Parameter(mxfp4_scale, requires_grad=False),
)
def get_online_weight_loader(self, layer, original_weight_loader):
"""
Wrap the original weight loader to perform online MXFP4 quantization for MoE layers.
@@ -2,9 +2,19 @@
import re
from collections.abc import Iterable, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Any, Optional
@dataclass
class Nvfp4SourceConfig:
"""Dispatch marker for online NVFP4 -> MXFP4 re-quantization, carried on
`QuarkConfig.dequantization_config` to represent an NVFP4 source
Only ModelOpt / AMD Quark NVFP4 (per-tensor `weight_scale_2`
that multiplies the per-block scale) is supported."""
import torch
try:
+20 -2
View File
@@ -283,7 +283,6 @@ def get_quant_config(
if hf_quant_config is not None:
if not isinstance(hf_quant_config, dict):
hf_quant_config = hf_quant_config.to_dict()
# For modelopt_mixed, config.json's quantization_config may not
# contain all runtime metadata. Fall through to the file-based
# hf_quant_config.json path when the per-layer map or KV-cache
@@ -302,7 +301,7 @@ def get_quant_config(
hf_quant_config["packed_modules_mapping"] = packed_modules_mapping
hf_quant_config["hf_config"] = model_config.hf_config
# This is only used by quantization methods that support requantization (e.g. from fp8 to mxfp4).
# This is only used by quantization methods that support requantization (e.g. from nvfp4/fp8 to mxfp4).
if model_config.quantization in REQUANTIZATION_METHODS:
hf_quant_config["requantization_method"] = model_config.quantization
@@ -348,6 +347,25 @@ def get_quant_config(
quant_cls = Fp8Config
return quant_cls(use_mxfp8=True, is_checkpoint_fp8_serialized=False)
if model_config.quantization == "quark_mxfp4":
# Some ModelOpt NVFP4 checkpoints store quant metadata only in
# hf_quant_config.json; others duplicate it in config.json. Read
# hf_quant_config.json first when present and FP4-typed.
modelopt_quant_path = os.path.join(hf_folder, "hf_quant_config.json")
if os.path.isfile(modelopt_quant_path):
with open(modelopt_quant_path) as f:
raw_quant_config = json.load(f)
source_quant = raw_quant_config.get("quantization", raw_quant_config)
if "FP4" in (source_quant.get("quant_algo") or "").upper():
flat_quant_config = dict(source_quant)
flat_quant_config["quant_method"] = (
raw_quant_config.get("producer", {}).get("name") or "modelopt"
)
flat_quant_config["requantization_method"] = (
model_config.quantization
)
flat_quant_config["packed_modules_mapping"] = packed_modules_mapping
flat_quant_config["hf_config"] = model_config.hf_config
return quant_cls.from_config(flat_quant_config)
return quant_cls(
online_scheme=model_config.quantization,
hf_config=model_config.hf_config,
+1 -1
View File
@@ -167,7 +167,7 @@ QUANTIZATION_CHOICES = [
"mxfp_w4a8", # for NPU W4A8 (MXFP4 weights + MXFP8 activations)
"quark", # AMD Quark quantizer (FP8 / MXFP4 / Int4FP8 etc.)
"quark_int4fp8_moe",
"quark_mxfp4", # Online MOE + linear quantization.
"quark_mxfp4", # Online MOE + linear quantization (incl. NVFP4 -> MXFP4 requantization).
# Apple Silicon MLX backend — on-the-fly quantization of fp16 weights at load
# time via mlx.nn.quantize. Only takes effect when SGLANG_USE_MLX=1.
"mlx_q4", # 4 bits, group_size=64 (mlx-community default)
+61 -22
View File
@@ -6,6 +6,7 @@ import unittest
from sglang.test.ci.ci_register import register_amd_ci
register_amd_ci(est_time=106, suite="stage-b-test-1-gpu-small-amd-mi35x")
import os
import time
from types import SimpleNamespace
@@ -87,20 +88,10 @@ class TestOnlineQuantizationMemoryLoad(CustomTestCase):
raise RuntimeError(f"Server {url} failed to start in {timeout}s")
time.sleep(1)
# # Extract and display peak GPU memory from logs
combined_output = cls.stdout.getvalue() + cls.stderr.getvalue()
peak_memory_before_load = cls._extract_peak_memory_before_load(combined_output)
if is_cuda_alike() and not peak_memory_before_load:
raise ValueError("Should have found peak memory")
cls.peak_memory_before_load = float(peak_memory_before_load)
memory_increase_load_weights = cls._extract_memory_increase_load_weights(
combined_output
)
if is_cuda_alike() and not memory_increase_load_weights:
raise ValueError("Should have found memory increase in load_weights")
cls.memory_increase_load_weights = float(memory_increase_load_weights)
# Keep the raw server for memory numbers, which are parsed lazily by
# _test_peak_memory so subclasses that don't test memory (e.g. the
# NVFP4->MXFP4 accuracy-only class) don't require these log lines.
cls.combined_output = cls.stdout.getvalue() + cls.stderr.getvalue()
@classmethod
def _extract_peak_memory_before_load(cls, log_output):
@@ -115,8 +106,11 @@ class TestOnlineQuantizationMemoryLoad(CustomTestCase):
@classmethod
def _extract_memory_increase_load_weights(cls, log_output):
"""Extract memory increase during load_weights call."""
# Search for the log message pattern
pattern = r"Memory increase during load_weights:\s+([\d.]+)\s+GiB"
# Signed: the value is (free_before - free_after) around load_weights.
# When the on-device source representation is larger than the loaded
# result (e.g. requantizing to a more compact format), loading frees
# net memory and the reported increase is negative.
pattern = r"Memory increase during load_weights:\s+(-?[\d.]+)\s+GiB"
match = re.search(pattern, log_output)
if match:
return match.group(1)
@@ -138,21 +132,33 @@ class TestOnlineQuantizationMemoryLoad(CustomTestCase):
if not is_cuda_alike():
self.skipTest("not is_cuda_alike")
peak_memory_before_load = self._extract_peak_memory_before_load(
self.combined_output
)
if not peak_memory_before_load:
raise ValueError("Should have found peak memory")
peak_memory_before_load = float(peak_memory_before_load)
memory_increase_load_weights = self._extract_memory_increase_load_weights(
self.combined_output
)
if not memory_increase_load_weights:
raise ValueError("Should have found memory increase in load_weights")
memory_increase_load_weights = float(memory_increase_load_weights)
# NOTE: We can not simply rely on peak memory after `load_weights` as functions used
# in-between (e.g. FP8->MXFP4 requantization) during weight loading may have a higher peak memory footprint
# in-between (e.g. NVFP4->MXFP4 requantization) during weight loading may have a higher peak memory footprint
# than simply the allocated weights.
if add_peak_memory_before_load:
reference_gib = (
self.memory_increase_load_weights + self.peak_memory_before_load
)
reference_gib = memory_increase_load_weights + peak_memory_before_load
else:
reference_gib = self.memory_increase_load_weights
reference_gib = memory_increase_load_weights
assert reference_gib < threshold
if test_start:
# Weights initialized on meta device (not for dense BF16->MXFP4)
assert self.peak_memory_before_load < 5
assert peak_memory_before_load < 5
def _test_gsm8k(self, accuracy_threshold):
"""Helper method to test GSM8K accuracy against a threshold."""
@@ -205,6 +211,39 @@ class TestOnlineQuantizationMemoryLoadMOE(TestOnlineQuantizationMemoryLoad):
self._test_gsm8k(accuracy_threshold=0.89)
class TestNVFP4ToMXFP4MOETP1(TestOnlineQuantizationMemoryLoad):
# ModelOpt NVFP4 export (quant_method="modelopt", quant_algo="NVFP4") =>
# Nvfp4SourceConfig(). Exercises the NVFP4 -> MXFP4 MoE requantization path:
# the per-expert dequantize_nvfp4 + dynamic_mxfp4_quant requant, and the w13
# gate/up weight_scale_2 split in _requantize_nvfp4_to_mxfp4.
model = "nvidia/Qwen3-30B-A3B-NVFP4" # NVFP4 model
tp = 1
def test_gsm8k(self):
# Requantized NVFP4 -> MXFP4 observed accuracy: ~0.88
# (BF16 Qwen/Qwen3-30B-A3B reference: ~0.94).
self._test_gsm8k(accuracy_threshold=0.85)
@unittest.skipIf(is_in_ci(), "local test only")
class TestDeepSeekR10528NVFP4ToMXFP4(TestOnlineQuantizationMemoryLoad):
# NVFP4 to MXFP4 online requantization for DeepSeek-R1-0528-NVFP4 on TP=8.
# Exercises the MLA attention path (attention_backend=aiter), multi-threaded
# weight loading, and the per-expert NVFP4 MoE requantization path.
model = "nvidia/DeepSeek-R1-0528-NVFP4" # NVFP4 model
tp = 8
runner_args = [
"--attention-backend",
"aiter",
"--model-loader-extra-config",
'{"enable_multithread_load": true}',
]
def test_gsm8k(self):
# Requantized NVFP4 -> MXFP4 observed accuracy: ~0.95.
self._test_gsm8k(accuracy_threshold=0.90)
class TestFP8ToMXFP4DenseTP1(TestOnlineQuantizationMemoryLoad):
tp = 1
model = "Qwen/Qwen3-8B-FP8"
@@ -7,7 +7,15 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
import unittest
from unittest.mock import patch
from sglang.srt.layers.quantization.quark.quark import QuarkConfig
import torch
from sglang.srt.layers.quantization.quark.quark import (
QuarkConfig,
_build_mixed_precision_layer_quant_config,
_mixed_precision_layer_map,
_parse_nvfp4_excludes,
)
from sglang.srt.layers.quantization.quark.utils import check_equal_or_regex_match
from sglang.test.test_utils import CustomTestCase
_GET_CAP = "sglang.srt.layers.quantization.quark.quark.get_device_capability"
@@ -83,5 +91,121 @@ class TestCheckSchemeSupportedError(CustomTestCase):
self.assertFalse(ok)
class TestMixedPrecisionLayerConfig(CustomTestCase):
"""NVFP4-only-experts + FP8-elsewhere online requant (quark_mxfp4).
A MIXED_PRECISION NVFP4 checkpoint (e.g. nvidia/Qwen3.5-397B-A17B-NVFP4-V2)
keeps some layers in NVFP4 while others in FP8. Online requant must send
only the NVFP4 layers through the dequant->MXFP4 path and load the FP8 layers
as FP8.
"""
_LAYER_MAP_SRC = {
"quant_algo": "MIXED_PRECISION",
"quantized_layers": {
"model.language_model.layers.0.self_attn.q_proj": {"quant_algo": "FP8"},
"model.language_model.layers.0.self_attn.k_proj": {"quant_algo": "FP8"},
"model.language_model.layers.0.self_attn.v_proj": {"quant_algo": "FP8"},
"model.language_model.layers.0.self_attn.o_proj": {"quant_algo": "FP8"},
"model.language_model.layers.0.mlp.shared_expert.gate_proj": {
"quant_algo": "FP8"
},
"model.language_model.layers.0.mlp.shared_expert.down_proj": {
"quant_algo": "FP8"
},
"model.language_model.layers.0.mlp.experts": {
"quant_algo": "NVFP4",
"group_size": 16,
},
"model.language_model.layers.1.mlp.experts": {
"quant_algo": "NVFP4",
"group_size": 16,
},
"model.language_model.layers.1.self_attn.q_proj": {"quant_algo": "FP8"},
},
}
def _build_bare_config(self) -> QuarkConfig:
layer_map = _mixed_precision_layer_map(self._LAYER_MAP_SRC)
layer_quant_config, has_nvfp4 = _build_mixed_precision_layer_quant_config(
layer_map
)
self.assertTrue(has_nvfp4)
synth_config = QuarkConfig._create_online_mxfp4_config(
model_type="qwen3_5_moe",
layer_quant_config=layer_quant_config,
)
synth_config["packed_modules_mapping"] = {
"qkv_proj": ["q_proj", "k_proj", "v_proj"],
}
quark_config = _bare_config()
quark_config.quant_config = synth_config
quark_config.packed_modules_mapping = synth_config["packed_modules_mapping"]
quark_config.exclude_layers = synth_config["exclude"]
return quark_config
def test_experts_route_to_mxfp4_requant(self):
# fnmatch keys (not `re:`) must match the sglang module path so experts
# hit the fp4 target, not fall through to the global config
quark_config = self._build_bare_config()
matched = quark_config._find_matched_config(
"model.layers.0.mlp.experts", torch.nn.Module()
)
self.assertEqual(matched["weight"]["dtype"], "fp4")
self.assertEqual(matched["weight"]["group_size"], 32)
def test_fp8_layers_not_requantized(self):
quark_config = self._build_bare_config()
for name in (
"model.layers.0.self_attn.o_proj",
"model.layers.0.mlp.shared_expert.gate_proj",
"model.layers.0.mlp.shared_expert.down_proj",
):
matched = quark_config._find_matched_config(name, torch.nn.Module())
self.assertEqual(matched["weight"]["dtype"], "fp8_e4m3", msg=name)
self.assertEqual(matched["weight"]["qscheme"], "per_tensor", msg=name)
def test_fused_qkv_shards_share_fp8_scheme(self):
# _find_matched_config expands qkv_proj -> q/k/v shards and requires a
# consistent scheme; all three are FP8 so this must resolve
quark_config = self._build_bare_config()
matched = quark_config._find_matched_config(
"model.layers.0.self_attn.qkv_proj", torch.nn.Module()
)
self.assertEqual(matched["weight"]["dtype"], "fp8_e4m3")
def test_shared_expert_fusion_disabled_on_precision_mismatch(self):
quark_config = self._build_bare_config()
self.assertFalse(quark_config.can_fuse_shared_expert())
def test_mixed_precision_skips_model_type_default_excludes(self):
quark_config = self._build_bare_config()
self.assertNotIn("re:.*shared_expert", quark_config.exclude_layers)
self.assertNotIn("re:.*o_proj", quark_config.exclude_layers)
def test_non_mixed_config_returns_none(self):
self.assertIsNone(_mixed_precision_layer_map({"quant_algo": "NVFP4"}))
class TestParseNvfp4Excludes(CustomTestCase):
"""ModelOpt `ignore` lists mix `re:`-prefixed regexes with fnmatch globs."""
def test_already_regex_entries_pass_through_and_match(self):
# wrapping an already-`re:` entry with another `re:` +
# fnmatch.translate produced `re:(?s:re:\\..*...)` which never matches,
excludes = _parse_nvfp4_excludes(
{"ignore": [r"re:.*linear_attn\.in_proj_a$", "mtp*"]}
)
self.assertTrue(
check_equal_or_regex_match("model.layers.0.linear_attn.in_proj_a", excludes)
)
# fnmatch glob still translated and matches.
self.assertTrue(check_equal_or_regex_match("mtp.layers.0.foo", excludes))
# A quantized layer stays un-excluded.
self.assertFalse(
check_equal_or_regex_match("model.layers.0.mlp.experts", excludes)
)
if __name__ == "__main__":
unittest.main()