diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index c195e2b26..d492edcbe 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -191,10 +191,33 @@ logger = logging.getLogger(__name__) def _get_quantization_config( model_config: ModelConfig, load_config: LoadConfig, - packed_modules_mapping: Dict[str, List[str]], - remap_prefix: Dict[str, str] | None = None, ) -> Optional[QuantizationConfig]: """Get the quantization config.""" + model_class, _ = get_model_architecture(model_config) + packed_modules_mapping = getattr(model_class, "packed_modules_mapping", {}) + remap_prefix = getattr(model_class, "remap_prefix", None) + if _is_npu: + packed_modules_mapping.update( + { + "visual": { + "qkv_proj": ["qkv"], + "gate_up_proj": ["gate_proj", "up_proj"], + }, + "vision_model": { + "qkv_proj": ["q_proj", "k_proj", "v_proj"], + "proj": ["out_proj"], + }, + "model": { + "qkv_proj": ["q_proj", "k_proj", "v_proj"], + "gate_up_proj": ["gate_proj", "up_proj"], + "fused_qkv_a_proj_with_mqa": [ + "q_a_proj", + "kv_a_proj_with_mqa", + ], + }, + } + ) + if model_config.quantization is not None: quant_config = get_quant_config( model_config, load_config, packed_modules_mapping, remap_prefix @@ -222,6 +245,10 @@ def _get_quantization_config( f"method {model_config.quantization}. Supported dtypes: " f"{supported_dtypes}" ) + hf_to_sglang_mapper = getattr(model_class, "hf_to_sglang_mapper", None) + # pass mappings by reference to quant_config + if hf_to_sglang_mapper is not None and quant_config is not None: + quant_config.apply_weight_name_mapper(hf_to_sglang_mapper) return quant_config return None @@ -229,42 +256,10 @@ def _get_quantization_config( def _initialize_model( model_config: ModelConfig, load_config: LoadConfig, + quant_config: Optional[QuantizationConfig] = None, ) -> nn.Module: """Initialize a model with the given configurations.""" model_class, _ = get_model_architecture(model_config) - packed_modules_mapping = getattr(model_class, "packed_modules_mapping", {}) - remap_prefix = getattr(model_class, "remap_prefix", None) - if _is_npu: - packed_modules_mapping.update( - { - "visual": { - "qkv_proj": ["qkv"], - "gate_up_proj": ["gate_proj", "up_proj"], - }, - "vision_model": { - "qkv_proj": ["q_proj", "k_proj", "v_proj"], - "proj": ["out_proj"], - }, - "model": { - "qkv_proj": ["q_proj", "k_proj", "v_proj"], - "gate_up_proj": ["gate_proj", "up_proj"], - "fused_qkv_a_proj_with_mqa": [ - "q_a_proj", - "kv_a_proj_with_mqa", - ], - }, - } - ) - - quant_config = _get_quantization_config( - model_config, load_config, packed_modules_mapping, remap_prefix - ) - hf_to_sglang_mapper = getattr(model_class, "hf_to_sglang_mapper", None) - # pass mappings by reference to quant_config - if hf_to_sglang_mapper is not None and quant_config is not None: - quant_config.apply_weight_name_mapper(hf_to_sglang_mapper) - - # Build kwargs conditionally kwargs = { "config": model_config.hf_config, "quant_config": quant_config, @@ -661,11 +656,13 @@ class DefaultModelLoader(BaseModelLoader): return model.eval() target_device = torch.device(device_config.device) + quant_config = _get_quantization_config(model_config, self.load_config) with set_default_torch_dtype(model_config.dtype): with target_device: model = _initialize_model( model_config, self.load_config, + quant_config, ) self.load_weights_and_postprocess( @@ -717,6 +714,7 @@ class LayeredModelLoader(DefaultModelLoader): torchao_config = get_global_server_args().torchao_config target_device = torch.device(device_config.device) + quant_config = _get_quantization_config(model_config, self.load_config) with set_default_torch_dtype(model_config.dtype): # Create model on meta device @@ -724,6 +722,7 @@ class LayeredModelLoader(DefaultModelLoader): model = _initialize_model( model_config, self.load_config, + quant_config, ) # Check model's layered load support @@ -1268,11 +1267,14 @@ class DummyModelLoader(BaseModelLoader): self, model_config=model_config, device_config=device_config ) + quant_config = _get_quantization_config(model_config, self.load_config) + with set_default_torch_dtype(model_config.dtype): with torch.device(device_config.device): model = _initialize_model( model_config, self.load_config, + quant_config, ) for _, module in model.named_modules(): @@ -1386,9 +1388,11 @@ class ShardedStateLoader(BaseModelLoader): model_config.model_path, model_config.revision ) + quant_config = _get_quantization_config(model_config, self.load_config) + with set_default_torch_dtype(model_config.dtype): with torch.device(device_config.device): - model = _initialize_model(model_config, self.load_config) + model = _initialize_model(model_config, self.load_config, quant_config) for _, module in model.named_modules(): quant_method = getattr(module, "quant_method", None) if quant_method is not None: @@ -1938,11 +1942,13 @@ class BitsAndBytesModelLoader(BaseModelLoader): model_config: ModelConfig, device_config: DeviceConfig, ) -> nn.Module: + quant_config = _get_quantization_config(model_config, self.load_config) with set_default_torch_dtype(model_config.dtype): with torch.device(device_config.device): model = _initialize_model( model_config, self.load_config, + quant_config, ) self._load_weights(model_config, model) @@ -2040,9 +2046,10 @@ class GGUFModelLoader(BaseModelLoader): model_config.hf_config.update({"tie_word_embeddings": True}) target_device = torch.device(device_config.device) + quant_config = _get_quantization_config(model_config, self.load_config) with set_default_torch_dtype(model_config.dtype): with target_device: - model = _initialize_model(model_config, self.load_config) + model = _initialize_model(model_config, self.load_config, quant_config) model.load_weights( self._get_weights_iterator(local_model_path, gguf_weights_map) ) @@ -2084,9 +2091,10 @@ class RemoteInstanceModelLoader(BaseModelLoader): f"load format {load_config.load_format}" ) + quant_config = _get_quantization_config(model_config, self.load_config) with set_default_torch_dtype(model_config.dtype): with torch.device(device_config.device): - model = _initialize_model(model_config, self.load_config) + model = _initialize_model(model_config, self.load_config, quant_config) if ( load_config.remote_instance_weight_loader_backend @@ -2374,9 +2382,11 @@ class RemoteModelLoader(BaseModelLoader): if hasattr(model_config, "model_weights"): model_weights = model_config.model_weights + quant_config = _get_quantization_config(model_config, self.load_config) + with set_default_torch_dtype(model_config.dtype): with torch.device(device_config.device): - model = _initialize_model(model_config, self.load_config) + model = _initialize_model(model_config, self.load_config, quant_config) with create_remote_connector( model_weights, device=device_config.device @@ -2401,10 +2411,12 @@ def load_model_with_cpu_quantization( device_config: DeviceConfig, ) -> nn.Module: target_device = torch.device(device_config.device) + quant_config = _get_quantization_config(model_config, self.load_config) with set_default_torch_dtype(model_config.dtype): model = _initialize_model( model_config, self.load_config, + quant_config, ) if not isinstance(self, DummyModelLoader):