diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 57fc93734..88310ba6e 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -202,37 +202,6 @@ class OffloadHook(accelerate.hooks.ModelHook): def init_hook(self, module): return module - def offload_module(self, module, op='unk'): - try: - used_gpu, used_ram = devices.torch_gc(fast=True) - perc_gpu = used_gpu / shared.gpu_memory - prev_gpu = used_gpu - module_size = self.model_size() - module_cls = module.__class__.__name__ - op = f'{op}:skip' - if module_cls in self.offload_never: - op = f'{op}:never' - elif module_cls in self.offload_always: - op = f'{op}:always' - module = module.to(devices.cpu) - used_gpu -= module_size - elif perc_gpu > shared.opts.diffusers_offload_min_gpu_memory: - op = f'{op}:mem' - module = module.to(devices.cpu) - used_gpu -= module_size - if debug: - quant = getattr(module, "quantization_method", None) - debug_move(f'Offload: type=balanced op={op} gpu={prev_gpu:.3f}:{used_gpu:.3f} perc={perc_gpu:.2f} ram={used_ram:.3f} current={module.device} dtype={module.dtype} quant={quant} module={module_cls} size={module_size:.3f}') - except Exception as e: - if 'out of memory' in str(e): - devices.torch_gc(fast=True, force=True, reason='oom') - elif 'bitsandbytes' in str(e): - pass - else: - shared.log.error(f'Offload: type=balanced op=apply module={getattr(module, "__name__", None)} cls={module.__class__ if inspect.isclass(module) else None} {e}') - if os.environ.get('SD_MOVE_DEBUG', None): - errors.display(e, f'Offload: type=balanced op=apply module={getattr(module, "__name__", None)}') - def pre_forward(self, module, *args, **kwargs): if self.last_pre != id(module): # offload every other module first time when new module starts pre-forward self.last_pre = id(module) @@ -323,9 +292,41 @@ def get_module_sizes(pipe=None, exclude=[]): return modules -def apply_balanced_offload_to_module(module, module_name=None, checkpoint_name="", op="apply"): - if module_name is None: - module_name = module.__class__.__name__ +def move_module_to_cpu(module, op='unk'): + try: + module_name = getattr(module, "module_name", module.__class__.__name__) + module_size = offload_hook_instance.offload_map.get(module_name, offload_hook_instance.model_size()) + used_gpu, used_ram = devices.torch_gc(fast=True) + perc_gpu = used_gpu / shared.gpu_memory + prev_gpu = used_gpu + module_cls = module.__class__.__name__ + op = f'{op}:skip' + if module_cls in offload_hook_instance.offload_never: + op = f'{op}:never' + elif module_cls in offload_hook_instance.offload_always: + op = f'{op}:always' + module = module.to(devices.cpu) + used_gpu -= module_size + elif perc_gpu > shared.opts.diffusers_offload_min_gpu_memory: + op = f'{op}:mem' + module = module.to(devices.cpu) + used_gpu -= module_size + if debug: + quant = getattr(module, "quantization_method", None) + debug_move(f'Offload: type=balanced op={op} gpu={prev_gpu:.3f}:{used_gpu:.3f} perc={perc_gpu:.2f} ram={used_ram:.3f} current={module.device} dtype={module.dtype} quant={quant} module={module_cls} size={module_size:.3f}') + except Exception as e: + if 'out of memory' in str(e): + devices.torch_gc(fast=True, force=True, reason='oom') + elif 'bitsandbytes' in str(e): + pass + else: + shared.log.error(f'Offload: type=balanced op=apply module={getattr(module, "__name__", None)} cls={module.__class__ if inspect.isclass(module) else None} {e}') + if os.environ.get('SD_MOVE_DEBUG', None): + errors.display(e, f'Offload: type=balanced op=apply module={getattr(module, "__name__", None)}') + + +def apply_balanced_offload_to_module(module, checkpoint_name="", op="apply"): + module_name = getattr(module, "module_name", module.__class__.__name__) 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) @@ -333,9 +334,7 @@ def apply_balanced_offload_to_module(module, module_name=None, checkpoint_name=" module = accelerate.hooks.remove_hook_from_module(module, recurse=True) except Exception as e: shared.log.warning(f'Offload remove hook: module={module_name} {e}') - offload_hook_instance.offload_module(module, op=op) - if not hasattr(module, "offload_dir"): - module.offload_dir = os.path.join(shared.opts.accelerate_offload_path, checkpoint_name, module_name) + move_module_to_cpu(module, op=op) try: module = accelerate.hooks.add_hook_to_module(module, offload_hook_instance, append=True) except Exception as e: @@ -347,6 +346,8 @@ def apply_balanced_offload_to_module(module, module_name=None, checkpoint_name=" module.balanced_offload_device_map = device_map module.balanced_offload_max_memory = max_memory module.offload_post = shared.sd_model_type in offload_post and shared.opts.te_hijack and module_name.startswith("text_encoder") + if shared.opts.layerwise_quantization or getattr(module, 'quantization_method', None) == 'LayerWise': + model_quant.apply_layerwise(module, quiet=True) # need to reapply since hooks were removed/readded devices.torch_gc(fast=True, force=True, reason='offload') @@ -378,9 +379,9 @@ def apply_balanced_offload(sd_model=None, exclude=[]): module = getattr(pipe, module_name, None) if module is None: continue - apply_balanced_offload_to_module(module, module_name=module_name, checkpoint_name=checkpoint_name) - if shared.opts.layerwise_quantization or (hasattr(sd_model, "transformer") and getattr(sd_model.transformer, 'quantization_method', None) == 'LayerWise'): - model_quant.apply_layerwise(sd_model, quiet=True) # need to reapply since hooks were removed/readded + module.module_name = module_name + module.offload_dir = os.path.join(shared.opts.accelerate_offload_path, checkpoint_name, module_name) + apply_balanced_offload_to_module(module, checkpoint_name=checkpoint_name) set_accelerate(sd_model) t = time.time() - t0 process_timer.add('offload', t)