diff --git a/extensions-builtin/Lora/lora_convert.py b/extensions-builtin/Lora/lora_convert.py index 5843c7ad8..fb314f258 100644 --- a/extensions-builtin/Lora/lora_convert.py +++ b/extensions-builtin/Lora/lora_convert.py @@ -112,11 +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.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) @@ -142,13 +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 + "_", "") + 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 diff --git a/extensions-builtin/Lora/network_oft.py b/extensions-builtin/Lora/network_oft.py new file mode 100644 index 000000000..6d350671a --- /dev/null +++ b/extensions-builtin/Lora/network_oft.py @@ -0,0 +1,49 @@ +import torch +import diffusers.models.lora as diffusers_lora +import network +from modules import devices + +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) + + self.weights = weights.w.get("oft_blocks").to(device=devices.device) + self.dim = self.weights.shape[0] # num blocks + self.alpha = self.multiplier() + self.block_size = self.weights.shape[-1] + + def get_weight(self): + 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) + 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() + + return self.finalize_updown(updown, orig_weight, orig_weight.shape) 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(),