Remove unset_config_on_save

This commit is contained in:
Disty0
2026-06-07 17:28:17 +03:00
parent 1247f871d3
commit 809f24aa13
-22
View File
@@ -19,14 +19,6 @@ def get_module_names(model: ModelMixin) -> list:
return modules_names
def unset_config_on_save(quantization_config: SDNQConfig) -> SDNQConfig:
quantization_config.quantization_device = None
quantization_config.return_device = None
quantization_config.non_blocking = False
quantization_config.add_skip_keys = False
return quantization_config
def normalize_tied_weights_keys_for_save(model: ModelMixin, is_pipeline: bool = False) -> list[tuple[torch.nn.Module, object]]:
normalized_modules = []
modules_to_walk = []
@@ -53,19 +45,6 @@ def restore_tied_weights_keys_after_save(normalized_modules: list[tuple[torch.nn
def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "5GB", is_pipeline: bool = False, sdnq_config: SDNQConfig | None = None) -> None:
if is_pipeline:
for module_name in get_module_names(model):
module = getattr(model, module_name, None)
if hasattr(module, "config") and hasattr(module.config, "quantization_config") and isinstance(module.config.quantization_config, SDNQConfig):
module.config.quantization_config = unset_config_on_save(module.config.quantization_config)
if hasattr(module, "quantization_config") and isinstance(module.quantization_config, SDNQConfig):
module.quantization_config = unset_config_on_save(module.quantization_config)
else:
if hasattr(model, "config") and hasattr(model.config, "quantization_config") and isinstance(model.config.quantization_config, SDNQConfig):
model.config.quantization_config = unset_config_on_save(model.config.quantization_config)
if hasattr(model, "quantization_config") and isinstance(model.quantization_config, SDNQConfig):
model.quantization_config = unset_config_on_save(model.quantization_config)
normalized_modules = normalize_tied_weights_keys_for_save(model, is_pipeline=is_pipeline)
try:
model.save_pretrained(model_path, max_shard_size=max_shard_size) # actual save
@@ -74,7 +53,6 @@ def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "5
quantization_config_path = os.path.join(model_path, "quantization_config.json")
if sdnq_config is not None: # if provided, save global config
sdnq_config = unset_config_on_save(sdnq_config)
sdnq_config.to_json_file(quantization_config_path)
if is_pipeline: