Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 8 additions & 4 deletions python/sglang/srt/configs/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -1146,10 +1146,14 @@ def _parse_modelopt_quant_config(self, quant_config_dict: dict) -> Optional[dict
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")
Comment thread
mmangkad marked this conversation as resolved.
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):
Expand Down
71 changes: 71 additions & 0 deletions python/sglang/srt/layers/logits_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,74 @@
_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:
Comment thread
mmangkad marked this conversation as resolved.
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 --
Expand Down Expand Up @@ -883,9 +951,12 @@ def _compute_lm_head(
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:
Expand Down
Loading
Loading