From e52019104d822c156a5b7dfb9c8a734bd897a4a3 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Fri, 29 Nov 2024 22:55:15 +0300 Subject: [PATCH] revert sd_models.py --- modules/sd_models.py | 14 +++++--------- 1 file changed, 5 insertions(+), 9 deletions(-) diff --git a/modules/sd_models.py b/modules/sd_models.py index 361f6375b..68446bdd3 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -405,7 +405,7 @@ def apply_balanced_offload(sd_model): if hasattr(pipe, "_internal_dict"): keys = pipe._internal_dict.keys() # pylint: disable=protected-access else: - keys = get_signature(pipe).keys() + keys = get_signature(shared.sd_model).keys() for module_name in keys: # pylint: disable=protected-access module = getattr(pipe, module_name, None) if isinstance(module, torch.nn.Module): @@ -1448,14 +1448,10 @@ def disable_offload(sd_model): from accelerate.hooks import remove_hook_from_module if not getattr(sd_model, 'has_accelerate', False): return - if hasattr(sd_model, "_internal_dict"): - keys = sd_model._internal_dict.keys() # pylint: disable=protected-access - else: - keys = get_signature(sd_model).keys() - for module_name in keys: # pylint: disable=protected-access - module = getattr(sd_model, module_name, None) - if isinstance(module, torch.nn.Module): - module = remove_hook_from_module(module, recurse=True) + if hasattr(sd_model, 'components'): + for _name, model in sd_model.components.items(): + if isinstance(model, torch.nn.Module): + remove_hook_from_module(model, recurse=True) sd_model.has_accelerate = False