diff --git a/python/sglang/srt/configs/nemotron_h.py b/python/sglang/srt/configs/nemotron_h.py index ead16aeae49f..3132c65f0e4e 100644 --- a/python/sglang/srt/configs/nemotron_h.py +++ b/python/sglang/srt/configs/nemotron_h.py @@ -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" ] diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 5b60bb40ad1a..e33d155868d6 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -12,6 +12,7 @@ ) from sglang.srt.configs.hybrid_arch import ( hybrid_gdn_config, + hybrid_lightning_config, kimi_linear_config, linear_attn_model_spec, ) @@ -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 ): diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py index f9d3e2ea0203..a848298e9e08 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py @@ -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): @@ -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 @@ -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 @@ -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 @@ -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 @@ -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( @@ -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 @@ -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: @@ -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() @@ -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) @@ -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") @@ -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 @@ -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 @@ -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 @@ -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) @@ -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, diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py index cc3624a7b5ef..2a1e0ee4a934 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_mxint4_moe.py @@ -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 @@ -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" diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py index 84bd4ae8fffb..40e99923366a 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a4_nvfp4_moe.py @@ -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, diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a8_fp8_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a8_fp8_moe.py index ef8d56abe13d..7463f18071c7 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a8_fp8_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a8_fp8_moe.py @@ -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 @@ -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: diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py index cbdfe11446b8..534b696412e9 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py @@ -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 @@ -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( diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/utils.py b/python/sglang/srt/layers/quantization/compressed_tensors/utils.py index d6c2ca3d208b..5dbb100a664f 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/utils.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/utils.py @@ -25,6 +25,8 @@ def should_ignore_layer( layer_name: Optional[str], ignore: Iterable[str] = tuple(), fused_mapping: Mapping[str, List[str]] = MappingProxyType({}), + *, + check_contains: bool = True, ) -> bool: if layer_name is None: return False @@ -50,7 +52,9 @@ def should_ignore_layer( should_ignore_layer = None for shard_name in shard_names: should_ignore_shard = check_equal_or_regex_match( - layer_name=shard_name, targets=ignore + layer_name=shard_name, + targets=ignore, + check_contains=check_contains, ) # If shard_idx=0, set layer ignore to match shard. @@ -69,22 +73,34 @@ def should_ignore_layer( # the safetensors checkpoint already. else: should_ignore_layer = check_equal_or_regex_match( - layer_name=layer_name, targets=ignore + layer_name=layer_name, + targets=ignore, + check_contains=check_contains, ) assert should_ignore_layer is not None return should_ignore_layer -def check_equal_or_regex_match(layer_name: str, targets: Iterable[str]) -> bool: +def check_equal_or_regex_match( + layer_name: str, + targets: Iterable[str], + *, + check_contains: bool = True, +) -> bool: """ - Checks whether a layer_name is exactly equal or a regex match for - if target starts with 're:' to any target in list. + Checks whether a layer_name is exactly equal to a target or matches a + target prefixed with ``re:``. When ``check_contains`` is true, plain + targets also match as case-insensitive substrings. """ - for target in targets: - if _is_equal_or_regex_match(layer_name, target, check_contains=True): - return True - return False + return any( + _is_equal_or_regex_match( + layer_name, + target, + check_contains=check_contains, + ) + for target in targets + ) def find_matched_target( diff --git a/test/registered/unit/configs/test_nemotron_h_puzzle_config.py b/test/registered/unit/configs/test_nemotron_h_puzzle_config.py new file mode 100644 index 000000000000..31bd7874b9b9 --- /dev/null +++ b/test/registered/unit/configs/test_nemotron_h_puzzle_config.py @@ -0,0 +1,39 @@ +import unittest + +from sglang.srt.configs.nemotron_h import NemotronHPuzzleConfig +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestNemotronHPuzzleConfig(unittest.TestCase): + def make_config(self, block_configs, n_routed_experts=512): + config = object.__new__(NemotronHPuzzleConfig) + config.block_configs = block_configs + config.n_routed_experts = n_routed_experts + return config + + def test_max_experts_falls_back_to_global_config(self): + config = self.make_config( + [ + {"block_type": "mamba"}, + {"block_type": "moe"}, + {"block_type": "moe", "n_routed_experts": 256}, + ] + ) + + self.assertEqual(config.max_n_routed_experts, 512) + + def test_max_experts_preserves_per_block_overrides(self): + config = self.make_config( + [ + {"block_type": "moe", "n_routed_experts": 128}, + {"block_type": "moe", "n_routed_experts": 768}, + ] + ) + + self.assertEqual(config.max_n_routed_experts, 768) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/layers/quantization/test_compressed_tensors_mixed_precision.py b/test/registered/unit/layers/quantization/test_compressed_tensors_mixed_precision.py new file mode 100644 index 000000000000..4e8d8380b059 --- /dev/null +++ b/test/registered/unit/layers/quantization/test_compressed_tensors_mixed_precision.py @@ -0,0 +1,259 @@ +"""CPU regressions for mixed-precision compressed-tensors configs.""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +import unittest +from unittest.mock import Mock, patch + +import torch + +from sglang.srt.layers.quantization.compressed_tensors.compressed_tensors import ( + CompressedTensorsConfig, + CompressedTensorsFusedMoEMethod, + CompressedTensorsLinearMethod, +) +from sglang.srt.layers.quantization.compressed_tensors.schemes import ( + CompressedTensorsW4A4Fp4, + CompressedTensorsW4A4Nvfp4MoE, + CompressedTensorsW4AFP8MoE, + CompressedTensorsW8A8Fp8, + CompressedTensorsWNA16, + CompressedTensorsWNA16MoE, +) +from sglang.srt.layers.quantization.compressed_tensors.utils import ( + check_equal_or_regex_match, + should_ignore_layer, +) +from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead +from sglang.test.test_utils import CustomTestCase + +FP8_TARGET = "re:.*(q_proj|lm_head)$" +NVFP4_TARGET = "re:.*mlp.experts.*" + + +def _mixed_precision_config(): + return { + "quant_method": "compressed-tensors", + "format": "mixed-precision", + "config_groups": { + "fp8": { + "format": "float-quantized", + "targets": [FP8_TARGET], + "weights": { + "num_bits": 8, + "type": "float", + "symmetric": True, + "strategy": "channel", + "dynamic": False, + }, + "input_activations": { + "num_bits": 8, + "type": "float", + "symmetric": True, + "strategy": "token", + "dynamic": True, + }, + }, + "nvfp4": { + "format": "nvfp4-pack-quantized", + "targets": [NVFP4_TARGET], + "weights": { + "num_bits": 4, + "type": "float", + "symmetric": True, + "strategy": "tensor_group", + "group_size": 16, + "dynamic": False, + }, + "input_activations": { + "num_bits": 4, + "type": "float", + "symmetric": True, + "strategy": "tensor_group", + "group_size": 16, + "dynamic": True, + }, + }, + }, + "ignore": [], + } + + +class TestCompressedTensorsMixedPrecision(CustomTestCase): + def test_nvfp4_moe_uses_backend_specific_w13_load_order(self): + scheme = CompressedTensorsW4A4Nvfp4MoE.__new__( + CompressedTensorsW4A4Nvfp4MoE + ) + + scheme.use_flashinfer_trtllm = False + self.assertTrue(scheme.load_up_proj_weight_first) + + scheme.use_flashinfer_trtllm = True + self.assertFalse(scheme.load_up_proj_weight_first) + + def test_fused_moe_method_exposes_layer_w13_load_order(self): + method = CompressedTensorsFusedMoEMethod(Mock()) + scheme = Mock(load_up_proj_weight_first=True) + layer = Mock(scheme=scheme) + + method.create_weights( + layer=layer, + num_experts=2, + hidden_size=16, + intermediate_size_per_partition=8, + params_dtype=torch.bfloat16, + ) + + self.assertTrue(method.load_up_proj_weight_first) + scheme.create_weights.assert_called_once() + + def test_parses_per_group_activation_quantization(self): + quant_config = CompressedTensorsConfig.from_config(_mixed_precision_config()) + + fp8_input = quant_config.target_scheme_map[FP8_TARGET]["input_activations"] + nvfp4_input = quant_config.target_scheme_map[NVFP4_TARGET]["input_activations"] + + self.assertIsNotNone(fp8_input) + self.assertEqual(fp8_input.num_bits, 8) + self.assertIsNotNone(nvfp4_input) + self.assertEqual(nvfp4_input.num_bits, 4) + + def test_selects_activation_scheme_for_mixed_precision_groups(self): + quant_config = CompressedTensorsConfig.from_config(_mixed_precision_config()) + + with patch.object(quant_config, "_check_scheme_supported", return_value=True): + fp8 = quant_config.target_scheme_map[FP8_TARGET] + nvfp4 = quant_config.target_scheme_map[NVFP4_TARGET] + + self.assertIsInstance( + quant_config._get_scheme_from_parts( + fp8["weights"], fp8["input_activations"] + ), + CompressedTensorsW8A8Fp8, + ) + self.assertIsInstance( + quant_config._get_scheme_from_parts( + nvfp4["weights"], nvfp4["input_activations"] + ), + CompressedTensorsW4A4Fp4, + ) + + def test_keeps_weight_only_pack_quantized_groups_valid(self): + config = _mixed_precision_config() + config["config_groups"] = { + "weight_only": { + "format": "pack-quantized", + "targets": ["Linear"], + "weights": { + "num_bits": 4, + "type": "int", + "symmetric": True, + "strategy": "group", + "group_size": 128, + "dynamic": False, + }, + "input_activations": None, + } + } + + quant_config = CompressedTensorsConfig.from_config(config) + scheme_dict = quant_config.target_scheme_map["Linear"] + + self.assertIsNone(scheme_dict["input_activations"]) + self.assertEqual(scheme_dict["format"], "pack-quantized") + with patch.object(quant_config, "_check_scheme_supported", return_value=True): + scheme = quant_config.get_linear_scheme( + torch.nn.Linear(1, 1), layer_name="model.linear" + ) + self.assertIsInstance(scheme, CompressedTensorsWNA16) + + def test_selects_weight_only_pack_quantized_moe_group(self): + config = _mixed_precision_config() + config["config_groups"] = { + "weight_only": { + "format": "pack-quantized", + "targets": [r"re:.*mlp\.experts.*"], + "weights": { + "num_bits": 4, + "type": "int", + "symmetric": True, + "strategy": "group", + "group_size": 128, + "dynamic": False, + }, + "input_activations": None, + } + } + quant_config = CompressedTensorsConfig.from_config(config) + + scheme = quant_config.get_moe_scheme( + torch.nn.Module(), layer_name="model.layers.0.mlp.experts" + ) + + self.assertIsInstance(scheme, CompressedTensorsWNA16MoE) + + def test_selects_pack_quantized_w4afp8_moe_group(self): + config = _mixed_precision_config() + config["config_groups"] = { + "w4afp8": { + "format": "pack-quantized", + "targets": [r"re:.*mlp\.experts.*"], + "weights": { + "num_bits": 4, + "type": "int", + "symmetric": True, + "strategy": "group", + "group_size": 128, + "dynamic": False, + }, + "input_activations": { + "num_bits": 8, + "type": "float", + "symmetric": True, + "strategy": "token", + "dynamic": True, + }, + } + } + quant_config = CompressedTensorsConfig.from_config(config) + + scheme = quant_config.get_moe_scheme( + torch.nn.Module(), layer_name="model.layers.0.mlp.experts" + ) + + self.assertIsInstance(scheme, CompressedTensorsW4AFP8MoE) + + def test_ignore_matching_supports_exact_scoping(self): + parent = "model.layers.0.linear_attn" + child = f"{parent}.in_proj_qkv" + + self.assertTrue(check_equal_or_regex_match(parent, [parent])) + self.assertTrue(check_equal_or_regex_match(child, [parent])) + self.assertFalse( + check_equal_or_regex_match(child, [parent], check_contains=False) + ) + self.assertTrue(check_equal_or_regex_match(child, [r"re:.*in_proj_qkv$"])) + self.assertTrue(should_ignore_layer(child, ignore=[parent])) + self.assertFalse( + should_ignore_layer(child, ignore=[parent], check_contains=False) + ) + + def test_quantizes_parallel_lm_head_when_targeted(self): + quant_config = CompressedTensorsConfig.from_config(_mixed_precision_config()) + layer = Mock(spec=ParallelLMHead) + scheme = object() + + with patch.object( + quant_config, "get_linear_scheme", return_value=scheme + ) as get_scheme: + method = quant_config.get_quant_method(layer, prefix="model.lm_head") + + self.assertIsInstance(method, CompressedTensorsLinearMethod) + self.assertIs(layer.scheme, scheme) + get_scheme.assert_called_once_with(layer=layer, layer_name="model.lm_head") + + +if __name__ == "__main__": + unittest.main()