From a4a68f33e37da74f16beca791963e01c07d25a85 Mon Sep 17 00:00:00 2001 From: Kubuxu Date: Sat, 29 Jul 2023 20:13:36 +0100 Subject: [PATCH] Fix model offload by not focring the model to GPU --- modules/sd_models.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/modules/sd_models.py b/modules/sd_models.py index 341083477..b0fe7d1a2 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -702,7 +702,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No base_sent_to_cpu=False if shared.opts.cuda_compile and torch.cuda.is_available(): - if op == 'refiner' and not shared.opts.diffusers_seq_cpu_offload: + if op == 'refiner' and not shared.opts.diffusers_seq_cpu_offload and not shared.opts.diffusers_model_cpu_offload: gpu_vram = memory_stats().get('gpu', {}) free_vram = gpu_vram.get('total', 0) - gpu_vram.get('used', 0) refiner_enough_vram = free_vram >= 7 if "StableDiffusionXL" in sd_model.__class__.__name__ else 3 @@ -721,7 +721,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No devices.torch_gc(force=True) sd_model.to(devices.device) base_sent_to_cpu=True - elif not shared.opts.diffusers_seq_cpu_offload: + elif not shared.opts.diffusers_seq_cpu_offload and not shared.opts.diffusers_model_cpu_offload: sd_model.to(devices.device) try: shared.log.info(f"Compiling pipeline={sd_model.__class__.__name__} shape={8 * sd_model.unet.config.sample_size} mode={shared.opts.cuda_compile_mode}") @@ -752,7 +752,8 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No if op == 'refiner' and shared.opts.diffusers_move_refiner: shared.log.debug('Moving refiner model to CPU') sd_model.to("cpu") - elif not shared.opts.diffusers_seq_cpu_offload: + elif not shared.opts.diffusers_seq_cpu_offload and not shared.opts.diffusers_model_cpu_offload: + # In offload modes, accelerate will move models around. sd_model.to(devices.device) if op == 'refiner' and base_sent_to_cpu: shared.log.debug('Moving base model back to GPU')