From dc6ffaef7bccae6de1616f2c0e6828d0349ed07e Mon Sep 17 00:00:00 2001 From: Andrew Tischenko <31860133+antis0007@users.noreply.github.com> Date: Thu, 19 Oct 2023 20:39:11 -0600 Subject: [PATCH 1/4] Adding OFT support A WIP adaptation of the OFT implementation from the Kohya repo Co-Authored-By: AI-Casanova <54461896+AI-Casanova@users.noreply.github.com> --- extensions-builtin/Lora/lora_convert.py | 10 +- extensions-builtin/Lora/network_oft.py | 196 ++++++++++++++++++++++++ extensions-builtin/Lora/networks.py | 2 + 3 files changed, 205 insertions(+), 3 deletions(-) create mode 100644 extensions-builtin/Lora/network_oft.py diff --git a/extensions-builtin/Lora/lora_convert.py b/extensions-builtin/Lora/lora_convert.py index 5843c7ad8..2afdcf258 100644 --- a/extensions-builtin/Lora/lora_convert.py +++ b/extensions-builtin/Lora/lora_convert.py @@ -114,6 +114,7 @@ class KeyConvert: self.UNET_CONVERSION_MAP = make_unet_conversion_map() if self.is_sdxl else None self.LORA_PREFIX_UNET = "lora_unet" self.LORA_PREFIX_TEXT_ENCODER = "lora_te" + self.OFT_PREFIX_UNET = "oft_unet" # SDXL: must starts with LORA_PREFIX_TEXT_ENCODER self.LORA_PREFIX_TEXT_ENCODER1 = "lora_te1" self.LORA_PREFIX_TEXT_ENCODER2 = "lora_te2" @@ -142,9 +143,12 @@ class KeyConvert: if self.is_sdxl: map_keys = list(self.UNET_CONVERSION_MAP.keys()) # prefix of U-Net modules map_keys.sort() - search_key = key.replace(self.LORA_PREFIX_UNET + "_", "").replace(self.LORA_PREFIX_TEXT_ENCODER1 + "_", - "").replace( - self.LORA_PREFIX_TEXT_ENCODER2 + "_", "") + oft_prefix = self.OFT_PREFIX_UNET + "_" + lora_prefix = self.LORA_PREFIX_UNET + "_" + te1_prefix = self.LORA_PREFIX_TEXT_ENCODER1 + "_" + te2_prefix = self.LORA_PREFIX_TEXT_ENCODER2 + "_" + search_key = key.replace(lora_prefix, "").replace(oft_prefix, "").replace(te1_prefix, "").replace(te2_prefix, "") + position = bisect.bisect_right(map_keys, search_key) map_key = map_keys[position - 1] if search_key.startswith(map_key): diff --git a/extensions-builtin/Lora/network_oft.py b/extensions-builtin/Lora/network_oft.py new file mode 100644 index 000000000..ec8e9273f --- /dev/null +++ b/extensions-builtin/Lora/network_oft.py @@ -0,0 +1,196 @@ +import torch + +import diffusers.models.lora as diffusers_lora +import lyco_helpers +import network +from modules import devices +import math +import os +from typing import Dict, List, Optional, Tuple, Type, Union +from diffusers import AutoencoderKL +from transformers import CLIPTextModel +import numpy as np +import re +#Lot of these imports are likely redundant, will refactor and remove + +#Unused regex within the original oft.py? +RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_") + +class ModuleTypeOFT(network.ModuleType): + def create_module(self, net: network.Network, weights: network.NetworkWeights): + """ + weights.w.items() + + alpha : tensor(0.0010, dtype=torch.bfloat16) + oft_blocks : tensor([[[ 0.0000e+00, 1.4400e-04, 1.7319e-03, ..., -8.8882e-04, + 5.7373e-03, -4.4250e-03], + [-1.4400e-04, 0.0000e+00, 8.6594e-04, ..., 1.5945e-03, + -8.5449e-04, 1.9684e-03], ...etc... + , dtype=torch.bfloat16)""" + + if "oft_blocks" in weights.w.keys(): + module = NetworkModuleOFT(net, weights) + return module + else: + return None + +class NetworkModuleOFT(network.NetworkModule): + def __init__(self, net: network.Network, weights: network.NetworkWeights): + super().__init__(net, weights) + """ + dim -> num blocks + alpha -> constraint + + alpha is equal to eps-deviation: eps + (only with the constrained variant COFT) + """ + self.weights = weights.w.get("oft_blocks").to(device=devices.device) + self.net = net + self.alpha = self.multiplier() + self.dim = self.weights.shape[0] #num blocks + + # old way of calculating out_features, not technically correct: + #self.out_dim = max(self.weights.shape[1],self.weights.shape[2])*self.dim + + self.is_linear = type(self.sd_module) in [torch.nn.Linear, torch.nn.modules.linear.NonDynamicallyQuantizableLinear, torch.nn.MultiheadAttention, diffusers_lora.LoRACompatibleLinear] + self.is_conv = type(self.sd_module) in [torch.nn.Conv2d, diffusers_lora.LoRACompatibleConv] + if self.is_linear == True: + self.out_dim = self.sd_module.out_features + if self.is_conv == True: + self.out_dim = self.sd_module.out_channels + #The is_conv check should be redundant? I havent seen any conv layers in my testing + + self.block_size = self.out_dim // self.dim + + #Initialize to zeros: + #self.oft_blocks = torch.nn.Parameter(torch.zeros(self.dim, self.block_size, self.block_size)).to(device=devices.device) + #Load from weights + self.oft_blocks = torch.nn.Parameter(self.weights) + #self.oft_blocks = torch.nn.Parameter(self.weights*self.alpha) #not sure if I need to apply alpha here but I just do anyway, should weaken + + #eps constraint value, calculate by using (alpha in weights) * (out_dim) + self.constraint = weights.w.get("alpha").to(device=devices.device)*self.out_dim + + def get_weight(self): + try: + self.alpha = self.multiplier() #update alpha? Not sure if necessary. + #get_weight implementation: + block_Q = self.weights - self.weights.transpose(1, 2) + norm_Q = torch.norm(block_Q.flatten()) + new_norm_Q = torch.clamp(norm_Q, max=self.constraint) + block_Q = (block_Q * ((new_norm_Q + 1e-8) / (norm_Q + 1e-8))).to(device=devices.device) + I = torch.eye(self.block_size, device=devices.device).unsqueeze(0).repeat(self.dim, 1, 1) + block_R = torch.matmul(I + block_Q, (I - block_Q).inverse()) + block_R_weighted = self.alpha*block_R + (1 - self.alpha) * I + R = torch.block_diag(*block_R_weighted) + R = R * self.alpha #Added this line, seems to make the results better, less overbaked + return R + except Exception as e: + print("ERROR:") + print(e) + + + + def calc_updown(self, orig_weight): + self.alpha = self.multiplier() #update alpha? Not sure if necessary. + output_shape = self.weights.shape + R = self.get_weight().to(device=devices.device, dtype=orig_weight.dtype) + + try: + #if orig_weight.shape[0] < orig_weight.shape[1]: + #attempt 1 + #R_expanded = torch.zeros(output_shape, device=devices.device, dtype=orig_weight.dtype) + #R_expanded[:, :R.shape[1]] = R + #R = R_expanded + #temp = orig_weight[:, :R.shape[0]] + #updown = torch.matmul(temp, R) + + #attempt 2 + #blocks = torch.split(orig_weight, split_size_or_sections=orig_weight.shape[1]//self.dim, dim=1) + #results = [torch.matmul(block,R) for block in blocks] + #updown = torch.cat(results, dim=1) + + #attempt 3 + #blocks = torch.split(orig_weight, split_size_or_sections=orig_weight.shape[1]//self.dim, dim=1) + #print("R.shape:") + #print(R.shape) + #transformed_blocks = [torch.matmul(block.transpose(1,0),R) for block in blocks] + #for i in range(0, len(transformed_blocks)): + #transformed_blocks[i] = transformed_blocks[i].transpose(1,0) + #print("END_UPDOWN") + #updown = torch.cat(transformed_blocks, dim=1) + #else: + #updown = torch.matmul(orig_weight, R) + + #Attempt 4: + if self.is_linear: + if orig_weight.shape[0] < orig_weight.shape[1]: + #check for irregular linear sizes, if dim1 is larger than dim0, that means: + #we have dim1 composed of self.dim elements (blocks) + #in order to apply batched matmul, we need to view this differently, add a dimension for our blocks + x = orig_weight.view(self.dim, orig_weight.shape[0], orig_weight.shape[1]//self.dim) + #x = orig_weight.view(self.dim, orig_weight.shape[1]//self.dim, orig_weight.shape[0]) + + #Since our size is irregular, I've made some assumptions here that may not be correct. + #I still do not fully understand what "orig_weight" represents relative to "x" in the original oft.py forward() + + #PROBLEM EXPLANATION: + #We need to do a matmul between x and R + #That means that x columns = R rows + #R will always end up a square matrix of size 640x640, or 1280x1280 (something like that) + + # However, in THESE cases, where orig_weight.shape[0] < orig_weight.shape[1]: + # x = [640,2048] or some other similar size + # We would then divide 2048 into self.dim chunks (in this case 4), and get 512 + # Thus we end up with: [4, 640, 512] where 2048 got split up into 4 channels (aka our dim) + + # Unfortunately, we cannot apply R as a matmul on this since we have unmatched dimensions + # to make this calculation possible, we need to take the transpose dim(1,2) of [4, 640, 512] to get [4, 512, 640] + + # We repeat R to fill our 4 channels, and do a batch matmul between x and R: + # [4, 512, 640](x) * [4, 640, 640](R) + + # Now after that, just torch.cat the 4 channels together back into the same shape as the beginning + + # This is just an example calculation, but one like this does happen many times + # Well, now we can kinda "calculate" something, but im honestly not sure if this is applying R properly at all. + # Here is the original forward from kohya's oft.py: + # If we could figure out a way to apply this same operation (permute/matmul for 4 dimensional input), but to our orig_weight instead of x, that would work perfect + + # Note: x.dim() == 4 is related to our (orig_weight.shape[0] < orig_weight.shape[1]) check + # If the sizes are not the same, then orig_weight.shape[1]//self.dim is the new size of our block (in that one dimension) + """ + def forward(self, x, scale=None): + x = self.org_forward(x) + if self.multiplier == 0.0: + return x + + R = self.get_weight().to(x.device, dtype=x.dtype) + if x.dim() == 4: + x = x.permute(0, 2, 3, 1) + x = torch.matmul(x, R) + x = x.permute(0, 3, 1, 2) + else: + x = torch.matmul(x, R) + return x + """ + + x = x.transpose(1,2) + #R_expanded = R.unsqueeze(0).expand(x.shape[0], -1, -1) + R_expanded = R.unsqueeze(0).repeat(x.shape[0], 1, 1) + #x = torch.bmm(x, R_expanded) + x = torch.matmul(x, R_expanded) + #x = x.transpose(1,2) + updown = torch.cat(x.unbind(0), dim=1) + #updown = x.view(orig_weight.shape[0], orig_weight.shape[1]) + else: + updown = torch.matmul(orig_weight, R) + elif self.is_conv: + updown = torch.matmul(orig_weight, R) + return(self.finalize_updown(updown, orig_weight, output_shape)) + + except Exception as e: + print("ERROR:") + print(e) + + \ No newline at end of file diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 2132c8ed5..2cb1db628 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -7,6 +7,7 @@ import network import network_lora import network_hada import network_ia3 +import network_oft import network_lokr import network_full import network_norm @@ -32,6 +33,7 @@ module_types = [ network_lora.ModuleTypeLora(), network_hada.ModuleTypeHada(), network_ia3.ModuleTypeIa3(), + network_oft.ModuleTypeOFT(), network_lokr.ModuleTypeLokr(), network_full.ModuleTypeFull(), network_norm.ModuleTypeNorm(), From 5b707e1747956fb4a41c6324af464273f72e68d1 Mon Sep 17 00:00:00 2001 From: AI-Casanova <54461896+AI-Casanova@users.noreply.github.com> Date: Thu, 19 Oct 2023 23:41:16 -0500 Subject: [PATCH 2/4] Fix and cleanup lora_convert.py --- extensions-builtin/Lora/lora_convert.py | 18 +++++++----------- 1 file changed, 7 insertions(+), 11 deletions(-) diff --git a/extensions-builtin/Lora/lora_convert.py b/extensions-builtin/Lora/lora_convert.py index 2afdcf258..fb314f258 100644 --- a/extensions-builtin/Lora/lora_convert.py +++ b/extensions-builtin/Lora/lora_convert.py @@ -112,12 +112,12 @@ class KeyConvert: self.converter = self.diffusers self.is_sdxl = True if shared.sd_model_type == "sdxl" else False self.UNET_CONVERSION_MAP = make_unet_conversion_map() if self.is_sdxl else None - self.LORA_PREFIX_UNET = "lora_unet" - self.LORA_PREFIX_TEXT_ENCODER = "lora_te" - self.OFT_PREFIX_UNET = "oft_unet" + self.LORA_PREFIX_UNET = "lora_unet_" + self.LORA_PREFIX_TEXT_ENCODER = "lora_te_" + self.OFT_PREFIX_UNET = "oft_unet_" # SDXL: must starts with LORA_PREFIX_TEXT_ENCODER - self.LORA_PREFIX_TEXT_ENCODER1 = "lora_te1" - self.LORA_PREFIX_TEXT_ENCODER2 = "lora_te2" + self.LORA_PREFIX_TEXT_ENCODER1 = "lora_te1_" + self.LORA_PREFIX_TEXT_ENCODER2 = "lora_te2_" def original(self, key): key = convert_diffusers_name_to_compvis(key, self.is_sd2) @@ -143,16 +143,12 @@ class KeyConvert: if self.is_sdxl: map_keys = list(self.UNET_CONVERSION_MAP.keys()) # prefix of U-Net modules map_keys.sort() - oft_prefix = self.OFT_PREFIX_UNET + "_" - lora_prefix = self.LORA_PREFIX_UNET + "_" - te1_prefix = self.LORA_PREFIX_TEXT_ENCODER1 + "_" - te2_prefix = self.LORA_PREFIX_TEXT_ENCODER2 + "_" - search_key = key.replace(lora_prefix, "").replace(oft_prefix, "").replace(te1_prefix, "").replace(te2_prefix, "") + search_key = key.replace(self.LORA_PREFIX_UNET, "").replace(self.OFT_PREFIX_UNET, "").replace(self.LORA_PREFIX_TEXT_ENCODER1, "").replace(self.LORA_PREFIX_TEXT_ENCODER2, "") position = bisect.bisect_right(map_keys, search_key) map_key = map_keys[position - 1] if search_key.startswith(map_key): - key = key.replace(map_key, self.UNET_CONVERSION_MAP[map_key]) # pylint: disable=unsubscriptable-object + key = key.replace(map_key, self.UNET_CONVERSION_MAP[map_key]).replace("oft","lora") # pylint: disable=unsubscriptable-object sd_module = shared.sd_model.network_layer_mapping.get(key, None) return key, sd_module From 929f73100af0b27c1fda4e867cd157cd9924441c Mon Sep 17 00:00:00 2001 From: Andrew Tischenko <31860133+antis0007@users.noreply.github.com> Date: Fri, 20 Oct 2023 09:29:36 -0600 Subject: [PATCH 3/4] Update network_oft.py Some updates to make the matmul slightly more correct? Still WIP, needs to be tested against the quality of the future einsum version. --- extensions-builtin/Lora/network_oft.py | 41 +++++++++++++++++--------- 1 file changed, 27 insertions(+), 14 deletions(-) diff --git a/extensions-builtin/Lora/network_oft.py b/extensions-builtin/Lora/network_oft.py index ec8e9273f..a2291ca4b 100644 --- a/extensions-builtin/Lora/network_oft.py +++ b/extensions-builtin/Lora/network_oft.py @@ -46,8 +46,9 @@ class NetworkModuleOFT(network.NetworkModule): """ self.weights = weights.w.get("oft_blocks").to(device=devices.device) self.net = net - self.alpha = self.multiplier() self.dim = self.weights.shape[0] #num blocks + self.num_blocks = self.dim + self.alpha = self.multiplier() # old way of calculating out_features, not technically correct: #self.out_dim = max(self.weights.shape[1],self.weights.shape[2])*self.dim @@ -73,17 +74,16 @@ class NetworkModuleOFT(network.NetworkModule): def get_weight(self): try: - self.alpha = self.multiplier() #update alpha? Not sure if necessary. #get_weight implementation: block_Q = self.weights - self.weights.transpose(1, 2) norm_Q = torch.norm(block_Q.flatten()) new_norm_Q = torch.clamp(norm_Q, max=self.constraint) - block_Q = (block_Q * ((new_norm_Q + 1e-8) / (norm_Q + 1e-8))).to(device=devices.device) + block_Q = (block_Q * ((new_norm_Q + self.constraint) / (norm_Q + self.constraint))).to(device=devices.device) I = torch.eye(self.block_size, device=devices.device).unsqueeze(0).repeat(self.dim, 1, 1) block_R = torch.matmul(I + block_Q, (I - block_Q).inverse()) block_R_weighted = self.alpha*block_R + (1 - self.alpha) * I R = torch.block_diag(*block_R_weighted) - R = R * self.alpha #Added this line, seems to make the results better, less overbaked + #R = R * self.alpha #Added this line, seems to make the results better, less overbaked return R except Exception as e: print("ERROR:") @@ -92,8 +92,6 @@ class NetworkModuleOFT(network.NetworkModule): def calc_updown(self, orig_weight): - self.alpha = self.multiplier() #update alpha? Not sure if necessary. - output_shape = self.weights.shape R = self.get_weight().to(device=devices.device, dtype=orig_weight.dtype) try: @@ -128,9 +126,23 @@ class NetworkModuleOFT(network.NetworkModule): #check for irregular linear sizes, if dim1 is larger than dim0, that means: #we have dim1 composed of self.dim elements (blocks) #in order to apply batched matmul, we need to view this differently, add a dimension for our blocks - x = orig_weight.view(self.dim, orig_weight.shape[0], orig_weight.shape[1]//self.dim) + #x = orig_weight.view(self.dim, orig_weight.shape[0], orig_weight.shape[1]//self.dim) #x = orig_weight.view(self.dim, orig_weight.shape[1]//self.dim, orig_weight.shape[0]) + channels_per_block = orig_weight.shape[1] // self.dim + + # Segment orig_weight into blocks along the second dimension (channels) + weight_segments = torch.split(orig_weight, split_size_or_sections=channels_per_block, dim=1) + # Apply the transformation to each segment + transformed_segments = [] + for segment in weight_segments: + # Reshape the segment to ensure matrix multiplication is feasible + reshaped_segment = segment.reshape(channels_per_block, -1) + transformed_segment = torch.matmul(reshaped_segment, R) + # Reshape the transformed segment back to its original shape + transformed_segment = transformed_segment.reshape(orig_weight.shape[0], channels_per_block) + transformed_segments.append(transformed_segment) + updown = torch.cat(transformed_segments, dim=1) #Since our size is irregular, I've made some assumptions here that may not be correct. #I still do not fully understand what "orig_weight" represents relative to "x" in the original oft.py forward() @@ -175,19 +187,20 @@ class NetworkModuleOFT(network.NetworkModule): return x """ - x = x.transpose(1,2) - #R_expanded = R.unsqueeze(0).expand(x.shape[0], -1, -1) - R_expanded = R.unsqueeze(0).repeat(x.shape[0], 1, 1) - #x = torch.bmm(x, R_expanded) - x = torch.matmul(x, R_expanded) #x = x.transpose(1,2) - updown = torch.cat(x.unbind(0), dim=1) + #R_expanded = R.unsqueeze(0).expand(x.shape[0], -1, -1) + #R_expanded = R.unsqueeze(0).repeat(x.shape[0], 1, 1) + #x = torch.bmm(x, R_expanded) + #x = torch.matmul(x, R_expanded) + #x = x.transpose(1,2) + #updown = torch.cat(x.unbind(0), dim=1) #updown = x.view(orig_weight.shape[0], orig_weight.shape[1]) else: updown = torch.matmul(orig_weight, R) elif self.is_conv: updown = torch.matmul(orig_weight, R) - return(self.finalize_updown(updown, orig_weight, output_shape)) + updown = updown * self.multiplier() * self.calc_scale() + return(self.finalize_updown(updown, orig_weight, orig_weight.shape)) except Exception as e: print("ERROR:") From 8f84e968d5b97cd7bcf5c5ae95712eec440af050 Mon Sep 17 00:00:00 2001 From: AI-Casanova <54461896+AI-Casanova@users.noreply.github.com> Date: Sat, 21 Oct 2023 12:11:38 -0500 Subject: [PATCH 4/4] Refactor and clean up network_oft.py Switch from matmul to einsum --- extensions-builtin/Lora/network_oft.py | 194 +++---------------------- 1 file changed, 17 insertions(+), 177 deletions(-) diff --git a/extensions-builtin/Lora/network_oft.py b/extensions-builtin/Lora/network_oft.py index a2291ca4b..6d350671a 100644 --- a/extensions-builtin/Lora/network_oft.py +++ b/extensions-builtin/Lora/network_oft.py @@ -1,20 +1,7 @@ import torch - import diffusers.models.lora as diffusers_lora -import lyco_helpers import network from modules import devices -import math -import os -from typing import Dict, List, Optional, Tuple, Type, Union -from diffusers import AutoencoderKL -from transformers import CLIPTextModel -import numpy as np -import re -#Lot of these imports are likely redundant, will refactor and remove - -#Unused regex within the original oft.py? -RE_UPDOWN = re.compile(r"(up|down)_blocks_(\d+)_(resnets|upsamplers|downsamplers|attentions)_(\d+)_") class ModuleTypeOFT(network.ModuleType): def create_module(self, net: network.Network, weights: network.NetworkWeights): @@ -27,183 +14,36 @@ class ModuleTypeOFT(network.ModuleType): [-1.4400e-04, 0.0000e+00, 8.6594e-04, ..., 1.5945e-03, -8.5449e-04, 1.9684e-03], ...etc... , dtype=torch.bfloat16)""" - + if "oft_blocks" in weights.w.keys(): module = NetworkModuleOFT(net, weights) return module else: return None + class NetworkModuleOFT(network.NetworkModule): - def __init__(self, net: network.Network, weights: network.NetworkWeights): + def __init__(self, net: network.Network, weights: network.NetworkWeights): super().__init__(net, weights) - """ - dim -> num blocks - alpha -> constraint - alpha is equal to eps-deviation: eps - (only with the constrained variant COFT) - """ self.weights = weights.w.get("oft_blocks").to(device=devices.device) - self.net = net - self.dim = self.weights.shape[0] #num blocks - self.num_blocks = self.dim + self.dim = self.weights.shape[0] # num blocks self.alpha = self.multiplier() - - # old way of calculating out_features, not technically correct: - #self.out_dim = max(self.weights.shape[1],self.weights.shape[2])*self.dim - - self.is_linear = type(self.sd_module) in [torch.nn.Linear, torch.nn.modules.linear.NonDynamicallyQuantizableLinear, torch.nn.MultiheadAttention, diffusers_lora.LoRACompatibleLinear] - self.is_conv = type(self.sd_module) in [torch.nn.Conv2d, diffusers_lora.LoRACompatibleConv] - if self.is_linear == True: - self.out_dim = self.sd_module.out_features - if self.is_conv == True: - self.out_dim = self.sd_module.out_channels - #The is_conv check should be redundant? I havent seen any conv layers in my testing - - self.block_size = self.out_dim // self.dim - - #Initialize to zeros: - #self.oft_blocks = torch.nn.Parameter(torch.zeros(self.dim, self.block_size, self.block_size)).to(device=devices.device) - #Load from weights - self.oft_blocks = torch.nn.Parameter(self.weights) - #self.oft_blocks = torch.nn.Parameter(self.weights*self.alpha) #not sure if I need to apply alpha here but I just do anyway, should weaken - - #eps constraint value, calculate by using (alpha in weights) * (out_dim) - self.constraint = weights.w.get("alpha").to(device=devices.device)*self.out_dim + self.block_size = self.weights.shape[-1] def get_weight(self): - try: - #get_weight implementation: - block_Q = self.weights - self.weights.transpose(1, 2) - norm_Q = torch.norm(block_Q.flatten()) - new_norm_Q = torch.clamp(norm_Q, max=self.constraint) - block_Q = (block_Q * ((new_norm_Q + self.constraint) / (norm_Q + self.constraint))).to(device=devices.device) - I = torch.eye(self.block_size, device=devices.device).unsqueeze(0).repeat(self.dim, 1, 1) - block_R = torch.matmul(I + block_Q, (I - block_Q).inverse()) - block_R_weighted = self.alpha*block_R + (1 - self.alpha) * I - R = torch.block_diag(*block_R_weighted) - #R = R * self.alpha #Added this line, seems to make the results better, less overbaked - return R - except Exception as e: - print("ERROR:") - print(e) - - - + block_Q = self.weights - self.weights.transpose(1, 2) + I = torch.eye(self.block_size, device=devices.device).unsqueeze(0).repeat(self.dim, 1, 1) + block_R = torch.matmul(I + block_Q, (I - block_Q).inverse()) + block_R_weighted = self.alpha * block_R + (1 - self.alpha) * I + R = torch.block_diag(*block_R_weighted) + return R + def calc_updown(self, orig_weight): R = self.get_weight().to(device=devices.device, dtype=orig_weight.dtype) - - try: - #if orig_weight.shape[0] < orig_weight.shape[1]: - #attempt 1 - #R_expanded = torch.zeros(output_shape, device=devices.device, dtype=orig_weight.dtype) - #R_expanded[:, :R.shape[1]] = R - #R = R_expanded - #temp = orig_weight[:, :R.shape[0]] - #updown = torch.matmul(temp, R) + if orig_weight.dim() == 4: + updown = torch.einsum("oihw, op -> pihw", orig_weight, R) * self.calc_scale() + else: + updown = torch.einsum("oi, op -> pi", orig_weight, R) * self.calc_scale() - #attempt 2 - #blocks = torch.split(orig_weight, split_size_or_sections=orig_weight.shape[1]//self.dim, dim=1) - #results = [torch.matmul(block,R) for block in blocks] - #updown = torch.cat(results, dim=1) - - #attempt 3 - #blocks = torch.split(orig_weight, split_size_or_sections=orig_weight.shape[1]//self.dim, dim=1) - #print("R.shape:") - #print(R.shape) - #transformed_blocks = [torch.matmul(block.transpose(1,0),R) for block in blocks] - #for i in range(0, len(transformed_blocks)): - #transformed_blocks[i] = transformed_blocks[i].transpose(1,0) - #print("END_UPDOWN") - #updown = torch.cat(transformed_blocks, dim=1) - #else: - #updown = torch.matmul(orig_weight, R) - - #Attempt 4: - if self.is_linear: - if orig_weight.shape[0] < orig_weight.shape[1]: - #check for irregular linear sizes, if dim1 is larger than dim0, that means: - #we have dim1 composed of self.dim elements (blocks) - #in order to apply batched matmul, we need to view this differently, add a dimension for our blocks - #x = orig_weight.view(self.dim, orig_weight.shape[0], orig_weight.shape[1]//self.dim) - #x = orig_weight.view(self.dim, orig_weight.shape[1]//self.dim, orig_weight.shape[0]) - channels_per_block = orig_weight.shape[1] // self.dim - - # Segment orig_weight into blocks along the second dimension (channels) - weight_segments = torch.split(orig_weight, split_size_or_sections=channels_per_block, dim=1) - - # Apply the transformation to each segment - transformed_segments = [] - for segment in weight_segments: - # Reshape the segment to ensure matrix multiplication is feasible - reshaped_segment = segment.reshape(channels_per_block, -1) - transformed_segment = torch.matmul(reshaped_segment, R) - # Reshape the transformed segment back to its original shape - transformed_segment = transformed_segment.reshape(orig_weight.shape[0], channels_per_block) - transformed_segments.append(transformed_segment) - updown = torch.cat(transformed_segments, dim=1) - #Since our size is irregular, I've made some assumptions here that may not be correct. - #I still do not fully understand what "orig_weight" represents relative to "x" in the original oft.py forward() - - #PROBLEM EXPLANATION: - #We need to do a matmul between x and R - #That means that x columns = R rows - #R will always end up a square matrix of size 640x640, or 1280x1280 (something like that) - - # However, in THESE cases, where orig_weight.shape[0] < orig_weight.shape[1]: - # x = [640,2048] or some other similar size - # We would then divide 2048 into self.dim chunks (in this case 4), and get 512 - # Thus we end up with: [4, 640, 512] where 2048 got split up into 4 channels (aka our dim) - - # Unfortunately, we cannot apply R as a matmul on this since we have unmatched dimensions - # to make this calculation possible, we need to take the transpose dim(1,2) of [4, 640, 512] to get [4, 512, 640] - - # We repeat R to fill our 4 channels, and do a batch matmul between x and R: - # [4, 512, 640](x) * [4, 640, 640](R) - - # Now after that, just torch.cat the 4 channels together back into the same shape as the beginning - - # This is just an example calculation, but one like this does happen many times - # Well, now we can kinda "calculate" something, but im honestly not sure if this is applying R properly at all. - # Here is the original forward from kohya's oft.py: - # If we could figure out a way to apply this same operation (permute/matmul for 4 dimensional input), but to our orig_weight instead of x, that would work perfect - - # Note: x.dim() == 4 is related to our (orig_weight.shape[0] < orig_weight.shape[1]) check - # If the sizes are not the same, then orig_weight.shape[1]//self.dim is the new size of our block (in that one dimension) - """ - def forward(self, x, scale=None): - x = self.org_forward(x) - if self.multiplier == 0.0: - return x - - R = self.get_weight().to(x.device, dtype=x.dtype) - if x.dim() == 4: - x = x.permute(0, 2, 3, 1) - x = torch.matmul(x, R) - x = x.permute(0, 3, 1, 2) - else: - x = torch.matmul(x, R) - return x - """ - - #x = x.transpose(1,2) - #R_expanded = R.unsqueeze(0).expand(x.shape[0], -1, -1) - #R_expanded = R.unsqueeze(0).repeat(x.shape[0], 1, 1) - #x = torch.bmm(x, R_expanded) - #x = torch.matmul(x, R_expanded) - #x = x.transpose(1,2) - #updown = torch.cat(x.unbind(0), dim=1) - #updown = x.view(orig_weight.shape[0], orig_weight.shape[1]) - else: - updown = torch.matmul(orig_weight, R) - elif self.is_conv: - updown = torch.matmul(orig_weight, R) - updown = updown * self.multiplier() * self.calc_scale() - return(self.finalize_updown(updown, orig_weight, orig_weight.shape)) - - except Exception as e: - print("ERROR:") - print(e) - - \ No newline at end of file + return self.finalize_updown(updown, orig_weight, orig_weight.shape)