Pass quantize_config to _initialize_model (#18273)
This commit is contained in:
@@ -191,10 +191,33 @@ logger = logging.getLogger(__name__)
|
|||||||
def _get_quantization_config(
|
def _get_quantization_config(
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
load_config: LoadConfig,
|
load_config: LoadConfig,
|
||||||
packed_modules_mapping: Dict[str, List[str]],
|
|
||||||
remap_prefix: Dict[str, str] | None = None,
|
|
||||||
) -> Optional[QuantizationConfig]:
|
) -> Optional[QuantizationConfig]:
|
||||||
"""Get the quantization config."""
|
"""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:
|
if model_config.quantization is not None:
|
||||||
quant_config = get_quant_config(
|
quant_config = get_quant_config(
|
||||||
model_config, load_config, packed_modules_mapping, remap_prefix
|
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"method {model_config.quantization}. Supported dtypes: "
|
||||||
f"{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 quant_config
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -229,42 +256,10 @@ def _get_quantization_config(
|
|||||||
def _initialize_model(
|
def _initialize_model(
|
||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
load_config: LoadConfig,
|
load_config: LoadConfig,
|
||||||
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
"""Initialize a model with the given configurations."""
|
"""Initialize a model with the given configurations."""
|
||||||
model_class, _ = get_model_architecture(model_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",
|
|
||||||
],
|
|
||||||
},
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
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 = {
|
kwargs = {
|
||||||
"config": model_config.hf_config,
|
"config": model_config.hf_config,
|
||||||
"quant_config": quant_config,
|
"quant_config": quant_config,
|
||||||
@@ -661,11 +656,13 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
return model.eval()
|
return model.eval()
|
||||||
|
|
||||||
target_device = torch.device(device_config.device)
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
with target_device:
|
with target_device:
|
||||||
model = _initialize_model(
|
model = _initialize_model(
|
||||||
model_config,
|
model_config,
|
||||||
self.load_config,
|
self.load_config,
|
||||||
|
quant_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.load_weights_and_postprocess(
|
self.load_weights_and_postprocess(
|
||||||
@@ -717,6 +714,7 @@ class LayeredModelLoader(DefaultModelLoader):
|
|||||||
|
|
||||||
torchao_config = get_global_server_args().torchao_config
|
torchao_config = get_global_server_args().torchao_config
|
||||||
target_device = torch.device(device_config.device)
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
# Create model on meta device
|
# Create model on meta device
|
||||||
@@ -724,6 +722,7 @@ class LayeredModelLoader(DefaultModelLoader):
|
|||||||
model = _initialize_model(
|
model = _initialize_model(
|
||||||
model_config,
|
model_config,
|
||||||
self.load_config,
|
self.load_config,
|
||||||
|
quant_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check model's layered load support
|
# Check model's layered load support
|
||||||
@@ -1268,11 +1267,14 @@ class DummyModelLoader(BaseModelLoader):
|
|||||||
self, model_config=model_config, device_config=device_config
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
with torch.device(device_config.device):
|
with torch.device(device_config.device):
|
||||||
model = _initialize_model(
|
model = _initialize_model(
|
||||||
model_config,
|
model_config,
|
||||||
self.load_config,
|
self.load_config,
|
||||||
|
quant_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
for _, module in model.named_modules():
|
for _, module in model.named_modules():
|
||||||
@@ -1386,9 +1388,11 @@ class ShardedStateLoader(BaseModelLoader):
|
|||||||
model_config.model_path, model_config.revision
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
with torch.device(device_config.device):
|
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():
|
for _, module in model.named_modules():
|
||||||
quant_method = getattr(module, "quant_method", None)
|
quant_method = getattr(module, "quant_method", None)
|
||||||
if quant_method is not None:
|
if quant_method is not None:
|
||||||
@@ -1938,11 +1942,13 @@ class BitsAndBytesModelLoader(BaseModelLoader):
|
|||||||
model_config: ModelConfig,
|
model_config: ModelConfig,
|
||||||
device_config: DeviceConfig,
|
device_config: DeviceConfig,
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
|
quant_config = _get_quantization_config(model_config, self.load_config)
|
||||||
with set_default_torch_dtype(model_config.dtype):
|
with set_default_torch_dtype(model_config.dtype):
|
||||||
with torch.device(device_config.device):
|
with torch.device(device_config.device):
|
||||||
model = _initialize_model(
|
model = _initialize_model(
|
||||||
model_config,
|
model_config,
|
||||||
self.load_config,
|
self.load_config,
|
||||||
|
quant_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._load_weights(model_config, model)
|
self._load_weights(model_config, model)
|
||||||
@@ -2040,9 +2046,10 @@ class GGUFModelLoader(BaseModelLoader):
|
|||||||
model_config.hf_config.update({"tie_word_embeddings": True})
|
model_config.hf_config.update({"tie_word_embeddings": True})
|
||||||
|
|
||||||
target_device = torch.device(device_config.device)
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
with target_device:
|
with target_device:
|
||||||
model = _initialize_model(model_config, self.load_config)
|
model = _initialize_model(model_config, self.load_config, quant_config)
|
||||||
model.load_weights(
|
model.load_weights(
|
||||||
self._get_weights_iterator(local_model_path, gguf_weights_map)
|
self._get_weights_iterator(local_model_path, gguf_weights_map)
|
||||||
)
|
)
|
||||||
@@ -2084,9 +2091,10 @@ class RemoteInstanceModelLoader(BaseModelLoader):
|
|||||||
f"load format {load_config.load_format}"
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
with torch.device(device_config.device):
|
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 (
|
if (
|
||||||
load_config.remote_instance_weight_loader_backend
|
load_config.remote_instance_weight_loader_backend
|
||||||
@@ -2374,9 +2382,11 @@ class RemoteModelLoader(BaseModelLoader):
|
|||||||
if hasattr(model_config, "model_weights"):
|
if hasattr(model_config, "model_weights"):
|
||||||
model_weights = 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 set_default_torch_dtype(model_config.dtype):
|
||||||
with torch.device(device_config.device):
|
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(
|
with create_remote_connector(
|
||||||
model_weights, device=device_config.device
|
model_weights, device=device_config.device
|
||||||
@@ -2401,10 +2411,12 @@ def load_model_with_cpu_quantization(
|
|||||||
device_config: DeviceConfig,
|
device_config: DeviceConfig,
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
target_device = torch.device(device_config.device)
|
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 set_default_torch_dtype(model_config.dtype):
|
||||||
model = _initialize_model(
|
model = _initialize_model(
|
||||||
model_config,
|
model_config,
|
||||||
self.load_config,
|
self.load_config,
|
||||||
|
quant_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not isinstance(self, DummyModelLoader):
|
if not isinstance(self, DummyModelLoader):
|
||||||
|
|||||||
Reference in New Issue
Block a user