diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index a08eec192..6fc34551e 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -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: