mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user