Improve balanced offload pre forward performance

This commit is contained in:
Disty0
2024-11-30 20:49:57 +03:00
parent eee85e5a4e
commit b55e746fca
2 changed files with 15 additions and 11 deletions
+3 -5
View File
@@ -330,14 +330,12 @@ class StableCascadeDecoderPipelineFixed(diffusers.StableCascadeDecoderPipeline):
elif output_type == "pil":
images = images.permute(0, 2, 3, 1).cpu().float().numpy() # float() as bfloat16-> numpy doesnt work
images = self.numpy_to_pil(images)
if shared.opts.diffusers_offload_mode == "balanced":
shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model)
else:
images = latents
# Offload all models
if shared.opts.diffusers_offload_mode == "balanced":
shared.sd_model = sd_models.apply_balanced_offload(shared.sd_model)
else:
self.maybe_free_model_hooks()
self.maybe_free_model_hooks()
if not return_dict:
return images
+12 -6
View File
@@ -378,15 +378,17 @@ class OffloadHook(accelerate.hooks.ModelHook):
if device_index is None:
device_index = 0
max_memory = {
device_index: f"{shared.opts.diffusers_offload_max_gpu_memory}GiB",
"cpu": f"{shared.opts.diffusers_offload_max_cpu_memory}GiB",
device_index: int(shared.opts.diffusers_offload_max_gpu_memory * 1024*1024*1024),
"cpu": int(shared.opts.diffusers_offload_max_cpu_memory * 1024*1024*1024),
}
device_map = accelerate.infer_auto_device_map(module, max_memory=max_memory)
module = accelerate.hooks.remove_hook_from_module(module, recurse=True)
device_map = getattr(module, "balanced_offload_device_map", None)
if device_map is None or max_memory != getattr(module, "balanced_offload_max_memory", None):
device_map = accelerate.infer_auto_device_map(module, max_memory=max_memory)
offload_dir = getattr(module, "offload_dir", os.path.join(shared.opts.accelerate_offload_path, module.__class__.__name__))
module = accelerate.dispatch_model(module, device_map=device_map, offload_dir=offload_dir)
module = accelerate.hooks.add_hook_to_module(module, OffloadHook(), append=True)
module._hf_hook.execution_device = torch.device(devices.device) # pylint: disable=protected-access
module.balanced_offload_device_map = device_map
module.balanced_offload_max_memory = max_memory
return args, kwargs
def post_forward(self, module, output):
@@ -421,15 +423,19 @@ def apply_balanced_offload(sd_model):
module = getattr(pipe, module_name, None)
if isinstance(module, torch.nn.Module):
network_layer_name = getattr(module, "network_layer_name", None)
device_map = getattr(module, "balanced_offload_device_map", None)
max_memory = getattr(module, "balanced_offload_max_memory", None)
module = accelerate.hooks.remove_hook_from_module(module, recurse=True)
try:
module = module.to(devices.cpu, non_blocking=True)
module.offload_dir = os.path.join(shared.opts.accelerate_offload_path, checkpoint_name, module_name)
# module = accelerate.hooks.add_hook_to_module(module, OffloadHook(), append=True)
module = accelerate.hooks.add_hook_to_module(module, offload_hook_instance, append=True)
module._hf_hook.execution_device = torch.device(devices.device) # pylint: disable=protected-access
if network_layer_name:
module.network_layer_name = network_layer_name
if device_map and max_memory:
module.balanced_offload_device_map = device_map
module.balanced_offload_max_memory = max_memory
except Exception as e:
if 'bitsandbytes' not in str(e):
shared.log.error(f'Balanced offload: module={module_name} {e}')