add offload pre-forward

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-08-01 10:49:03 -04:00
parent b291c337a1
commit fbad5d87d6
2 changed files with 6 additions and 7 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
# Change Log for SD.Next
## Update for 2025-07-31
## Update for 2025-08-01
- **Models**
- [FLUX.1-Krea-Dev](https://www.krea.ai/blog/flux-krea-open-source-release)
+5 -6
View File
@@ -234,11 +234,10 @@ class OffloadHook(accelerate.hooks.ModelHook):
errors.display(e, f'Offload: type=balanced op=apply module={module.__name__}')
def pre_forward(self, module, *args, **kwargs):
print('HERE', id(module), module.__class__.__name__)
if self.last_pre != module.__class__.__name__: # offload every other module first time when new module starts pre-forward
self.last_pre = module.__class__.__name__
if self.last_pre != id(module): # offload every other module first time when new module starts pre-forward
self.last_pre = id(module)
if shared.opts.diffusers_offload_pre:
debug_move(f'Offload: type=balanced op=pre module={self.last_pre}')
debug_move(f'Offload: type=balanced op=pre module={module.__class__.__name__}')
for pipe in get_pipe_variants():
for module_name in get_module_names(pipe):
module_instance = getattr(pipe, module_name, None)
@@ -268,8 +267,8 @@ class OffloadHook(accelerate.hooks.ModelHook):
return args, kwargs
def post_forward(self, module, output):
if self.last_post != module.__class__.__name__:
self.last_post = module.__class__.__name__
if self.last_post != id(module):
self.last_post = id(module)
if getattr(module, "offload_post", False) and (module.device != devices.cpu):
self.offload_module(module, op='post')
return output