diff options
Diffstat (limited to 'modules/sd_unet.py')
-rw-r--r-- | modules/sd_unet.py | 16 |
1 files changed, 9 insertions, 7 deletions
diff --git a/modules/sd_unet.py b/modules/sd_unet.py index 6d708ad2..a771849c 100644 --- a/modules/sd_unet.py +++ b/modules/sd_unet.py @@ -1,12 +1,11 @@ import torch.nn
-import ldm.modules.diffusionmodules.openaimodel
from modules import script_callbacks, shared, devices
unet_options = []
current_unet_option = None
current_unet = None
-
+original_forward = None # not used, only left temporarily for compatibility
def list_unets():
new_unets = script_callbacks.list_unets_callback()
@@ -47,7 +46,7 @@ def apply_unet(option=None): if current_unet_option is None:
current_unet = None
- if not (shared.cmd_opts.lowvram or shared.cmd_opts.medvram):
+ if not shared.sd_model.lowvram:
shared.sd_model.model.diffusion_model.to(devices.device)
return
@@ -84,9 +83,12 @@ class SdUnet(torch.nn.Module): pass
-def UNetModel_forward(self, x, timesteps=None, context=None, *args, **kwargs):
- if current_unet is not None:
- return current_unet.forward(x, timesteps, context, *args, **kwargs)
+def create_unet_forward(original_forward):
+ def UNetModel_forward(self, x, timesteps=None, context=None, *args, **kwargs):
+ if current_unet is not None:
+ return current_unet.forward(x, timesteps, context, *args, **kwargs)
+
+ return original_forward(self, x, timesteps, context, *args, **kwargs)
- return ldm.modules.diffusionmodules.openaimodel.copy_of_UNetModel_forward_for_webui(self, x, timesteps, context, *args, **kwargs)
+ return UNetModel_forward
|