Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
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
2 changes: 1 addition & 1 deletion python/sglang/srt/configs/nemotron_h.py
Original file line number Diff line number Diff line change
Expand Up @@ -550,7 +550,7 @@ def get_mtp_config(self) -> NemotronHConfig:
@property
def max_n_routed_experts(self) -> int:
block_n_routed_experts = [
block["n_routed_experts"]
block.get("n_routed_experts", self.n_routed_experts)
for block in self.block_configs
if block["block_type"] == "moe"
]
Expand Down
2 changes: 2 additions & 0 deletions python/sglang/srt/layers/attention/triton_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
)
from sglang.srt.configs.hybrid_arch import (
hybrid_gdn_config,
hybrid_lightning_config,
kimi_linear_config,
linear_attn_model_spec,
)
Expand Down Expand Up @@ -188,6 +189,7 @@ def __init__(
self.swa_v_head_dim = swa_v_head_dim
elif (
hybrid_gdn_config(model_runner.model_config) is not None
or hybrid_lightning_config(model_runner.model_config) is not None
or kimi_linear_config(model_runner.model_config) is not None
or linear_attn_model_spec(model_runner.model_config) is not None
):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -85,7 +85,7 @@
__all__ = ["CompressedTensorsLinearMethod"]

SPARSITY_CONFIG_NAME: Literal["sparsity_config"] = "sparsity_config"
QUANTIZATION_SCHEME_MAP_TYPE = Dict[str, Optional[Dict[str, QuantizationArgs]]]
QUANTIZATION_SCHEME_MAP_TYPE = Dict[str, Optional[Dict[str, Any]]]


class DeviceCapability(NamedTuple):
Expand Down Expand Up @@ -164,6 +164,16 @@ def get_quant_method(
prefix: str,
) -> Optional[QuantizeMethodBase]:
from sglang.srt.layers.linear import LinearBase
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead

if isinstance(layer, ParallelLMHead):
try:
scheme = self.get_linear_scheme(layer=layer, layer_name=prefix)
except ValueError:
scheme = None
if scheme is not None:
layer.scheme = scheme
return CompressedTensorsLinearMethod(self)

if isinstance(layer, LinearBase):
# If linear_fp8_config is set, use FP8 for linear layers
Expand Down Expand Up @@ -305,8 +315,17 @@ def _quantization_scheme_map_from_config(
)

target_scheme_map[target]["input_activations"] = None
if is_activation_quantization_format(quant_format):
input_activations = quant_config.get("input_activations")
group_format = quant_config.get("format")
target_scheme_map[target]["format"] = (
group_format if group_format is not None else quant_format
)
activation_quantized = (
is_activation_quantization_format(group_format)
if group_format is not None
else is_activation_quantization_format(quant_format)
)
input_activations = quant_config.get("input_activations")
if activation_quantized or input_activations:
# When activation quant format is set but no
# input_activations provided: valid for w8a16fp8 (FLOAT
# weights) and pack-quantized without activation quant
Expand Down Expand Up @@ -367,12 +386,18 @@ def _is_dynamic_token_w4a8(
and is_dynamic
)

def _is_wint4afp8(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool:
def _is_wint4afp8(
self,
weight_quant: BaseModel,
input_quant: BaseModel,
quant_format: Optional[str] = None,
) -> bool:
"""Detect W4AFP8: packed INT4 weights + 8-bit dynamic per-token activations."""
if weight_quant is None or input_quant is None:
return False
quant_format = quant_format or self.quant_format
return (
self.quant_format == CompressionFormat.pack_quantized.value
quant_format == CompressionFormat.pack_quantized.value
and weight_quant.num_bits == 4
and weight_quant.type == QuantizationType.INT
and weight_quant.symmetric
Expand All @@ -382,12 +407,18 @@ def _is_wint4afp8(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool
and input_quant.dynamic # currently not support static input scales
)

def _is_wint4abf16(self, weight_quant: BaseModel, input_quant: BaseModel) -> bool:
def _is_wint4abf16(
self,
weight_quant: BaseModel,
input_quant: BaseModel,
quant_format: Optional[str] = None,
) -> bool:
"""Detect W4A16: packed INT4 weights with no activation quantization (activations stay BF16)."""
if weight_quant is None or input_quant is not None:
return False
quant_format = quant_format or self.quant_format
return (
self.quant_format == CompressionFormat.pack_quantized.value
quant_format == CompressionFormat.pack_quantized.value
and weight_quant.num_bits == 4
and weight_quant.type == QuantizationType.INT
and weight_quant.symmetric
Expand Down Expand Up @@ -572,13 +603,17 @@ def _is_dynamic_token_w4(
return is_w4 and weight_quant.symmetric and is_token and is_dynamic

def _get_scheme_from_parts(
self, weight_quant: BaseModel, input_quant: BaseModel
self,
weight_quant: BaseModel,
input_quant: BaseModel,
quant_format: Optional[str] = None,
) -> CompressedTensorsLinearScheme:
quant_format = quant_format or self.quant_format

# Detect If Mixed Precision
if self._is_wNa16_group_channel(weight_quant, input_quant):
if (
self.quant_format == CompressionFormat.pack_quantized.value
quant_format == CompressionFormat.pack_quantized.value
and weight_quant.num_bits in WNA16_SUPPORTED_BITS
):
return CompressedTensorsWNA16(
Expand All @@ -593,7 +628,7 @@ def _get_scheme_from_parts(
"Other method (CompressedTensorsW4A16Sparse24) is not supported now"
)

if is_activation_quantization_format(self.quant_format):
if input_quant is not None or is_activation_quantization_format(quant_format):
if self._is_fp4a4_nvfp4(weight_quant, input_quant):
is_fp4a4_nvfp4_supported = self._check_scheme_supported(
CompressedTensorsW4A4Fp4.get_min_capability(), error=False
Expand Down Expand Up @@ -701,6 +736,7 @@ def get_moe_scheme(

weight_quant = scheme_dict.get("weights")
input_quant = scheme_dict.get("input_activations")
quant_format = scheme_dict.get("format") or self.quant_format

if self._is_wNa16_group_channel(weight_quant, input_quant):
if not _is_npu:
Expand All @@ -711,11 +747,17 @@ def get_moe_scheme(
logger.info_once(
"Using CompressedTensorsMxInt4MoE with flashinfer_trtllm backend"
)
return CompressedTensorsMxInt4MoE(self, weight_quant=weight_quant)
return CompressedTensorsMxInt4MoE(
self,
weight_quant=weight_quant,
quant_format=quant_format,
)
elif _is_hip:
logger.info_once("Using CompressedTensorsWNA16TritonMoE (ROCm)")
return CompressedTensorsWNA16TritonMoE(
self, weight_quant=weight_quant
self,
weight_quant=weight_quant,
quant_format=quant_format,
)
else:
moe_backend = get_moe_runner_backend()
Expand All @@ -725,10 +767,16 @@ def get_moe_scheme(
"(moe_runner_backend=triton)"
)
return CompressedTensorsWNA16TritonMoE(
self, weight_quant=weight_quant
self,
weight_quant=weight_quant,
quant_format=quant_format,
)
logger.info_once("Using CompressedTensorsWNA16MarlinMoEMethod")
return CompressedTensorsWNA16MoE(self, weight_quant=weight_quant)
return CompressedTensorsWNA16MoE(
self,
weight_quant=weight_quant,
quant_format=quant_format,
)
else:
if (
self._is_dynamic_token_w4(weight_quant, input_quant)
Expand All @@ -750,13 +798,18 @@ def get_moe_scheme(
raise NotImplementedError(
f"The W8A8Int8 Fused MoE scheme is implemented only for NPU for now."
)
elif self._is_wint4afp8(weight_quant, input_quant):
elif self._is_wint4afp8(weight_quant, input_quant, quant_format):
# On NPU prefer the dedicated NPU W4A8Int8 path when activations are INT8.
if _is_npu and self._is_dynamic_token_w4a8(weight_quant, input_quant):
logger.info_once("Using NPUCompressedTensorsW4A8Int8DynamicMoE")
return NPUCompressedTensorsW4A8Int8DynamicMoE(self)
logger.info_once("Using CompressedTensorsW4AFP8MoE")
return CompressedTensorsW4AFP8MoE(self, weight_quant, input_quant)
return CompressedTensorsW4AFP8MoE(
self,
weight_quant,
input_quant,
quant_format=quant_format,
)
elif self._is_dynamic_token_w4a8(weight_quant, input_quant):
if _is_npu:
logger.info_once("Using NPUCompressedTensorsW4A8Int8DynamicMoE")
Expand Down Expand Up @@ -796,9 +849,11 @@ def get_linear_scheme(
scheme_dict = self.get_scheme_dict(layer, layer_name)
weight_quant = None
input_quant = None
quant_format = self.quant_format
if scheme_dict:
weight_quant = scheme_dict.get("weights")
input_quant = scheme_dict.get("input_activations")
quant_format = scheme_dict.get("format") or self.quant_format

# Find the sparsity scheme of the layer
# assume that fused layers inerhit first component's sparsity scheme
Expand Down Expand Up @@ -834,6 +889,7 @@ def get_linear_scheme(
scheme = self._get_scheme_from_parts( # type: ignore
weight_quant=weight_quant,
input_quant=input_quant,
quant_format=quant_format,
)

# Raise error if device does not support the scheme
Expand All @@ -858,7 +914,10 @@ def get_scheme_dict(
} | None
"""
if should_ignore_layer(
layer_name, ignore=self.ignore, fused_mapping=self.packed_modules_mapping
layer_name,
ignore=self.ignore,
fused_mapping=self.packed_modules_mapping,
check_contains=False,
):
return None

Expand Down Expand Up @@ -1017,6 +1076,7 @@ class CompressedTensorsFusedMoEMethod(FusedMoEMethodBase):
def __init__(self, quantization_config: CompressedTensorsConfig):
self.quantization_config = quantization_config
self.quant_config = quantization_config
self.load_up_proj_weight_first = False

def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
layer.scheme.process_weights_after_loading(layer)
Expand All @@ -1035,6 +1095,12 @@ def create_weights(
the necessary parameters for the layer. See LinearMethodBase for param
details
"""
# FusedMoE's checkpoint loader reads this flag from the quant method,
# while compressed-tensors resolves the backend-specific contract on
# the per-layer scheme.
self.load_up_proj_weight_first = getattr(
layer.scheme, "load_up_proj_weight_first", False
)
layer.scheme.create_weights(
layer=layer,
num_experts=num_experts,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -49,9 +49,13 @@

class CompressedTensorsMxInt4MoE(CompressedTensorsMoEScheme):
def __init__(
self, quant_config: CompressedTensorsConfig, weight_quant: QuantizationArgs
self,
quant_config: CompressedTensorsConfig,
weight_quant: QuantizationArgs,
quant_format: str | None = None,
):
self.quant_config = quant_config
quant_format = quant_format or self.quant_config.quant_format
# Per-layer scheme already resolved by get_moe_scheme(); reuse it directly
# (mixed-precision MoE has no "Linear" config group to fall back on).
config = weight_quant
Expand All @@ -74,7 +78,7 @@ def __init__(
), "Actorder is not supported by flashinfer_trtllm backend"
self.moe_ep_rank = get_parallel().moe_ep_rank

if self.quant_config.quant_format != CompressionFormat.pack_quantized.value:
if quant_format != CompressionFormat.pack_quantized.value:
raise ValueError(
f"For Fused MoE layers, only {CompressionFormat.pack_quantized.value} "
"is supported for the mxint4"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,16 @@ def get_min_capability(cls) -> int:
# Requires sm100(blackwell) architecture
return 100

@property
def load_up_proj_weight_first(self) -> bool:
"""Use the W13 ordering required by the selected FlashInfer kernel.

FlashInfer CUTLASS consumes fused gated weights as ``[up, gate]``.
The TRT-LLM path consumes ``[gate, up]`` at load time and reorders the
tensors, including their block scales, during post-processing below.
"""
return not self.use_flashinfer_trtllm

def create_weights(
self,
layer: torch.nn.Module,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -82,9 +82,11 @@ def __init__(
quant_config: CompressedTensorsConfig,
weight_quant,
input_quant,
quant_format: str | None = None,
):
self.quant_config = quant_config
config = self.quant_config.target_scheme_map["Linear"].get("weights")
quant_format = quant_format or self.quant_config.quant_format
config = weight_quant
self.num_bits = config.num_bits
self.packed_factor = 32 // config.num_bits
self.group_size = config.group_size
Expand All @@ -93,8 +95,8 @@ def __init__(

assert config.symmetric, "Only symmetric quantization is supported"
assert (
self.quant_config.quant_format == CompressionFormat.pack_quantized.value
), f"W4AFP8MoE requires pack-quantized format, got {self.quant_config.quant_format}"
quant_format == CompressionFormat.pack_quantized.value
), f"W4AFP8MoE requires pack-quantized format, got {quant_format}"

@classmethod
def get_min_capability(cls) -> int:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -69,8 +69,10 @@ def __init__(
quant_config: CompressedTensorsConfig,
weight_quant: QuantizationArgs,
num_gpu_experts: int = -1,
quant_format: str | None = None,
):
self.quant_config = quant_config
quant_format = quant_format or self.quant_config.quant_format
# Per-layer scheme already resolved by get_moe_scheme(); reuse it directly
# (mixed-precision MoE has no "Linear" config group to fall back on).
config = weight_quant
Expand All @@ -82,7 +84,7 @@ def __init__(
self.sym = config.symmetric

if not (
self.quant_config.quant_format == CompressionFormat.pack_quantized.value
quant_format == CompressionFormat.pack_quantized.value
and self.num_bits in WNA16_SUPPORTED_BITS
):
raise ValueError(
Expand Down
Loading