From 8725cfc4887b30730c7a1759dc4b9c543388a296 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 3 Apr 2025 14:31:37 -0400 Subject: [PATCH] lora obey device Signed-off-by: Vladimir Mandic --- extensions-builtin/sdnext-modernui | 2 +- modules/lora/lora_apply.py | 34 ++++++++++++++---------------- modules/lora/lora_load.py | 10 +++++---- modules/lora/network_lora.py | 5 +++-- 4 files changed, 26 insertions(+), 25 deletions(-) diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 770db0076..9bed415dc 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 770db007688d5be9df0def02af64a1fe6449c04e +Subproject commit 9bed415dccad1041e0573134d2f214e40bd310f1 diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index 8bf327151..7164543d1 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -74,12 +74,11 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. continue try: t0 = time.time() - try: - weight = self.weight.to(devices.device) - except Exception: - weight = self.weight + weight = self.weight.to(devices.device) # must perform calc on gpu due to performance updown, ex_bias = module.calc_updown(weight) - del module + weight = None + del weight + if updown is not None: if batch_updown is not None: batch_updown += updown.to(batch_updown.device) @@ -91,6 +90,7 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. else: batch_ex_bias = ex_bias.to(devices.device) l.timer.calc += time.time() - t0 + if shared.opts.diffusers_offload_mode == "sequential": t0 = time.time() if batch_updown is not None: @@ -110,20 +110,21 @@ def network_calc_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. return batch_updown, batch_ex_bias -def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False, bias: bool = False): +def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], model_weights: Union[None, torch.Tensor] = None, lora_weights: torch.Tensor = None, deactivate: bool = False, device: torch.device = None, bias: bool = False): if lora_weights is None: - return None + return if deactivate: lora_weights *= -1 if model_weights is None: # weights are used if provided-from-backup else use self.weight model_weights = self.weight + weight, new_weight = None, None + device = device or devices.device # TODO lora: add other quantization types - weight = None if self.__class__.__name__ == 'Linear4bit' and bnb is not None: try: dequant_weight = bnb.functional.dequantize_4bit(model_weights.to(devices.device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize) new_weight = dequant_weight.to(devices.device) + lora_weights.to(devices.device) - weight = bnb.nn.Params4bit(new_weight, quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize, requires_grad=False) + weight = bnb.nn.Params4bit(new_weight.to(device), quant_state=self.quant_state, quant_type=self.quant_type, blocksize=self.blocksize, requires_grad=False) # TODO lora: maybe force imediate quantization # weight._quantize(devices.device) / weight.to(device=device) except Exception as e: @@ -134,16 +135,13 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G except Exception as e: shared.log.warning(f'Network load: {e}') new_weight = model_weights + lora_weights # try without device cast - del model_weights - del lora_weights - weight = torch.nn.Parameter(new_weight, requires_grad=False) - del new_weight # without this its a massive memory leak + weight = torch.nn.Parameter(new_weight.to(device), requires_grad=False) if weight is not None: if not bias: self.weight = weight else: self.bias = weight - return weight + del model_weights, lora_weights, new_weight, weight # required to avoid memory leak def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.GroupNorm, torch.nn.LayerNorm, diffusers.models.lora.LoRACompatibleLinear, diffusers.models.lora.LoRACompatibleConv], updown: torch.Tensor, ex_bias: torch.Tensor, deactivate: bool = False, device: torch.device = devices.device): @@ -162,11 +160,11 @@ def network_apply_direct(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn. if updown is not None and len(self.weight.shape) == 4 and self.weight.shape[1] == 9: # inpainting model so zero pad updown to make channel 4 to 9 updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if updown is not None: - network_add_weights(self, lora_weights=updown, deactivate=deactivate, bias=False) + network_add_weights(self, lora_weights=updown, deactivate=deactivate, device=device, bias=False) if bias_backup: if ex_bias is not None: - network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate, bias=True) + network_add_weights(self, lora_weights=ex_bias, deactivate=deactivate, device=device, bias=True) if hasattr(self, "qweight") and hasattr(self, "freeze"): self.freeze() @@ -186,14 +184,14 @@ def network_apply_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn if updown is not None and len(weights_backup.shape) == 4 and weights_backup.shape[1] == 9: # inpainting model. zero pad updown to make channel[1] 4 to 9 updown = torch.nn.functional.pad(updown, (0, 0, 0, 0, 0, 5)) # pylint: disable=not-callable if updown is not None: - network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate, bias=False) + network_add_weights(self, model_weights=weights_backup, lora_weights=updown, deactivate=deactivate, device=device, bias=False) else: self.weight = torch.nn.Parameter(weights_backup.to(device), requires_grad=False) if bias_backup is not None: self.bias = None if ex_bias is not None: - network_add_weights(self, model_weights=bias_backup, lora_weights=ex_bias, deactivate=deactivate, bias=True) + network_add_weights(self, model_weights=bias_backup, lora_weights=ex_bias, deactivate=deactivate, device=device, bias=True) else: self.bias = torch.nn.Parameter(bias_backup.to(device), requires_grad=False) diff --git a/modules/lora/lora_load.py b/modules/lora/lora_load.py index 110cc0e46..be7b127fc 100644 --- a/modules/lora/lora_load.py +++ b/modules/lora/lora_load.py @@ -58,12 +58,12 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: return cached net = network.Network(name, network_on_disk) net.mtime = os.path.getmtime(network_on_disk.filename) - sd = sd_models.read_state_dict(network_on_disk.filename, what='network') + state_dict = sd_models.read_state_dict(network_on_disk.filename, what='network') if shared.sd_model_type == 'f1': # if kohya flux lora, convert state_dict - sd = lora_convert._convert_kohya_flux_lora_to_diffusers(sd) or sd # pylint: disable=protected-access + state_dict = lora_convert._convert_kohya_flux_lora_to_diffusers(state_dict) or state_dict # pylint: disable=protected-access if shared.sd_model_type == 'sd3': # if kohya flux lora, convert state_dict try: - sd = lora_convert._convert_kohya_sd3_lora_to_diffusers(sd) or sd # pylint: disable=protected-access + state_dict = lora_convert._convert_kohya_sd3_lora_to_diffusers(state_dict) or state_dict # pylint: disable=protected-access except ValueError: # EAFP for diffusers PEFT keys pass lora_convert.assign_network_names_to_compvis_modules(shared.sd_model) @@ -73,7 +73,7 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: dtypes = [] convert = lora_convert.KeyConvert() device = devices.device if shared.opts.lora_apply_gpu else devices.cpu - for key_network, weight in sd.items(): + for key_network, weight in state_dict.items(): parts = key_network.split('.') if parts[0] == "bundle_emb": emb_name, vec_name = parts[1], key_network.split(".", 2)[-1] @@ -98,6 +98,8 @@ def load_safetensors(name, network_on_disk) -> Union[network.Network, None]: if weight.dtype not in dtypes: dtypes.append(weight.dtype) network_types = [] + state_dict = None + del state_dict for key, weights in matched_networks.items(): net_module = None for nettype in l.module_types: diff --git a/modules/lora/network_lora.py b/modules/lora/network_lora.py index 3604e059d..1043bc8c2 100644 --- a/modules/lora/network_lora.py +++ b/modules/lora/network_lora.py @@ -57,14 +57,15 @@ class NetworkModuleLora(network.NetworkModule): down = self.down_model.weight.to(target.device, dtype=target_dtype) output_shape = [up.size(0), down.size(1)] if self.mid_model is not None: - # cp-decomposition mid = self.mid_model.weight.to(target.device, dtype=target_dtype) - updown = lyco_helpers.rebuild_cp_decomposition(up, down, mid) + updown = lyco_helpers.rebuild_cp_decomposition(up, down, mid) # cp-decomposition output_shape += mid.shape[2:] else: + mid = None if len(down.shape) == 4: output_shape += down.shape[2:] updown = lyco_helpers.rebuild_conventional(up, down, output_shape, self.network.dyn_dim) + del up, down, mid return self.finalize_updown(updown, target, output_shape) def forward(self, x, y):