mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 09:14:35 +02:00
Remove unset_config_on_save
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user