diff --git a/docs/docs/advanced_features/quantization.mdx b/docs/docs/advanced_features/quantization.mdx index 044e33000..92abe9989 100644 --- a/docs/docs/advanced_features/quantization.mdx +++ b/docs/docs/advanced_features/quantization.mdx @@ -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, diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index d28987651..0b24cd87c 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -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: diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index d42d7578f..e133358b3 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -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 ), ) diff --git a/python/sglang/srt/layers/quantization/base_config.py b/python/sglang/srt/layers/quantization/base_config.py index 6a81f42b5..509e2cc07 100644 --- a/python/sglang/srt/layers/quantization/base_config.py +++ b/python/sglang/srt/layers/quantization/base_config.py @@ -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 diff --git a/python/sglang/srt/layers/quantization/base_scheme.py b/python/sglang/srt/layers/quantization/base_scheme.py index 4155ad07f..2e62d6c26 100644 --- a/python/sglang/srt/layers/quantization/base_scheme.py +++ b/python/sglang/srt/layers/quantization/base_scheme.py @@ -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): """ diff --git a/python/sglang/srt/layers/quantization/dequantization.py b/python/sglang/srt/layers/quantization/dequantization.py index c2de905a4..5c7e7b34d 100644 --- a/python/sglang/srt/layers/quantization/dequantization.py +++ b/python/sglang/srt/layers/quantization/dequantization.py @@ -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) diff --git a/python/sglang/srt/layers/quantization/quark/quark.py b/python/sglang/srt/layers/quantization/quark/quark.py index 69a7177f0..fed208fb5 100644 --- a/python/sglang/srt/layers/quantization/quark/quark.py +++ b/python/sglang/srt/layers/quantization/quark/quark.py @@ -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..` (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, diff --git a/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4.py b/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4.py index 0b0711ab2..fb25623de 100644 --- a/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4.py +++ b/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4.py @@ -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, diff --git a/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4_moe.py b/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4_moe.py index 0e1557c0a..f05a64627 100644 --- a/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4_moe.py +++ b/python/sglang/srt/layers/quantization/quark/schemes/quark_w4a4_mxfp4_moe.py @@ -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. diff --git a/python/sglang/srt/layers/quantization/quark/utils.py b/python/sglang/srt/layers/quantization/quark/utils.py index 948a79dbb..dcd3aeb59 100644 --- a/python/sglang/srt/layers/quantization/quark/utils.py +++ b/python/sglang/srt/layers/quantization/quark/utils.py @@ -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: diff --git a/python/sglang/srt/model_loader/weight_utils.py b/python/sglang/srt/model_loader/weight_utils.py index fd91d5a21..2d25f2add 100644 --- a/python/sglang/srt/model_loader/weight_utils.py +++ b/python/sglang/srt/model_loader/weight_utils.py @@ -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, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index ea01528c9..0a9bbf2cb 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -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) diff --git a/test/registered/quant/test_quark_mxfp4.py b/test/registered/quant/test_quark_mxfp4.py index 8c1f2f885..afe5118b7 100644 --- a/test/registered/quant/test_quark_mxfp4.py +++ b/test/registered/quant/test_quark_mxfp4.py @@ -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" diff --git a/test/registered/unit/layers/quantization/test_quark_config.py b/test/registered/unit/layers/quantization/test_quark_config.py index bd543e55e..c2990d79b 100644 --- a/test/registered/unit/layers/quantization/test_quark_config.py +++ b/test/registered/unit/layers/quantization/test_quark_config.py @@ -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()