Merge pull request #3853 from vladmandic/dev

refresh master
This commit is contained in:
Vladimir Mandic
2025-04-03 14:48:28 -04:00
committed by GitHub
4 changed files with 26 additions and 25 deletions
+16 -18
View File
@@ -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)
+6 -4
View File
@@ -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:
+3 -2
View File
@@ -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):