[Model] Support Qwen3.6 ModelOpt mixed NVFP4 (#27906)
This commit is contained in:
@@ -1146,10 +1146,14 @@ class ModelConfig:
|
||||
quant_algo = json_quant_configs.get("quant_algo", None)
|
||||
|
||||
if quant_algo == "MIXED_PRECISION":
|
||||
architectures = getattr(self.hf_config, "architectures", []) or []
|
||||
if getattr(self.hf_config, "model_type", None) == "nemotron_h" or any(
|
||||
arch.startswith("NemotronH") for arch in architectures
|
||||
):
|
||||
quantized_layers = json_quant_configs.get("quantized_layers") or {}
|
||||
has_modelopt_nvfp4_layers = any(
|
||||
str(layer_info.get("quant_algo", "")).upper()
|
||||
in ("NVFP4", "W4A16_NVFP4")
|
||||
for layer_info in quantized_layers.values()
|
||||
if isinstance(layer_info, dict)
|
||||
)
|
||||
if has_modelopt_nvfp4_layers:
|
||||
return {"quant_method": "modelopt_mixed", "quant_algo": quant_algo}
|
||||
return {"quant_method": "w4afp8", "quant_algo": quant_algo}
|
||||
elif quant_algo and ("FP4" in quant_algo or "NVFP4" in quant_algo):
|
||||
|
||||
@@ -61,6 +61,74 @@ logger = logging.getLogger(__name__)
|
||||
_is_npu = is_npu()
|
||||
_is_cpu = is_cpu()
|
||||
|
||||
_UNQUANTIZED_LM_HEAD_METHODS = {
|
||||
"UnquantizedEmbeddingMethod",
|
||||
"UnquantizedLinearMethod",
|
||||
"PackWeightMethod",
|
||||
}
|
||||
|
||||
|
||||
def _has_lm_head_runtime_attrs(lm_head, attr_names: Tuple[str, ...]) -> bool:
|
||||
return all(hasattr(lm_head, attr_name) for attr_name in attr_names)
|
||||
|
||||
|
||||
def should_apply_lm_head_quant_method(lm_head, quant_method) -> bool:
|
||||
if (
|
||||
quant_method is None
|
||||
or not hasattr(lm_head, "weight")
|
||||
or not callable(getattr(quant_method, "apply", None))
|
||||
):
|
||||
return False
|
||||
|
||||
method_name = type(quant_method).__name__
|
||||
if method_name in _UNQUANTIZED_LM_HEAD_METHODS:
|
||||
return False
|
||||
|
||||
# Some draft models share an unquantized target lm_head tensor while still
|
||||
# carrying the draft model's stale ModelOpt quant_method. Only use the
|
||||
# ModelOpt lm_head kernel when the runtime quantization state matches it.
|
||||
if method_name == "ModelOptFp4LinearMethod":
|
||||
if lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
||||
lm_head,
|
||||
(
|
||||
"weight_scale",
|
||||
"weight_global_scale",
|
||||
"workspace",
|
||||
"input_size_per_partition",
|
||||
"output_size_per_partition",
|
||||
),
|
||||
):
|
||||
return True
|
||||
return lm_head.weight.dtype == torch.uint8 and _has_lm_head_runtime_attrs(
|
||||
lm_head,
|
||||
(
|
||||
"weight_scale_interleaved",
|
||||
"alpha",
|
||||
"input_scale_inv",
|
||||
"input_size_per_partition",
|
||||
"output_size_per_partition",
|
||||
),
|
||||
)
|
||||
if method_name == "ModelOptNvFp4A16LinearMethod":
|
||||
return lm_head.weight.dtype == torch.int32 and _has_lm_head_runtime_attrs(
|
||||
lm_head,
|
||||
(
|
||||
"weight_scale",
|
||||
"weight_global_scale",
|
||||
"workspace",
|
||||
"input_size_per_partition",
|
||||
"output_size_per_partition",
|
||||
),
|
||||
)
|
||||
if method_name == "ModelOptFp8LinearMethod":
|
||||
return (
|
||||
lm_head.weight.dtype == torch.float8_e4m3fn
|
||||
and _has_lm_head_runtime_attrs(lm_head, ("weight_scale", "input_scale"))
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
# When set, LogitsProcessor.forward returns an empty output and skips the
|
||||
# LM head + tensor-parallel all-gather. FlashInfer autotune only profiles
|
||||
# attention/MoE/GEMM kernels, so the LM-head all-gather is wasted work --
|
||||
@@ -883,9 +951,12 @@ class LogitsProcessor(nn.Module):
|
||||
lm_head: VocabParallelEmbedding,
|
||||
embedding_bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
quant_method = getattr(lm_head, "quant_method", None)
|
||||
if hasattr(lm_head, "set_lora") and hasattr(lm_head, "apply_lora"):
|
||||
# This is a LoRA-wrapped module, use its forward method
|
||||
logits = lm_head(hidden_states)
|
||||
elif should_apply_lm_head_quant_method(lm_head, quant_method):
|
||||
logits = quant_method.apply(lm_head, hidden_states, embedding_bias)
|
||||
elif hasattr(lm_head, "weight"):
|
||||
# Normal linear layer
|
||||
if self.use_fp32_lm_head:
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
|
||||
|
||||
import regex as re
|
||||
import torch
|
||||
@@ -65,6 +65,7 @@ from sglang.srt.utils.common import (
|
||||
is_sm100_supported,
|
||||
is_sm120_supported,
|
||||
round_up,
|
||||
set_weight_attrs,
|
||||
)
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
from sglang.srt.utils.patch_torch import register_fake_if_exists
|
||||
@@ -77,6 +78,25 @@ if TYPE_CHECKING:
|
||||
)
|
||||
from sglang.srt.models.utils import WeightsMapper
|
||||
|
||||
|
||||
def _make_per_tensor_scale_parameter(
|
||||
shape,
|
||||
weight_loader,
|
||||
*,
|
||||
fill_value: Optional[float] = None,
|
||||
needs_scalar_to_array: bool = False,
|
||||
) -> PerTensorScaleParameter:
|
||||
data = (
|
||||
torch.empty(shape, dtype=torch.float32)
|
||||
if fill_value is None
|
||||
else torch.full(shape, fill_value, dtype=torch.float32)
|
||||
)
|
||||
scale = PerTensorScaleParameter(data=data, weight_loader=weight_loader)
|
||||
if needs_scalar_to_array:
|
||||
set_weight_attrs(scale, {"needs_scalar_to_array": True})
|
||||
return scale
|
||||
|
||||
|
||||
try:
|
||||
from flashinfer import mm_fp4 as flashinfer_fp4_gemm
|
||||
from flashinfer import reorder_rows_for_gated_act_gemm, shuffle_matrix_sf_a
|
||||
@@ -281,8 +301,9 @@ class ModelOptQuantConfig(QuantizationConfig):
|
||||
) -> Optional[QuantizeMethodBase]:
|
||||
from sglang.srt.layers.linear import LinearBase
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||
|
||||
if isinstance(layer, LinearBase):
|
||||
if isinstance(layer, (LinearBase, ParallelLMHead)):
|
||||
if is_layer_skipped(
|
||||
prefix, self.exclude_modules, self.packed_modules_mapping
|
||||
) or self.is_layer_excluded(prefix):
|
||||
@@ -433,6 +454,8 @@ class ModelOptFp8Config(ModelOptQuantConfig):
|
||||
and kv_cache_scheme.get("num_bits") == 8
|
||||
):
|
||||
kv_cache_quant_method = "FP8"
|
||||
else:
|
||||
kv_cache_quant_method = config.get("kv_cache_quant_algo")
|
||||
|
||||
# Map 'ignore' field to 'exclude_modules'
|
||||
exclude_modules = config.get("ignore")
|
||||
@@ -536,17 +559,13 @@ class ModelOptFp8LinearMethod(LinearMethodBase):
|
||||
if self.quant_config.is_checkpoint_fp8_serialized:
|
||||
# Register weight and input scales
|
||||
for scale_name in ["weight_scale", "input_scale"]:
|
||||
layer.register_parameter(
|
||||
scale_name,
|
||||
PerTensorScaleParameter(
|
||||
data=torch.full(
|
||||
(len(output_partition_sizes),),
|
||||
torch.finfo(torch.float32).min,
|
||||
dtype=torch.float32,
|
||||
),
|
||||
weight_loader=weight_loader,
|
||||
),
|
||||
scale = _make_per_tensor_scale_parameter(
|
||||
(len(output_partition_sizes),),
|
||||
weight_loader=weight_loader,
|
||||
fill_value=torch.finfo(torch.float32).min,
|
||||
needs_scalar_to_array=True,
|
||||
)
|
||||
layer.register_parameter(scale_name, scale)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
"""Requantizes weights after loading using the maximum scale."""
|
||||
@@ -604,11 +623,13 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
quantized_layers: Dict[str, Dict[str, Any]],
|
||||
fp8_config: ModelOptFp8Config,
|
||||
nvfp4_config: ModelOptFp4Config,
|
||||
nvfp4a16_config: ModelOptFp4Config,
|
||||
) -> None:
|
||||
super().__init__(kv_cache_quant_algo, exclude_modules, packed_modules_mapping)
|
||||
self.quantized_layers = quantized_layers
|
||||
self.fp8_config = fp8_config
|
||||
self.nvfp4_config = nvfp4_config
|
||||
self.nvfp4a16_config = nvfp4a16_config
|
||||
|
||||
@classmethod
|
||||
def override_quantization_method(cls, hf_quant_config, user_quant):
|
||||
@@ -652,6 +673,8 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
kv_cache_quant_algo = "NVFP4"
|
||||
else:
|
||||
kv_cache_quant_algo = "auto"
|
||||
else:
|
||||
kv_cache_quant_algo = config.get("kv_cache_quant_algo")
|
||||
exclude_modules = config.get("ignore")
|
||||
quantized_layers = config.get("quantized_layers", {})
|
||||
else:
|
||||
@@ -672,7 +695,10 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
|
||||
group_size = None
|
||||
for layer_info in quantized_layers.values():
|
||||
if layer_info.get("quant_algo", "").upper() == "NVFP4":
|
||||
if layer_info.get("quant_algo", "").upper() in (
|
||||
"NVFP4",
|
||||
"W4A16_NVFP4",
|
||||
):
|
||||
group_size = layer_info.get("group_size", 16)
|
||||
break
|
||||
if group_size is None:
|
||||
@@ -692,6 +718,14 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
packed_modules_mapping=packed_modules_mapping,
|
||||
group_size=group_size,
|
||||
)
|
||||
nvfp4a16_config = ModelOptFp4Config(
|
||||
is_checkpoint_nvfp4_serialized=True,
|
||||
kv_cache_quant_algo=kv_cache_quant_algo,
|
||||
exclude_modules=[],
|
||||
packed_modules_mapping=packed_modules_mapping,
|
||||
group_size=group_size,
|
||||
use_per_token_activation=False,
|
||||
)
|
||||
|
||||
return cls(
|
||||
kv_cache_quant_algo=kv_cache_quant_algo,
|
||||
@@ -700,6 +734,7 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
quantized_layers=quantized_layers,
|
||||
fp8_config=fp8_config,
|
||||
nvfp4_config=nvfp4_config,
|
||||
nvfp4a16_config=nvfp4a16_config,
|
||||
)
|
||||
|
||||
def apply_weight_name_mapper(self, hf_to_sglang_mapper: WeightsMapper):
|
||||
@@ -710,17 +745,21 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
)
|
||||
|
||||
def _resolve_quant_algo(self, prefix: str) -> Optional[str]:
|
||||
if prefix in self.quantized_layers:
|
||||
return self.quantized_layers[prefix]["quant_algo"].upper()
|
||||
for candidate in self._quantized_layer_prefix_candidates(prefix):
|
||||
if candidate in self.quantized_layers:
|
||||
return self.quantized_layers[candidate]["quant_algo"].upper()
|
||||
|
||||
proj_name = prefix.rsplit(".", 1)[-1]
|
||||
if self.packed_modules_mapping and proj_name in self.packed_modules_mapping:
|
||||
algos = set()
|
||||
base = prefix.rsplit(".", 1)[0]
|
||||
for shard_name in self.packed_modules_mapping[proj_name]:
|
||||
shard_prefix = f"{base}.{shard_name}"
|
||||
if shard_prefix in self.quantized_layers:
|
||||
algos.add(self.quantized_layers[shard_prefix]["quant_algo"].upper())
|
||||
for base_candidate in self._quantized_layer_prefix_candidates(base):
|
||||
for shard_name in self.packed_modules_mapping[proj_name]:
|
||||
shard_prefix = f"{base_candidate}.{shard_name}"
|
||||
if shard_prefix in self.quantized_layers:
|
||||
algos.add(
|
||||
self.quantized_layers[shard_prefix]["quant_algo"].upper()
|
||||
)
|
||||
if len(algos) == 1:
|
||||
return algos.pop()
|
||||
if len(algos) > 1:
|
||||
@@ -729,22 +768,42 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
"All shards must use the same quantization."
|
||||
)
|
||||
|
||||
prefix_dot = prefix + "."
|
||||
for key, info in self.quantized_layers.items():
|
||||
if key.startswith(prefix_dot):
|
||||
return info["quant_algo"].upper()
|
||||
for candidate in self._quantized_layer_prefix_candidates(prefix):
|
||||
prefix_dot = candidate + "."
|
||||
for key, info in self.quantized_layers.items():
|
||||
if key.startswith(prefix_dot):
|
||||
return info["quant_algo"].upper()
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _quantized_layer_prefix_candidates(prefix: str) -> Tuple[str, ...]:
|
||||
candidates = [prefix]
|
||||
|
||||
if prefix.endswith(".lm_head"):
|
||||
candidates.append("lm_head")
|
||||
|
||||
if prefix.startswith("language_model.model."):
|
||||
candidates.append(
|
||||
"model.language_model." + prefix[len("language_model.model.") :]
|
||||
)
|
||||
elif prefix.startswith("model.language_model."):
|
||||
candidates.append(
|
||||
"language_model.model." + prefix[len("model.language_model.") :]
|
||||
)
|
||||
|
||||
return tuple(dict.fromkeys(candidates))
|
||||
|
||||
def get_quant_method(
|
||||
self, layer: torch.nn.Module, prefix: str
|
||||
) -> Optional[QuantizeMethodBase]:
|
||||
from sglang.srt.layers.linear import LinearBase
|
||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||
|
||||
quant_algo = self._resolve_quant_algo(prefix)
|
||||
|
||||
if isinstance(layer, LinearBase):
|
||||
if isinstance(layer, (LinearBase, ParallelLMHead)):
|
||||
if is_layer_skipped(
|
||||
prefix, self.exclude_modules, self.packed_modules_mapping
|
||||
) or self.is_layer_excluded(prefix):
|
||||
@@ -753,6 +812,8 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
return ModelOptFp8LinearMethod(self.fp8_config)
|
||||
if quant_algo == "NVFP4":
|
||||
return ModelOptFp4LinearMethod(self.nvfp4_config)
|
||||
if quant_algo == "W4A16_NVFP4":
|
||||
return ModelOptNvFp4A16LinearMethod(self.nvfp4a16_config)
|
||||
return UnquantizedLinearMethod()
|
||||
|
||||
if self.kv_cache_quant_algo and isinstance(layer, RadixAttention):
|
||||
@@ -765,6 +826,8 @@ class ModelOptMixedPrecisionConfig(ModelOptQuantConfig):
|
||||
return ModelOptFp8MoEMethod(self.fp8_config)
|
||||
if quant_algo == "NVFP4":
|
||||
return ModelOptNvFp4FusedMoEMethod(self.nvfp4_config)
|
||||
if quant_algo == "W4A16_NVFP4":
|
||||
return ModelOptNvFp4FusedMoEMethod(self.nvfp4a16_config)
|
||||
return None
|
||||
|
||||
return None
|
||||
@@ -1139,8 +1202,9 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
||||
)
|
||||
self.group_size = group_size
|
||||
self.use_per_token_activation = (
|
||||
use_per_token_activation
|
||||
or envs.SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION.get()
|
||||
envs.SGLANG_FLASHINFER_NVFP4_PER_TOKEN_ACTIVATION.get()
|
||||
if use_per_token_activation is None
|
||||
else use_per_token_activation
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -1229,7 +1293,7 @@ class ModelOptFp4Config(ModelOptQuantConfig):
|
||||
else:
|
||||
kv_cache_quant_algo = "auto"
|
||||
else:
|
||||
kv_cache_quant_algo = "auto"
|
||||
kv_cache_quant_algo = config.get("kv_cache_quant_algo") or "auto"
|
||||
|
||||
group_size = config.get("group_size")
|
||||
# If group_size is not at top level, try to extract from config_groups
|
||||
@@ -1400,16 +1464,17 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
||||
)
|
||||
layer.register_parameter("weight", weight)
|
||||
|
||||
input_scale = PerTensorScaleParameter(
|
||||
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
|
||||
input_scale = _make_per_tensor_scale_parameter(
|
||||
(len(output_partition_sizes),),
|
||||
weight_loader=weight_loader,
|
||||
needs_scalar_to_array=True,
|
||||
)
|
||||
|
||||
layer.register_parameter("input_scale", input_scale)
|
||||
|
||||
weight_scale_2 = PerTensorScaleParameter(
|
||||
data=torch.empty(len(output_partition_sizes), dtype=torch.float32),
|
||||
weight_scale_2 = _make_per_tensor_scale_parameter(
|
||||
(len(output_partition_sizes),),
|
||||
weight_loader=weight_loader,
|
||||
needs_scalar_to_array=True,
|
||||
)
|
||||
layer.register_parameter("weight_scale_2", weight_scale_2)
|
||||
|
||||
@@ -1450,6 +1515,7 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
||||
)
|
||||
copy_or_rebind_param(layer, "input_global_scale", input_scale_2)
|
||||
copy_or_rebind_param(layer, "weight_global_scale", weight_scale_2)
|
||||
layer.quant_config = self.quant_config
|
||||
prepare_nvfp4_layer_for_marlin(layer)
|
||||
layer.weights_padding_cols = 0
|
||||
return
|
||||
@@ -1668,6 +1734,131 @@ class ModelOptFp4LinearMethod(LinearMethodBase):
|
||||
return out.view(*output_shape)
|
||||
|
||||
|
||||
class ModelOptNvFp4A16LinearMethod(LinearMethodBase):
|
||||
"""Linear method for ModelOpt NVFP4A16 checkpoints.
|
||||
|
||||
Loads packed NVFP4 weights with fp16/bf16 activations. ModelOpt may still
|
||||
provide input_scale tensors for fused loader compatibility; they are
|
||||
consumed during loading and discarded before runtime.
|
||||
"""
|
||||
|
||||
def __init__(self, quant_config: ModelOptFp4Config):
|
||||
self.quant_config = quant_config
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: List[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
del input_size, output_size
|
||||
if not self.quant_config.is_checkpoint_nvfp4_serialized:
|
||||
raise ValueError(
|
||||
"W4A16_NVFP4 quantization was selected, "
|
||||
"dynamic quantization is not supported."
|
||||
)
|
||||
|
||||
output_size_per_partition = sum(output_partition_sizes)
|
||||
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||
|
||||
layer.logical_widths = output_partition_sizes
|
||||
layer.input_size_per_partition = input_size_per_partition
|
||||
layer.output_size_per_partition = output_size_per_partition
|
||||
layer.params_dtype = params_dtype
|
||||
layer.quant_config = self.quant_config
|
||||
|
||||
if input_size_per_partition % 16 != 0:
|
||||
raise ValueError(
|
||||
"Unsupported model when input feature size is not a multiple of 16"
|
||||
)
|
||||
|
||||
weight = ModelWeightParameter(
|
||||
data=torch.empty(
|
||||
output_size_per_partition,
|
||||
input_size_per_partition // 2,
|
||||
dtype=torch.uint8,
|
||||
),
|
||||
input_dim=1,
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("weight", weight)
|
||||
|
||||
weight_scale_2 = _make_per_tensor_scale_parameter(
|
||||
(len(output_partition_sizes),),
|
||||
weight_loader=weight_loader,
|
||||
needs_scalar_to_array=True,
|
||||
)
|
||||
layer.register_parameter("weight_scale_2", weight_scale_2)
|
||||
|
||||
weight_scale = ModelWeightParameter(
|
||||
data=torch.empty(
|
||||
output_size_per_partition,
|
||||
input_size_per_partition // self.quant_config.group_size,
|
||||
dtype=torch.float8_e4m3fn,
|
||||
),
|
||||
input_dim=1,
|
||||
output_dim=0,
|
||||
weight_loader=weight_loader,
|
||||
)
|
||||
layer.register_parameter("weight_scale", weight_scale)
|
||||
|
||||
# Some ModelOpt checkpoints may still include input_scale entries in
|
||||
# fused-loader paths. NVFP4A16 does not use them, but registering the
|
||||
# placeholder lets the generic loader consume those tensors harmlessly.
|
||||
input_scale = _make_per_tensor_scale_parameter(
|
||||
(len(output_partition_sizes),),
|
||||
weight_loader=weight_loader,
|
||||
needs_scalar_to_array=True,
|
||||
)
|
||||
layer.register_parameter("input_scale", input_scale)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
if hasattr(layer, "input_scale"):
|
||||
del layer.input_scale
|
||||
|
||||
if torch.unique(layer.weight_scale_2).numel() != 1:
|
||||
logger.warning(
|
||||
"In NVFP4A16 linear, weight_scale_2 differs across fused "
|
||||
"parallel layers. Accuracy may be degraded."
|
||||
)
|
||||
|
||||
copy_or_rebind_param(
|
||||
layer,
|
||||
"weight_global_scale",
|
||||
layer.weight_scale_2.max().to(torch.float32),
|
||||
)
|
||||
del layer.weight_scale_2
|
||||
|
||||
if self.quant_config.group_size != 16:
|
||||
raise ValueError(
|
||||
f"NVFP4A16 Marlin requires group_size=16, got {self.quant_config.group_size}."
|
||||
)
|
||||
layer.quant_config = self.quant_config
|
||||
prepare_nvfp4_layer_for_marlin(layer)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
return apply_fp4_marlin_linear(
|
||||
input=x,
|
||||
weight=layer.weight,
|
||||
weight_scale=layer.weight_scale,
|
||||
weight_global_scale=layer.weight_global_scale,
|
||||
workspace=layer.workspace,
|
||||
size_n=layer.output_size_per_partition,
|
||||
size_k=layer.input_size_per_partition,
|
||||
bias=bias,
|
||||
)
|
||||
|
||||
|
||||
def _compute_gemm1_alphas(
|
||||
w13_weight_scale_2: torch.Tensor,
|
||||
w13_input_scale: torch.Tensor,
|
||||
|
||||
@@ -297,7 +297,7 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
# If we are making an embedding layer, then our quantization linear
|
||||
# method must implement the embedding operation. If we are another
|
||||
# layer type like ParallelLMHead, this is not important.
|
||||
is_embedding_layer = type(self.__class__) is VocabParallelEmbedding
|
||||
is_embedding_layer = type(self) is VocabParallelEmbedding
|
||||
quant_method_implements_embedding = method_has_implemented_embedding(
|
||||
type(quant_method)
|
||||
)
|
||||
@@ -459,6 +459,12 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
# If parameter does not have output dim, then it should
|
||||
# be copied onto all gpus (e.g. g_idx for act_order gptq).
|
||||
if output_dim is None:
|
||||
if (
|
||||
loaded_weight.ndim == 0
|
||||
and param.data.ndim == 1
|
||||
and param.data.numel() == 1
|
||||
):
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
assert param.data.shape == loaded_weight.shape
|
||||
param.data.copy_(loaded_weight)
|
||||
return
|
||||
|
||||
@@ -74,6 +74,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
RUNAI_STREAMER_TENSOR_ATTR = "_sglang_runai_streamer_tensor"
|
||||
|
||||
|
||||
# Matches routed-expert weight keys in both HF-style layouts
|
||||
# (``...mlp.experts.<N>.{gate,up,down}_proj.weight``) and DeepSeek V4
|
||||
# layouts (``...ffn.experts.<N>.w{1,2,3}.weight``). ``shared_experts`` is
|
||||
@@ -258,8 +259,24 @@ 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()
|
||||
hf_quant_config["packed_modules_mapping"] = packed_modules_mapping
|
||||
return quant_cls.from_config(hf_quant_config)
|
||||
|
||||
# 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
|
||||
# quantization metadata is missing.
|
||||
modelopt_mixed_config_incomplete = (
|
||||
model_config.quantization == "modelopt_mixed"
|
||||
and (
|
||||
"quantized_layers" not in hf_quant_config
|
||||
or (
|
||||
"kv_cache_quant_algo" not in hf_quant_config
|
||||
and "kv_cache_scheme" not in hf_quant_config
|
||||
)
|
||||
)
|
||||
)
|
||||
if not modelopt_mixed_config_incomplete:
|
||||
hf_quant_config["packed_modules_mapping"] = packed_modules_mapping
|
||||
return quant_cls.from_config(hf_quant_config)
|
||||
|
||||
# In case of bitsandbytes/QLoRA, get quant config from the adapter model.
|
||||
if model_config.quantization == "bitsandbytes":
|
||||
|
||||
@@ -59,7 +59,10 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
|
||||
config = copy.deepcopy(config)
|
||||
|
||||
# The MTP model is unquantized in the nvfp4 checkpoint.
|
||||
if quant_config and quant_config.get_name() == "modelopt_fp4":
|
||||
if quant_config and quant_config.get_name() in (
|
||||
"modelopt_fp4",
|
||||
"modelopt_mixed",
|
||||
):
|
||||
quant_config = None
|
||||
if (
|
||||
is_npu()
|
||||
@@ -135,6 +138,12 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
def set_lm_head_from_target(self, target_lm_head):
|
||||
if self.config.tie_word_embeddings:
|
||||
return
|
||||
|
||||
self.lm_head = target_lm_head
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
|
||||
@@ -316,6 +316,17 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
|
||||
def init_lm_head(self):
|
||||
embed, head = self.target_worker.model_runner.model.get_embed_and_head()
|
||||
target_lm_head = getattr(self.target_worker.model_runner.model, "lm_head", None)
|
||||
|
||||
def maybe_share_target_lm_head():
|
||||
if (
|
||||
target_lm_head is not None
|
||||
and self.hot_token_id is None
|
||||
and getattr(self.draft_runner.model, "hot_token_id", None) is None
|
||||
and hasattr(self.draft_runner.model, "set_lm_head_from_target")
|
||||
):
|
||||
self.draft_runner.model.set_lm_head_from_target(target_lm_head)
|
||||
|
||||
if self.speculative_algorithm.is_eagle3():
|
||||
# most cases EAGLE3 models don't share lm_head
|
||||
# but some models (e.g. nvidia/gpt-oss-120b-Eagle3) shares
|
||||
@@ -324,6 +335,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
and self.draft_runner.model.load_lm_head_from_target
|
||||
):
|
||||
self.draft_runner.model.set_embed_and_head(embed, head)
|
||||
maybe_share_target_lm_head()
|
||||
else:
|
||||
self.draft_runner.model.set_embed(embed)
|
||||
|
||||
@@ -341,6 +353,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
|
||||
# Share the embedding and lm_head
|
||||
self.draft_runner.model.set_embed_and_head(embed, head)
|
||||
maybe_share_target_lm_head()
|
||||
|
||||
def init_attention_backend(self):
|
||||
# Create multi-step attn backends and cuda graph runners
|
||||
|
||||
Reference in New Issue
Block a user