Pass quantize_config to _initialize_model (#18273)

This commit is contained in:
Kurt Shuster
2026-02-09 07:34:42 -08:00
committed by GitHub
parent ddbcfbaaab
commit 006da22268

View File

@@ -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):