From 998968f5ff6bee792ffd64c6d47111af5187fc00 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 16 Jan 2025 19:36:18 +0200 Subject: [PATCH] Support LyCORIS https://github.com/KohakuBlueleaf/Lycoris --- lycoris/__init__.py | 28 ++ lycoris/config.py | 151 +++++++ lycoris/functional/__init__.py | 9 + lycoris/functional/boft.py | 122 ++++++ lycoris/functional/diag_oft.py | 112 ++++++ lycoris/functional/general.py | 108 +++++ lycoris/functional/locon.py | 85 ++++ lycoris/functional/loha.py | 165 ++++++++ lycoris/functional/lokr.py | 247 ++++++++++++ lycoris/kohya.py | 676 ++++++++++++++++++++++++++++++++ lycoris/logging.py | 52 +++ lycoris/modules/__init__.py | 46 +++ lycoris/modules/base.py | 315 +++++++++++++++ lycoris/modules/boft.py | 255 ++++++++++++ lycoris/modules/diag_oft.py | 217 ++++++++++ lycoris/modules/dylora.py | 156 ++++++++ lycoris/modules/full.py | 214 ++++++++++ lycoris/modules/glora.py | 262 +++++++++++++ lycoris/modules/ia3.py | 142 +++++++ lycoris/modules/locon.py | 332 ++++++++++++++++ lycoris/modules/loha.py | 329 ++++++++++++++++ lycoris/modules/lokr.py | 609 ++++++++++++++++++++++++++++ lycoris/modules/norms.py | 161 ++++++++ lycoris/utils/__init__.py | 483 +++++++++++++++++++++++ lycoris/utils/general.py | 5 + lycoris/utils/logger.py | 35 ++ lycoris/utils/preset.py | 9 + lycoris/utils/quant.py | 88 +++++ lycoris/utils/xformers_utils.py | 13 + lycoris/wrapper.py | 640 ++++++++++++++++++++++++++++++ nodes.py | 54 ++- nodes_sdxl.py | 12 +- 32 files changed, 6124 insertions(+), 8 deletions(-) create mode 100644 lycoris/__init__.py create mode 100644 lycoris/config.py create mode 100644 lycoris/functional/__init__.py create mode 100644 lycoris/functional/boft.py create mode 100644 lycoris/functional/diag_oft.py create mode 100644 lycoris/functional/general.py create mode 100644 lycoris/functional/locon.py create mode 100644 lycoris/functional/loha.py create mode 100644 lycoris/functional/lokr.py create mode 100644 lycoris/kohya.py create mode 100644 lycoris/logging.py create mode 100644 lycoris/modules/__init__.py create mode 100644 lycoris/modules/base.py create mode 100644 lycoris/modules/boft.py create mode 100644 lycoris/modules/diag_oft.py create mode 100644 lycoris/modules/dylora.py create mode 100644 lycoris/modules/full.py create mode 100644 lycoris/modules/glora.py create mode 100644 lycoris/modules/ia3.py create mode 100644 lycoris/modules/locon.py create mode 100644 lycoris/modules/loha.py create mode 100644 lycoris/modules/lokr.py create mode 100644 lycoris/modules/norms.py create mode 100644 lycoris/utils/__init__.py create mode 100644 lycoris/utils/general.py create mode 100644 lycoris/utils/logger.py create mode 100644 lycoris/utils/preset.py create mode 100644 lycoris/utils/quant.py create mode 100644 lycoris/utils/xformers_utils.py create mode 100644 lycoris/wrapper.py diff --git a/lycoris/__init__.py b/lycoris/__init__.py new file mode 100644 index 0000000..7a1846f --- /dev/null +++ b/lycoris/__init__.py @@ -0,0 +1,28 @@ +#source https://github.com/KohakuBlueleaf/Lycoris + +# try: +# from . import kohya +# except Exception: +# pass +# from . import ( +# modules, +# utils, +# ) + +# from .modules.locon import LoConModule +# from .modules.loha import LohaModule +# from .modules.lokr import LokrModule +# from .modules.dylora import DyLoraModule +# from .modules.glora import GLoRAModule +# from .modules.norms import NormModule +# from .modules.full import FullModule +# from .modules.diag_oft import DiagOFTModule +# from .modules import make_module + +# from .wrapper import ( +# LycorisNetwork, +# create_lycoris, +# create_lycoris_from_weights, +# ) + +# from .logging import logger diff --git a/lycoris/config.py b/lycoris/config.py new file mode 100644 index 0000000..1f78296 --- /dev/null +++ b/lycoris/config.py @@ -0,0 +1,151 @@ +PRESET = { + "full": { + "enable_conv": True, + "unet_target_module": [ + "Transformer2DModel", + "ResnetBlock2D", + "Downsample2D", + "Upsample2D", + "HunYuanDiTBlock", #HunYuanDiT + "DoubleStreamBlock", #Flux + "SingleStreamBlock", #Flux + "SingleDiTBlock", #SD3.5 + "MMDoubleStreamBlock", #HunYuanVideo + "MMSingleStreamBlock", #HunYuanVideo + ], + "unet_target_name": [ + "conv_in", + "conv_out", + "time_embedding.linear_1", + "time_embedding.linear_2", + ], + "text_encoder_target_module": [ + "CLIPAttention", + "CLIPSdpaAttention", + "CLIPMLP", + "MT5Block", + "BertLayer", + ], + "text_encoder_target_name": [], + }, + "full-lin": { + "enable_conv": False, + "unet_target_module": [ + "Transformer2DModel", + "ResnetBlock2D", + "HunYuanDiTBlock", + "DoubleStreamBlock", + "SingleStreamBlock", + "SingleDiTBlock", + "MMDoubleStreamBlock", #HunYuanVideo + "MMSingleStreamBlock", #HunYuanVideo + ], + "unet_target_name": [ + "time_embedding.linear_1", + "time_embedding.linear_2", + ], + "text_encoder_target_module": [ + "CLIPAttention", + "CLIPSdpaAttention", + "CLIPMLP", + "MT5Block", + "BertLayer", + ], + "text_encoder_target_name": [], + }, + "attn-mlp": { + "enable_conv": False, + "unet_target_module": [ + "Transformer2DModel", + "HunYuanDiTBlock", + "DoubleStreamBlock", + "SingleStreamBlock", + "SingleDiTBlock", + "MMDoubleStreamBlock", #HunYuanVideo + "MMSingleStreamBlock", #HunYuanVideo + ], + "unet_target_name": [], + "text_encoder_target_module": [ + "CLIPAttention", + "CLIPSdpaAttention", + "CLIPMLP", + "MT5Block", + "BertLayer", + ], + "text_encoder_target_name": [], + }, + "attn-only": { + "enable_conv": False, + "unet_target_module": [ + "CrossAttention", + "SelfAttention", + ], + "unet_target_name": [], + "text_encoder_target_module": [ + "CLIPAttention", + "CLIPSdpaAttention", + "BertAttention", + "MT5LayerSelfAttention", + ], + "text_encoder_target_name": [], + }, + "unet-only": { + "enable_conv": True, + "unet_target_module": [ + "Transformer2DModel", + "ResnetBlock2D", + "Downsample2D", + "Upsample2D", + "HunYuanDiTBlock", + "DoubleStreamBlock", + "SingleStreamBlock", + "SingleDiTBlock", + "MMDoubleStreamBlock", #HunYuanVideo + "MMSingleStreamBlock", #HunYuanVideo + ], + "unet_target_name": [ + "conv_in", + "conv_out", + "time_embedding.linear_1", + "time_embedding.linear_2", + ], + "text_encoder_target_module": [], + "text_encoder_target_name": [], + }, + "unet-transformer-only": { + "enable_conv": False, + "unet_target_module": [ + "Transformer2DModel", + "HunYuanDiTBlock", + "DoubleStreamBlock", + "SingleStreamBlock", + "SingleDiTBlock", + "MMDoubleStreamBlock", #HunYuanVideo + "MMSingleStreamBlock", #HunYuanVideo + ], + "unet_target_name": [], + "text_encoder_target_module": [], + "text_encoder_target_name": [], + }, + "unet-convblock-only": { + "enable_conv": True, + "unet_target_module": ["ResnetBlock2D", "Downsample2D", "Upsample2D"], + "unet_target_name": [ + "conv_in", + "conv_out", + ], + "text_encoder_target_module": [], + "text_encoder_target_name": [], + }, + "ia3": { + "enable_conv": False, + "unet_target_module": [], + "unet_target_name": ["to_k", "to_v", "ff.net.2"], + "text_encoder_target_module": [], + "text_encoder_target_name": ["k_proj", "v_proj", "mlp.fc2"], + "name_algo_map": { + "mlp.fc2": {"train_on_input": True}, + "ff.net.2": {"train_on_input": True}, + }, + }, +} diff --git a/lycoris/functional/__init__.py b/lycoris/functional/__init__.py new file mode 100644 index 0000000..a489a16 --- /dev/null +++ b/lycoris/functional/__init__.py @@ -0,0 +1,9 @@ +from .general import ( + rebuild_tucker, + factorization, + power2factorization, + FUNC_LIST, + tucker_weight, + tucker_weight_from_conv, + apply_dora_scale, +) diff --git a/lycoris/functional/boft.py b/lycoris/functional/boft.py new file mode 100644 index 0000000..4eeac68 --- /dev/null +++ b/lycoris/functional/boft.py @@ -0,0 +1,122 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from einops import rearrange + +from .general import power2factorization, FUNC_LIST +from .diag_oft import get_r + + +def weight_gen(org_weight, max_block_size, boft_m=-1, rescale=False): + """### boft_weight_gen + + Args: + org_weight (torch.Tensor): the weight tensor + max_block_size (int): max block size + rescale (bool, optional): whether to rescale the weight. Defaults to False. + + Returns: + torch.Tensor: oft_blocks[, rescale_weight] + """ + out_dim, *rest = org_weight.shape + block_size, block_num = power2factorization(out_dim, max_block_size) + max_boft_m = sum(int(i) for i in f"{block_num-1:b}") + 1 + if boft_m == -1: + boft_m = max_boft_m + boft_m = min(boft_m, max_boft_m) + oft_blocks = torch.zeros(boft_m, block_num, block_size, block_size) + if rescale is not None: + return oft_blocks, torch.ones(out_dim, *[1] * len(rest)) + else: + return oft_blocks, None + + +def diff_weight(org_weight, *weights, constraint=None): + """### boft_diff_weight + + Args: + org_weight (torch.Tensor): the weight tensor of original model + weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight]) + constraint (float, optional): constraint for oft + + Returns: + torch.Tensor: ΔW + """ + oft_blocks, rescale = weights + m, num, b, _ = oft_blocks.shape + r_b = b // 2 + I = torch.eye(b, device=oft_blocks.device) + r = get_r(oft_blocks, I, constraint) + inp = org = org_weight.to(dtype=r.dtype) + + for i in range(m): + bi = r[i] # b_num, b_size, b_size + g = 2 + k = 2**i * r_b + inp = ( + inp.unflatten(-1, (-1, g, k)) + .transpose(-2, -1) + .flatten(-3) + .unflatten(-1, (-1, b)) + ) + inp = torch.einsum("b i j, b j ... -> b i ...", bi, inp) + inp = inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3) + + if rescale is not None: + inp = inp * rescale + + return inp - org + + +def bypass_forward_diff(org_out, *weights, constraint=None, need_transpose=False): + """### boft_bypass_forward_diff + + Args: + x (torch.Tensor): the input tensor for original model + org_out (torch.Tensor): the output tensor from original model + weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight]) + constraint (float, optional): constraint for oft + need_transpose (bool, optional): + whether to transpose the input and output, + set to `True` if the original model have "dim" not in the last axis. + For example: Convolution layers + + Returns: + torch.Tensor: output tensor + """ + oft_blocks, rescale = weights + m, num, b, _ = oft_blocks.shape + r_b = b // 2 + I = torch.eye(b, device=oft_blocks.device) + r = get_r(oft_blocks, I, constraint) + inp = org = org_out.to(dtype=r.dtype) + if need_transpose: + inp = org = inp.transpose(1, -1) + + for i in range(m): + bi = r[i] # b_num, b_size, b_size + g = 2 + k = 2**i * r_b + # ... (c g k) ->... (c k g) + # ... (d b) -> ... d b + inp = ( + inp.unflatten(-1, (-1, g, k)) + .transpose(-2, -1) + .flatten(-3) + .unflatten(-1, (-1, b)) + ) + inp = torch.einsum("b i j, ... b j -> ... b i", bi, inp) + # ... d b -> ... (d b) + # ... (c k g) -> ... (c g k) + inp = inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3) + + if rescale is not None: + inp = inp * rescale.transpose(0, -1) + + inp = inp - org + if need_transpose: + inp = inp.transpose(1, -1) + return inp diff --git a/lycoris/functional/diag_oft.py b/lycoris/functional/diag_oft.py new file mode 100644 index 0000000..233a609 --- /dev/null +++ b/lycoris/functional/diag_oft.py @@ -0,0 +1,112 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .general import factorization, FUNC_LIST + + +def get_r(oft_blocks, I=None, constraint=0): + if I is None: + I = torch.eye(oft_blocks.shape[-1], device=oft_blocks.device) + if I.ndim < oft_blocks.ndim: + for _ in range(oft_blocks.ndim - I.ndim): + I = I.unsqueeze(0) + # for Q = -Q^T + q = oft_blocks - oft_blocks.transpose(-1, -2) + normed_q = q + if constraint is not None and constraint > 0: + q_norm = torch.norm(q) + 1e-8 + if q_norm > constraint: + normed_q = q * constraint / q_norm + # use float() to prevent unsupported type + r = (I + normed_q) @ (I - normed_q).float().inverse() + return r + + +def weight_gen(org_weight, max_block_size=-1, rescale=False): + """### weight_gen + + Args: + org_weight (torch.Tensor): the weight tensor + max_block_size (int): max block size + rescale (bool, optional): whether to rescale the weight. Defaults to False. + + Returns: + torch.Tensor: oft_blocks[, rescale_weight] + """ + out_dim, *rest = org_weight.shape + block_size, block_num = factorization(out_dim, max_block_size) + oft_blocks = torch.zeros(block_num, block_size, block_size) + if rescale: + return oft_blocks, torch.ones(out_dim, *[1] * len(rest)) + else: + return oft_blocks, None + + +def diff_weight(org_weight, *weights, constraint=None): + """### diff_weight + + Args: + org_weight (torch.Tensor): the weight tensor of original model + weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight]) + constraint (float, optional): constraint for oft + + Returns: + torch.Tensor: ΔW + """ + oft_blocks, rescale = weights + I = torch.eye(oft_blocks.shape[1], device=oft_blocks.device) + r = get_r(oft_blocks, I, constraint) + + block_num, block_size, _ = oft_blocks.shape + _, *shape = org_weight.shape + org_weight = org_weight.to(dtype=r.dtype) + org_weight = org_weight.view(block_num, block_size, *shape) + # Init R=0, so add I on it to ensure the output of step0 is original model output + weight = torch.einsum( + "k n m, k n ... -> k m ...", + r - I, + org_weight, + ).view(-1, *shape) + if rescale is not None: + weight = rescale * weight + weight = weight + (rescale - 1) * org_weight + return weight + + +def bypass_forward_diff(x, org_out, *weights, constraint=None, need_transpose=False): + """### bypass_forward_diff + + Args: + x (torch.Tensor): the input tensor for original model + org_out (torch.Tensor): the output tensor from original model + weights (tuple[torch.Tensor]): (oft_blocks[, rescale_weight]) + constraint (float, optional): constraint for oft + need_transpose (bool, optional): + whether to transpose the input and output, + set to `True` if the original model have "dim" not in the last axis. + For example: Convolution layers + + Returns: + torch.Tensor: output tensor + """ + oft_blocks, rescale = weights + block_num, block_size, _ = oft_blocks.shape + I = torch.eye(block_size, device=oft_blocks.device) + r = get_r(oft_blocks, I, constraint) + if need_transpose: + org_out = org_out.transpose(1, -1) + org_out = org_out.to(dtype=r.dtype) + *shape, _ = org_out.shape + oft_out = torch.einsum( + "k n m, ... k n -> ... k m", r - I, org_out.view(*shape, block_num, block_size) + ) + out = oft_out.view(*shape, -1) + if rescale is not None: + out = rescale.transpose(-1, 0) * out + out = out + (rescale - 1).transpose(-1, 0) * org_out + if need_transpose: + out = out.transpose(1, -1) + return out diff --git a/lycoris/functional/general.py b/lycoris/functional/general.py new file mode 100644 index 0000000..dc0f8a0 --- /dev/null +++ b/lycoris/functional/general.py @@ -0,0 +1,108 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + + +FUNC_LIST = [None, None, F.linear, F.conv1d, F.conv2d, F.conv3d] + + +def rebuild_tucker(t, wa, wb): + rebuild2 = torch.einsum("i j ..., i p, j r -> p r ...", t, wa, wb) + return rebuild2 + + +def factorization(dimension: int, factor: int = -1) -> tuple[int, int]: + """ + return a tuple of two value of input dimension decomposed by the number closest to factor + second value is higher or equal than first value. + + In LoRA with Kroneckor Product, first value is a value for weight scale. + second value is a value for weight. + + Because of non-commutative property, A⊗B ≠ B⊗A. Meaning of two matrices is slightly different. + + examples) + factor + -1 2 4 8 16 ... + 127 -> 1, 127 127 -> 1, 127 127 -> 1, 127 127 -> 1, 127 127 -> 1, 127 + 128 -> 8, 16 128 -> 2, 64 128 -> 4, 32 128 -> 8, 16 128 -> 8, 16 + 250 -> 10, 25 250 -> 2, 125 250 -> 2, 125 250 -> 5, 50 250 -> 10, 25 + 360 -> 8, 45 360 -> 2, 180 360 -> 4, 90 360 -> 8, 45 360 -> 12, 30 + 512 -> 16, 32 512 -> 2, 256 512 -> 4, 128 512 -> 8, 64 512 -> 16, 32 + 1024 -> 32, 32 1024 -> 2, 512 1024 -> 4, 256 1024 -> 8, 128 1024 -> 16, 64 + """ + + if factor > 0 and (dimension % factor) == 0: + m = factor + n = dimension // factor + if m > n: + n, m = m, n + return m, n + if factor < 0: + factor = dimension + m, n = 1, dimension + length = m + n + while m < n: + new_m = m + 1 + while dimension % new_m != 0: + new_m += 1 + new_n = dimension // new_m + if new_m + new_n > length or new_m > factor: + break + else: + m, n = new_m, new_n + if m > n: + n, m = m, n + return m, n + + +def power2factorization(dimension: int, factor: int = -1) -> tuple[int, int]: + """ + m = 2k + n = 2**p + m*n = dim + """ + if factor == -1: + factor = dimension + + # Find the first solution and check if it is even doable + m = n = 0 + while m <= factor: + m += 2 + while dimension % m != 0 and m < dimension: + m += 2 + if m > factor: + break + if sum(int(i) for i in f"{dimension//m:b}") == 1: + n = dimension // m + + if n == 0: + return None, n + return dimension // n, n + + +def tucker_weight_from_conv(up, down, mid): + up = up.reshape(up.size(0), up.size(1)) + down = down.reshape(down.size(0), down.size(1)) + return torch.einsum("m n ..., i m, n j -> i j ...", mid, up, down) + + +def tucker_weight(wa, wb, t): + temp = torch.einsum("i j ..., j r -> i r ...", t, wb) + return torch.einsum("i j ..., i r -> r j ...", temp, wa) + + +def apply_dora_scale(org_weight, rebuild, dora_scale, scale): + dora_norm_dims = org_weight.dim() - 1 + weight = org_weight + rebuild + weight = weight.to(dora_scale.dtype) + weight_norm = ( + weight.transpose(0, 1) + .reshape(weight.shape[1], -1) + .norm(dim=1, keepdim=True) + .reshape(weight.shape[1], *[1] * dora_norm_dims) + .transpose(0, 1) + ) + merged_scale1 = weight / weight_norm * dora_scale + diff_weight = merged_scale1 - org_weight + return org_weight + diff_weight * scale diff --git a/lycoris/functional/locon.py b/lycoris/functional/locon.py new file mode 100644 index 0000000..756bbd8 --- /dev/null +++ b/lycoris/functional/locon.py @@ -0,0 +1,85 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .general import rebuild_tucker, FUNC_LIST + + +def weight_gen(org_weight, rank, tucker=True): + """### weight_gen + + Args: + org_weight (torch.Tensor): the weight tensor + rank (int): low rank + + Returns: + torch.Tensor: down, up[, mid] + """ + out_dim, in_dim, *k = org_weight.shape + if k and tucker: + down = torch.empty(rank, in_dim, *(1 for _ in k)) + up = torch.empty(out_dim, rank, *(1 for _ in k)) + mid = torch.empty(rank, rank, *k) + nn.init.kaiming_uniform_(down, a=math.sqrt(5)) + nn.init.constant_(up, 0) + nn.init.kaiming_uniform_(mid, a=math.sqrt(5)) + return down, up, mid + else: + down = torch.empty(rank, in_dim) + up = torch.empty(out_dim, rank) + nn.init.kaiming_uniform_(down, a=math.sqrt(5)) + nn.init.constant_(up, 0) + return down, up, None + + +def diff_weight(*weights: tuple[torch.Tensor], gamma=1.0): + """### diff_weight + + Get ΔW = BA, where BA is low rank decomposition + + Args: + weights (tuple[torch.Tensor]): (down, up[, mid]) + gamma (float, optional): scale factor, normally alpha/rank here + + Returns: + torch.Tensor: ΔW + """ + d, u, m = weights + R, I, *k = d.shape + O, R, *_ = u.shape + u = u * gamma + + if m is None: + result = u.reshape(-1, u.size(1)) @ d.reshape(d.size(0), -1) + else: + R, R, *k = m.shape + u = u.reshape(u.size(0), -1).transpose(0, 1) + d = d.reshape(d.size(0), -1) + result = rebuild_tucker(m, u, d) + return result.reshape(O, I, *k) + + +def bypass_forward_diff(x, org_out, *weights, gamma=1.0, extra_args={}): + """### bypass_forward_diff + + Args: + x (torch.Tensor): input tensor + weights (tuple[torch.Tensor]): (down, up[, mid]) + gamma (float, optional): scale factor, normally alpha/rank here + extra_args (dict, optional): extra args for forward func, \ + e.g. padding, stride for Conv1/2/3d + + Returns: + torch.Tensor: output tensor + """ + d, u, m = weights + if m is not None: + down = FUNC_LIST[d.dim()](x, d) + mid = FUNC_LIST[d.dim()](down, m, **extra_args) + up = FUNC_LIST[d.dim()](mid, u) + else: + down = FUNC_LIST[d.dim()](x, d, **extra_args) + up = FUNC_LIST[d.dim()](down, u) + return up * gamma diff --git a/lycoris/functional/loha.py b/lycoris/functional/loha.py new file mode 100644 index 0000000..042e56b --- /dev/null +++ b/lycoris/functional/loha.py @@ -0,0 +1,165 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .general import FUNC_LIST + + +class HadaWeight(torch.autograd.Function): + @staticmethod + def forward(ctx, w1d, w1u, w2d, w2u, scale=torch.tensor(1)): + ctx.save_for_backward(w1d, w1u, w2d, w2u, scale) + diff_weight = ((w1u @ w1d) * (w2u @ w2d)) * scale + return diff_weight + + @staticmethod + def backward(ctx, grad_out): + (w1d, w1u, w2d, w2u, scale) = ctx.saved_tensors + grad_out = grad_out * scale + temp = grad_out * (w2u @ w2d) + grad_w1u = temp @ w1d.T + grad_w1d = w1u.T @ temp + + temp = grad_out * (w1u @ w1d) + grad_w2u = temp @ w2d.T + grad_w2d = w2u.T @ temp + + del temp + return grad_w1d, grad_w1u, grad_w2d, grad_w2u, None + + +class HadaWeightTucker(torch.autograd.Function): + @staticmethod + def forward(ctx, t1, w1d, w1u, t2, w2d, w2u, scale=torch.tensor(1)): + ctx.save_for_backward(t1, w1d, w1u, t2, w2d, w2u, scale) + + rebuild1 = torch.einsum("i j ..., j r, i p -> p r ...", t1, w1d, w1u) + rebuild2 = torch.einsum("i j ..., j r, i p -> p r ...", t2, w2d, w2u) + + return rebuild1 * rebuild2 * scale + + @staticmethod + def backward(ctx, grad_out): + (t1, w1d, w1u, t2, w2d, w2u, scale) = ctx.saved_tensors + grad_out = grad_out * scale + + temp = torch.einsum("i j ..., j r -> i r ...", t2, w2d) + rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w2u) + + grad_w = rebuild * grad_out + del rebuild + + grad_w1u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w) + grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w1u.T) + del grad_w, temp + + grad_w1d = torch.einsum("i r ..., i j ... -> r j", t1, grad_temp) + grad_t1 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w1d.T) + del grad_temp + + temp = torch.einsum("i j ..., j r -> i r ...", t1, w1d) + rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w1u) + + grad_w = rebuild * grad_out + del rebuild + + grad_w2u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w) + grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w2u.T) + del grad_w, temp + + grad_w2d = torch.einsum("i r ..., i j ... -> r j", t2, grad_temp) + grad_t2 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w2d.T) + del grad_temp + return grad_t1, grad_w1d, grad_w1u, grad_t2, grad_w2d, grad_w2u, None + + +def make_weight(w1d, w1u, w2d, w2u, scale): + return HadaWeight.apply(w1d, w1u, w2d, w2u, scale) + + +def make_weight_tucker(t1, w1d, w1u, t2, w2d, w2u, scale): + return HadaWeightTucker.apply(t1, w1d, w1u, t2, w2d, w2u, scale) + + +def weight_gen(org_weight, rank, tucker=True): + """### weight_gen + + Args: + org_weight (torch.Tensor): the weight tensor + rank (int): low rank + + Returns: + torch.Tensor: w1d, w2d, w1u, w2u[, t1, t2] + """ + out_dim, in_dim, *k = org_weight.shape + if k and tucker: + w1d = torch.empty(rank, in_dim) + w1u = torch.empty(rank, out_dim) + t1 = torch.empty(rank, rank, *k) + w2d = torch.empty(rank, in_dim) + w2u = torch.empty(rank, out_dim) + t2 = torch.empty(rank, rank, *k) + nn.init.normal_(t1, std=0.1) + nn.init.normal_(t2, std=0.1) + else: + w1d = torch.empty(rank, in_dim) + w1u = torch.empty(out_dim, rank) + w2d = torch.empty(rank, in_dim) + w2u = torch.empty(out_dim, rank) + t1 = t2 = None + nn.init.normal_(w1d, std=1) + nn.init.constant_(w1u, 0) + nn.init.normal_(w2d, std=1) + nn.init.normal_(w2u, std=0.1) + return w1d, w1u, w2d, w2u, t1, t2 + + +def diff_weight(*weights, gamma=1.0): + """### diff_weight + + Get ΔW = BA, where BA is low rank decomposition + + Args: + wegihts (tuple[torch.Tensor]): (w1d, w2d, w1u, w2u[, t1, t2]) + gamma (float, optional): scale factor, normally alpha/rank here + + Returns: + torch.Tensor: ΔW + """ + w1d, w1u, w2d, w2u, t1, t2 = weights + if t1 is not None and t2 is not None: + R, I = w1d.shape + R, O = w1u.shape + R, R, *k = t1.shape + result = make_weight_tucker(t1, w1d, w1u, t2, w2d, w2u, gamma) + else: + R, I, *k = w1d.shape + O, R, *_ = w1u.shape + w1d = w1d.reshape(w1d.size(0), -1) + w1u = w1u.reshape(-1, w1u.size(1)) + w2d = w2d.reshape(w2d.size(0), -1) + w2u = w2u.reshape(-1, w2u.size(1)) + result = make_weight(w1d, w1u, w2d, w2u, gamma) + + result = result.reshape(O, I, *k) + return result + + +def bypass_forward_diff(x, org_out, *weights, gamma=1.0, extra_args={}): + """### bypass_forward_diff + + Args: + x (torch.Tensor): input tensor + weights (tuple[torch.Tensor]): (w1d, w2d, w1u, w2u[, t1, t2]) + gamma (float, optional): scale factor, normally alpha/rank here + extra_args (dict, optional): extra args for forward func, \ + e.g. padding, stride for Conv1/2/3d + + Returns: + torch.Tensor: output tensor + """ + w1d, w1u, w2d, w2u, t1, t2 = weights + diff_w = diff_weight(w1d, w1u, w2d, w2u, t1, t2, gamma) + return FUNC_LIST[w1d.dim() if t1 is None else t1.dim()](x, diff_w, **extra_args) diff --git a/lycoris/functional/lokr.py b/lycoris/functional/lokr.py new file mode 100644 index 0000000..75720ed --- /dev/null +++ b/lycoris/functional/lokr.py @@ -0,0 +1,247 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .general import rebuild_tucker, FUNC_LIST +from .general import factorization + + +def make_kron(w1, w2, scale): + for _ in range(w2.dim() - w1.dim()): + w1 = w1.unsqueeze(-1) + w2 = w2.contiguous() + rebuild = torch.kron(w1, w2) + + if scale != 1: + rebuild = rebuild * scale + + return rebuild + + +def weight_gen( + org_weight, + rank, + tucker=True, + factor=-1, + decompose_both=False, + full_matrix=False, + unbalanced_factorization=False, +): + """### weight_gen + + Args: + org_weight (torch.Tensor): the weight tensor + rank (int): low rank + + Returns: + torch.Tensor | None: w1, w1a, w1b, w2, w2a, w2b, t2 + """ + out_dim, in_dim, *k = org_weight.shape + w1 = w1a = w1b = None + w2 = w2a = w2b = None + t2 = None + use_w1 = use_w2 = False + + if k: + k_size = k + shape = (out_dim, in_dim, *k_size) + + in_m, in_n = factorization(in_dim, factor) + out_l, out_k = factorization(out_dim, factor) + if unbalanced_factorization: + out_l, out_k = out_k, out_l + shape = ((out_l, out_k), (in_m, in_n), *k_size) # ((a, b), (c, d), *k_size) + tucker = tucker and any(i != 1 for i in k_size) + if ( + decompose_both + and rank < max(shape[0][0], shape[1][0]) / 2 + and not full_matrix + ): + w1a = torch.empty(shape[0][0], rank) + w1b = torch.empty(rank, shape[1][0]) + else: + use_w1 = True + w1 = torch.empty(shape[0][0], shape[1][0]) # a*c, 1-mode + + if rank >= max(shape[0][1], shape[1][1]) / 2 or full_matrix: + use_w2 = True + w2 = torch.empty(shape[0][1], shape[1][1], *k_size) + elif tucker: + t2 = torch.empty(rank, rank, *shape[2:]) + w2a = torch.empty(rank, shape[0][1]) # b, 1-mode + w2b = torch.empty(rank, shape[1][1]) # d, 2-mode + else: # Conv2d not tucker + # bigger part. weight and LoRA. [b, dim] x [dim, d*k1*k2] + w2a = torch.empty(shape[0][1], rank) + w2b = torch.empty(rank, shape[1][1], *shape[2:]) + # w1 ⊗ (w2a x w2b) = (a, b)⊗((c, dim)x(dim, d*k1*k2)) = (a, b)⊗(c, d*k1*k2) = (ac, bd*k1*k2) + else: # Linear + shape = (out_dim, in_dim) + + in_m, in_n = factorization(in_dim, factor) + out_l, out_k = factorization(out_dim, factor) + if unbalanced_factorization: + out_l, out_k = out_k, out_l + shape = ( + (out_l, out_k), + (in_m, in_n), + ) # ((a, b), (c, d)), out_dim = a*c, in_dim = b*d + # smaller part. weight scale + if decompose_both and rank < max(shape[0][0], shape[1][0]) / 2: + w1a = torch.empty(shape[0][0], rank) + w1b = torch.empty(rank, shape[1][0]) + else: + use_w1 = True + w1 = torch.empty(shape[0][0], shape[1][0]) # a*c, 1-mode + if rank < max(shape[0][1], shape[1][1]) / 2: + # bigger part. weight and LoRA. [b, dim] x [dim, d] + w2a = torch.empty(shape[0][1], rank) + w2b = torch.empty(rank, shape[1][1]) + # w1 ⊗ (w2a x w2b) = (a, b)⊗((c, dim)x(dim, d)) = (a, b)⊗(c, d) = (ac, bd) + else: + use_w2 = True + w2 = torch.empty(shape[0][1], shape[1][1]) + + if use_w2: + torch.nn.init.constant_(w2, 1) + else: + if tucker: + torch.nn.init.kaiming_uniform_(t2, a=math.sqrt(5)) + torch.nn.init.kaiming_uniform_(w2a, a=math.sqrt(5)) + torch.nn.init.constant_(w2b, 1) + + if use_w1: + torch.nn.init.kaiming_uniform_(w1, a=math.sqrt(5)) + else: + torch.nn.init.kaiming_uniform_(w1a, a=math.sqrt(5)) + torch.nn.init.kaiming_uniform_(w1b, a=math.sqrt(5)) + + return w1, w1a, w1b, w2, w2a, w2b, t2 + + +def diff_weight(*weights, gamma=1.0): + """### diff_weight + + Args: + weights (tuple[torch.Tensor]): (w1, w1a, w1b, w2, w2a, w2b, t) + gamma (float, optional): scale factor, normally alpha/rank here + + Returns: + torch.Tensor: ΔW + """ + w1, w1a, w1b, w2, w2a, w2b, t = weights + if w1a is not None: + rank = w1a.shape[1] + elif w2a is not None: + rank = w2a.shape[1] + else: + rank = gamma + scale = gamma / rank + if w1 is None: + w1 = w1a @ w1b + if w2 is None: + if t is None: + r, o, *k = w2b.shape + w2 = w2a @ w2b.view(r, -1) + w2 = w2.view(-1, o, *k) + else: + w2 = rebuild_tucker(t, w2a, w2b) + return make_kron(w1, w2, scale) + + +def bypass_forward_diff(h, org_out, *weights, gamma=1.0, extra_args={}): + """### bypass_forward_diff + + Args: + weights (tuple[torch.Tensor]): (w1, w1a, w1b, w2, w2a, w2b, t) + gamma (float, optional): scale factor, normally alpha/rank here + extra_args (dict, optional): extra args for forward func, \ + e.g. padding, stride for Conv1/2/3d + + Returns: + torch.Tensor: output tensor + """ + w1, w1a, w1b, w2, w2a, w2b, t = weights + use_w1 = w1 is not None + use_w2 = w2 is not None + tucker = t is not None + dim = t.dim() if tucker else w2.dim() if w2 is not None else w2b.dim() + rank = w1b.size(0) if not use_w1 else w2b.size(0) if not use_w2 else gamma + scale = gamma / rank + is_conv = dim > 2 + op = FUNC_LIST[dim] + + if is_conv: + kw_dict = extra_args + else: + kw_dict = {} + + if use_w2: + ba = w2 + else: + a = w2b + b = w2a + + if t is not None: + a = a.view(*a.shape, *[1] * (dim - 2)) + b = b.view(*b.shape, *[1] * (dim - 2)) + elif is_conv: + b = b.view(*b.shape, *[1] * (dim - 2)) + + if use_w1: + c = w1 + else: + c = w1a @ w1b + uq = c.size(1) + + if is_conv: + # (b, uq), vq, ... + B, _, *rest = h.shape + h_in_group = h.reshape(B * uq, -1, *rest) + else: + # b, ..., uq, vq + h_in_group = h.reshape(*h.shape[:-1], uq, -1) + + if use_w2: + hb = op(h_in_group, ba, **kw_dict) + else: + if is_conv: + if tucker: + ha = op(h_in_group, a) + ht = op(ha, t, **kw_dict) + hb = op(ht, b) + else: + ha = op(h_in_group, a, **kw_dict) + hb = op(ha, b) + else: + ha = op(h_in_group, a, **kw_dict) + hb = op(ha, b) + + if is_conv: + # (b, uq), vp, ..., f + # -> b, uq, vp, ..., f + # -> b, f, vp, ..., uq + hb = hb.view(B, -1, *hb.shape[1:]) + h_cross_group = hb.transpose(1, -1) + else: + # b, ..., uq, vq + # -> b, ..., vq, uq + h_cross_group = hb.transpose(-1, -2) + + hc = F.linear(h_cross_group, c) + if is_conv: + # b, f, vp, ..., up + # -> b, up, vp, ... ,f + # -> b, c, ..., f + hc = hc.transpose(1, -1) + h = hc.reshape(B, -1, *hc.shape[3:]) + else: + # b, ..., vp, up + # -> b, ..., up, vp + # -> b, ..., c + hc = hc.transpose(-1, -2) + h = hc.reshape(*hc.shape[:-2], -1) + + return h * scale diff --git a/lycoris/kohya.py b/lycoris/kohya.py new file mode 100644 index 0000000..3f31872 --- /dev/null +++ b/lycoris/kohya.py @@ -0,0 +1,676 @@ +import os +import fnmatch +import re +import logging + +from typing import Any, List + +import torch + +from .utils import precalculate_safetensors_hashes +from .wrapper import LycorisNetwork, network_module_dict, deprecated_arg_dict +from .modules.locon import LoConModule +from .modules.loha import LohaModule +from .modules.ia3 import IA3Module +from .modules.lokr import LokrModule +from .modules.dylora import DyLoraModule +from .modules.glora import GLoRAModule +from .modules.norms import NormModule +from .modules.full import FullModule +from .modules.diag_oft import DiagOFTModule +from .modules.boft import ButterflyOFTModule +from .modules import make_module, get_module + +from .config import PRESET +from .utils.preset import read_preset +from .utils import str_bool +from .logging import logger + + +def create_network( + multiplier, network_dim, network_alpha, vae, text_encoder, unet, **kwargs +): + for key, value in list(kwargs.items()): + if key in deprecated_arg_dict: + logger.warning( + f"{key} is deprecated. Please use {deprecated_arg_dict[key]} instead.", + stacklevel=2, + ) + kwargs[deprecated_arg_dict[key]] = value + if network_dim is None: + network_dim = 4 # default + conv_dim = int(kwargs.get("conv_dim", network_dim) or network_dim) + conv_alpha = float(kwargs.get("conv_alpha", network_alpha) or network_alpha) + dropout = float(kwargs.get("dropout", 0.0) or 0.0) + rank_dropout = float(kwargs.get("rank_dropout", 0.0) or 0.0) + module_dropout = float(kwargs.get("module_dropout", 0.0) or 0.0) + algo = (kwargs.get("algo", "lora") or "lora").lower() + use_tucker = str_bool( + not kwargs.get("disable_conv_cp", True) + or kwargs.get("use_conv_cp", False) + or kwargs.get("use_cp", False) + or kwargs.get("use_tucker", False) + ) + use_scalar = str_bool(kwargs.get("use_scalar", False)) + block_size = int(kwargs.get("block_size", None) or 4) + train_norm = str_bool(kwargs.get("train_norm", False)) + constraint = float(kwargs.get("constraint", None) or 0) + rescaled = str_bool(kwargs.get("rescaled", False)) + weight_decompose = str_bool(kwargs.get("dora_wd", False)) + wd_on_output = str_bool(kwargs.get("wd_on_output", False)) + full_matrix = str_bool(kwargs.get("full_matrix", False)) + bypass_mode = str_bool(kwargs.get("bypass_mode", None)) + rs_lora = str_bool(kwargs.get("rs_lora", False)) + unbalanced_factorization = str_bool(kwargs.get("unbalanced_factorization", False)) + train_t5xxl = str_bool(kwargs.get("train_t5xxl", False)) + + if unbalanced_factorization: + logger.info("Unbalanced factorization for LoKr is enabled") + + if bypass_mode: + logger.info("Bypass mode is enabled") + + if weight_decompose: + logger.info("Weight decomposition is enabled") + + if full_matrix: + logger.info("Full matrix mode for LoKr is enabled") + + preset_str = kwargs.get("preset", "full") + if preset_str not in PRESET: + preset = read_preset(preset_str) + else: + preset = PRESET[preset_str] + assert preset is not None + LycorisNetworkKohya.apply_preset(preset) + + logger.info(f"Using rank adaptation algo: {algo}") + + if algo == "ia3" and preset_str != "ia3": + logger.warning("It is recommended to use preset ia3 for IA^3 algorithm") + + network = LycorisNetworkKohya( + text_encoder, + unet, + multiplier=multiplier, + lora_dim=network_dim, + conv_lora_dim=conv_dim, + alpha=network_alpha, + conv_alpha=conv_alpha, + dropout=dropout, + rank_dropout=rank_dropout, + module_dropout=module_dropout, + use_tucker=use_tucker, + use_scalar=use_scalar, + network_module=algo, + train_norm=train_norm, + decompose_both=kwargs.get("decompose_both", False), + factor=kwargs.get("factor", -1), + block_size=block_size, + constraint=constraint, + rescaled=rescaled, + weight_decompose=weight_decompose, + wd_on_out=wd_on_output, + full_matrix=full_matrix, + bypass_mode=bypass_mode, + rs_lora=rs_lora, + unbalanced_factorization=unbalanced_factorization, + train_t5xxl=train_t5xxl, + ) + + return network + + +def create_network_from_weights( + multiplier, + file, + vae, + text_encoder, + unet, + weights_sd=None, + for_inference=False, + **kwargs, +): + if weights_sd is None: + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import load_file, safe_open + + weights_sd = load_file(file) + else: + weights_sd = torch.load(file, map_location="cpu") + + # get dim/alpha mapping + unet_loras = {} + te_loras = {} + for key, value in weights_sd.items(): + if "." not in key: + continue + + lora_name = key.split(".")[0] + if lora_name.startswith(LycorisNetworkKohya.LORA_PREFIX_UNET): + unet_loras[lora_name] = None + elif lora_name.startswith(LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER): + te_loras[lora_name] = None + + for name, modules in unet.named_modules(): + lora_name = f"{LycorisNetworkKohya.LORA_PREFIX_UNET}_{name}".replace(".", "_") + if lora_name in unet_loras: + unet_loras[lora_name] = modules + + if isinstance(text_encoder, list): + text_encoders = text_encoder + use_index = True + else: + text_encoders = [text_encoder] + use_index = False + + for idx, te in enumerate(text_encoders): + if use_index: + prefix = f"{LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER}{idx+1}" + else: + prefix = LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER + for name, modules in te.named_modules(): + lora_name = f"{prefix}_{name}".replace(".", "_") + if lora_name in te_loras: + te_loras[lora_name] = modules + + original_level = logger.level + logger.setLevel(logging.ERROR) + network = LycorisNetworkKohya(text_encoder, unet) + network.unet_loras = [] + network.text_encoder_loras = [] + logger.setLevel(original_level) + + logger.info("Loading UNet Modules from state dict...") + for lora_name, orig_modules in unet_loras.items(): + if orig_modules is None: + continue + lyco_type, params = get_module(weights_sd, lora_name) + module = make_module(lyco_type, params, lora_name, orig_modules) + if module is not None: + network.unet_loras.append(module) + logger.info(f"{len(network.unet_loras)} Modules Loaded") + + logger.info("Loading TE Modules from state dict...") + for lora_name, orig_modules in te_loras.items(): + if orig_modules is None: + continue + lyco_type, params = get_module(weights_sd, lora_name) + module = make_module(lyco_type, params, lora_name, orig_modules) + if module is not None: + network.text_encoder_loras.append(module) + logger.info(f"{len(network.text_encoder_loras)} Modules Loaded") + + for lora in network.unet_loras + network.text_encoder_loras: + lora.multiplier = multiplier + + return network, weights_sd + + +class LycorisNetworkKohya(LycorisNetwork): + """ + LoRA + LoCon + """ + + # Ignore proj_in or proj_out, their channels is only a few. + ENABLE_CONV = True + UNET_TARGET_REPLACE_MODULE = [ + "Transformer2DModel", + "ResnetBlock2D", + "Downsample2D", + "Upsample2D", + "HunYuanDiTBlock", + "DoubleStreamBlock", + "SingleStreamBlock", + "SingleDiTBlock", + "MMDoubleStreamBlock", #HunYuanVideo + "MMSingleStreamBlock", #HunYuanVideo + ] + UNET_TARGET_REPLACE_NAME = [ + "conv_in", + "conv_out", + "time_embedding.linear_1", + "time_embedding.linear_2", + ] + TEXT_ENCODER_TARGET_REPLACE_MODULE = [ + "CLIPAttention", + "CLIPSdpaAttention", + "CLIPMLP", + "MT5Block", + "BertLayer", + ] + TEXT_ENCODER_TARGET_REPLACE_NAME = [] + LORA_PREFIX_UNET = "lora_unet" + LORA_PREFIX_TEXT_ENCODER = "lora_te" + MODULE_ALGO_MAP = {} + NAME_ALGO_MAP = {} + USE_FNMATCH = False + + @classmethod + def apply_preset(cls, preset): + if "enable_conv" in preset: + cls.ENABLE_CONV = preset["enable_conv"] + if "unet_target_module" in preset: + cls.UNET_TARGET_REPLACE_MODULE = preset["unet_target_module"] + if "unet_target_name" in preset: + cls.UNET_TARGET_REPLACE_NAME = preset["unet_target_name"] + if "text_encoder_target_module" in preset: + cls.TEXT_ENCODER_TARGET_REPLACE_MODULE = preset[ + "text_encoder_target_module" + ] + if "text_encoder_target_name" in preset: + cls.TEXT_ENCODER_TARGET_REPLACE_NAME = preset["text_encoder_target_name"] + if "module_algo_map" in preset: + cls.MODULE_ALGO_MAP = preset["module_algo_map"] + if "name_algo_map" in preset: + cls.NAME_ALGO_MAP = preset["name_algo_map"] + if "use_fnmatch" in preset: + cls.USE_FNMATCH = preset["use_fnmatch"] + return cls + + def __init__( + self, + text_encoder, + unet, + multiplier=1.0, + lora_dim=4, + conv_lora_dim=4, + alpha=1, + conv_alpha=1, + use_tucker=False, + dropout=0, + rank_dropout=0, + module_dropout=0, + network_module: str = "locon", + norm_modules=NormModule, + train_norm=False, + train_t5xxl=False, + **kwargs, + ) -> None: + torch.nn.Module.__init__(self) + root_kwargs = kwargs + self.multiplier = multiplier + self.lora_dim = lora_dim + self.train_t5xxl = train_t5xxl + + if not self.ENABLE_CONV: + conv_lora_dim = 0 + + self.conv_lora_dim = int(conv_lora_dim) + if self.conv_lora_dim and self.conv_lora_dim != self.lora_dim: + logger.info("Apply different lora dim for conv layer") + logger.info(f"Conv Dim: {conv_lora_dim}, Linear Dim: {lora_dim}") + elif self.conv_lora_dim == 0: + logger.info("Disable conv layer") + + self.alpha = alpha + self.conv_alpha = float(conv_alpha) + if self.conv_lora_dim and self.alpha != self.conv_alpha: + logger.info("Apply different alpha value for conv layer") + logger.info(f"Conv alpha: {conv_alpha}, Linear alpha: {alpha}") + + if 1 >= dropout >= 0: + logger.info(f"Use Dropout value: {dropout}") + self.dropout = dropout + self.rank_dropout = rank_dropout + self.module_dropout = module_dropout + + self.use_tucker = use_tucker + + def create_single_module( + lora_name: str, + module: torch.nn.Module, + algo_name, + dim=None, + alpha=None, + use_tucker=self.use_tucker, + **kwargs, + ): + for k, v in root_kwargs.items(): + if k in kwargs: + continue + kwargs[k] = v + + if train_norm and "Norm" in module.__class__.__name__: + return norm_modules( + lora_name, + module, + self.multiplier, + self.rank_dropout, + self.module_dropout, + **kwargs, + ) + lora = None + if isinstance(module, torch.nn.Linear) and lora_dim > 0: + dim = dim or lora_dim + alpha = alpha or self.alpha + elif isinstance( + module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d) + ): + k_size, *_ = module.kernel_size + if k_size == 1 and lora_dim > 0: + dim = dim or lora_dim + alpha = alpha or self.alpha + elif conv_lora_dim > 0 or dim: + dim = dim or conv_lora_dim + alpha = alpha or self.conv_alpha + else: + return None + else: + return None + lora = network_module_dict[algo_name]( + lora_name, + module, + self.multiplier, + dim, + alpha, + self.dropout, + self.rank_dropout, + self.module_dropout, + use_tucker, + **kwargs, + ) + return lora + + def create_modules_( + prefix: str, + root_module: torch.nn.Module, + algo, + configs={}, + ): + loras = {} + lora_names = [] + for name, module in root_module.named_modules(): + module_name = module.__class__.__name__ + if module_name in self.MODULE_ALGO_MAP and module is not root_module: + next_config = self.MODULE_ALGO_MAP[module_name] + next_algo = next_config.get("algo", algo) + new_loras, new_lora_names = create_modules_( + f"{prefix}_{name}", module, next_algo, next_config + ) + for lora_name, lora in zip(new_lora_names, new_loras): + if lora_name not in loras: + loras[lora_name] = lora + lora_names.append(lora_name) + continue + if name: + lora_name = prefix + "." + name + else: + lora_name = prefix + lora_name = lora_name.replace(".", "_") + if lora_name in loras: + continue + + lora = create_single_module(lora_name, module, algo, **configs) + if lora is not None: + loras[lora_name] = lora + lora_names.append(lora_name) + return [loras[lora_name] for lora_name in lora_names], lora_names + + # create module instances + def create_modules( + prefix, + root_module: torch.nn.Module, + target_replace_modules, + target_replace_names=[], + ) -> List: + logger.info("Create LyCORIS Module") + loras = [] + next_config = {} + for name, module in root_module.named_modules(): + module_name = module.__class__.__name__ + if module_name in target_replace_modules and not any( + self.match_fn(t, name) for t in target_replace_names + ): + if module_name in self.MODULE_ALGO_MAP: + next_config = self.MODULE_ALGO_MAP[module_name] + algo = next_config.get("algo", network_module) + else: + algo = network_module + loras.extend( + create_modules_(f"{prefix}_{name}", module, algo, next_config)[ + 0 + ] + ) + next_config = {} + elif name in target_replace_names or any( + self.match_fn(t, name) for t in target_replace_names + ): + conf_from_name = self.find_conf_for_name(name) + if conf_from_name is not None: + next_config = conf_from_name + algo = next_config.get("algo", network_module) + elif module_name in self.MODULE_ALGO_MAP: + next_config = self.MODULE_ALGO_MAP[module_name] + algo = next_config.get("algo", network_module) + else: + algo = network_module + lora_name = prefix + "." + name + lora_name = lora_name.replace(".", "_") + lora = create_single_module(lora_name, module, algo, **next_config) + next_config = {} + if lora is not None: + loras.append(lora) + return loras + + if network_module == GLoRAModule: + logger.info("GLoRA enabled, only train transformer") + # only train transformer (for GLoRA) + LycorisNetworkKohya.UNET_TARGET_REPLACE_MODULE = [ + "Transformer2DModel", + "Attention", + ] + LycorisNetworkKohya.UNET_TARGET_REPLACE_NAME = [] + + self.text_encoder_loras = [] + if text_encoder: + if isinstance(text_encoder, list): + text_encoders = text_encoder + use_index = True + else: + text_encoders = [text_encoder] + use_index = False + + for i, te in enumerate(text_encoders): + self.text_encoder_loras.extend( + create_modules( + LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER + + (f"{i+1}" if use_index else ""), + te, + LycorisNetworkKohya.TEXT_ENCODER_TARGET_REPLACE_MODULE, + LycorisNetworkKohya.TEXT_ENCODER_TARGET_REPLACE_NAME, + ) + ) + logger.info( + f"create LyCORIS for Text Encoder: {len(self.text_encoder_loras)} modules." + ) + + self.unet_loras = create_modules( + LycorisNetworkKohya.LORA_PREFIX_UNET, + unet, + LycorisNetworkKohya.UNET_TARGET_REPLACE_MODULE, + LycorisNetworkKohya.UNET_TARGET_REPLACE_NAME, + ) + logger.info(f"create LyCORIS for U-Net: {len(self.unet_loras)} modules.") + + algo_table = {} + for lora in self.text_encoder_loras + self.unet_loras: + algo_table[lora.__class__.__name__] = ( + algo_table.get(lora.__class__.__name__, 0) + 1 + ) + logger.info(f"module type table: {algo_table}") + + self.weights_sd = None + + self.loras = self.text_encoder_loras + self.unet_loras + # assertion + names = set() + for lora in self.loras: + assert ( + lora.lora_name not in names + ), f"duplicated lora name: {lora.lora_name}" + names.add(lora.lora_name) + + def match_fn(self, pattern: str, name: str) -> bool: + if self.USE_FNMATCH: + return fnmatch.fnmatch(name, pattern) + return re.match(pattern, name) + + def find_conf_for_name( + self, + name: str, + ) -> dict[str, Any]: + if name in self.NAME_ALGO_MAP.keys(): + return self.NAME_ALGO_MAP[name] + + for key, value in self.NAME_ALGO_MAP.items(): + if self.match_fn(key, name): + return value + + return None + + def load_weights(self, file): + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import load_file, safe_open + + self.weights_sd = load_file(file) + else: + self.weights_sd = torch.load(file, map_location="cpu") + missing, unexpected = self.load_state_dict(self.weights_sd, strict=False) + state = {} + if missing: + state["missing keys"] = missing + if unexpected: + state["unexpected keys"] = unexpected + return state + + def apply_to(self, text_encoder, unet, apply_text_encoder=None, apply_unet=None): + assert ( + apply_text_encoder is not None and apply_unet is not None + ), f"internal error: flag not set" + + if apply_text_encoder: + logger.info("enable LyCORIS for text encoder") + else: + self.text_encoder_loras = [] + + if apply_unet: + logger.info("enable LyCORIS for U-Net") + else: + self.unet_loras = [] + + self.loras = self.text_encoder_loras + self.unet_loras + + for lora in self.loras: + lora.apply_to() + self.add_module(lora.lora_name, lora) + + if self.weights_sd: + # if some weights are not in state dict, it is ok because initial LoRA does nothing (lora_up is initialized by zeros) + info = self.load_state_dict(self.weights_sd, False) + logger.info(f"weights are loaded: {info}") + + # TODO refactor to common function with apply_to + def merge_to(self, text_encoder, unet, weights_sd, dtype, device): + apply_text_encoder = apply_unet = False + for key in weights_sd.keys(): + if key.startswith(LycorisNetworkKohya.LORA_PREFIX_TEXT_ENCODER): + apply_text_encoder = True + elif key.startswith(LycorisNetworkKohya.LORA_PREFIX_UNET): + apply_unet = True + + if apply_text_encoder: + logger.info("enable LoRA for text encoder") + else: + self.text_encoder_loras = [] + + if apply_unet: + logger.info("enable LoRA for U-Net") + else: + self.unet_loras = [] + + self.loras = self.text_encoder_loras + self.unet_loras + super().merge_to(1) + + def apply_max_norm_regularization(self, max_norm_value, device): + key_scaled = 0 + norms = [] + for module in self.unet_loras + self.text_encoder_loras: + scaled, norm = module.apply_max_norm(max_norm_value, device) + if scaled is None: + continue + norms.append(norm) + key_scaled += scaled + + if key_scaled == 0: + return 0, 0, 0 + + return key_scaled, sum(norms) / len(norms), max(norms) + + def prepare_optimizer_params(self, text_encoder_lr=None, unet_lr: float = 1e-4, learning_rate=None): + def enumerate_params(loras): + params = [] + for lora in loras: + params.extend(lora.parameters()) + return params + + self.requires_grad_(True) + all_params = [] + lr_descriptions = [] + + if self.text_encoder_loras: + param_data = {"params": enumerate_params(self.text_encoder_loras)} + if text_encoder_lr is not None: + param_data["lr"] = text_encoder_lr + all_params.append(param_data) + lr_descriptions.append("text_encoder") + + if self.unet_loras: + param_data = {"params": enumerate_params(self.unet_loras)} + if unet_lr is not None: + param_data["lr"] = unet_lr + all_params.append(param_data) + lr_descriptions.append("unet") + + return all_params, lr_descriptions + + def enable_gradient_checkpointing(self): + # not supported + pass + + def prepare_grad_etc(self, text_encoder, unet): + self.requires_grad_(True) + + def on_epoch_start(self, text_encoder, unet): + self.train() + + #def on_step_start(self): + # pass + + def get_trainable_params(self): + return self.parameters() + + def save_weights(self, file, dtype, metadata): + if metadata is not None and len(metadata) == 0: + metadata = None + + state_dict = self.state_dict() + + if dtype is not None: + for key in list(state_dict.keys()): + v = state_dict[key] + v = v.detach().clone().to("cpu").to(dtype) + state_dict[key] = v + + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import save_file + + # Precalculate model hashes to save time on indexing + if metadata is None: + metadata = {} + model_hash = precalculate_safetensors_hashes(state_dict) + metadata["sshs_model_hash"] = model_hash + + save_file(state_dict, file, metadata) + else: + torch.save(state_dict, file) diff --git a/lycoris/logging.py b/lycoris/logging.py new file mode 100644 index 0000000..6c51ccc --- /dev/null +++ b/lycoris/logging.py @@ -0,0 +1,52 @@ +import sys +import copy +import logging +from functools import cache + + +class ColoredFormatter(logging.Formatter): + COLORS = { + "DEBUG": "\033[0;36m", # CYAN + "INFO": "\033[0;32m", # GREEN + "WARNING": "\033[0;33m", # YELLOW + "ERROR": "\033[0;31m", # RED + "CRITICAL": "\033[0;37;41m", # WHITE ON RED + "RESET": "\033[0m", # RESET COLOR + } + + def format(self, record): + colored_record = copy.copy(record) + levelname = colored_record.levelname + seq = self.COLORS.get(levelname, self.COLORS["RESET"]) + colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}" + return super().format(colored_record) + + +logger = logging.getLogger("LyCORIS") +logger.propagate = False +logger.setLevel(logging.INFO) + + +if not logger.handlers: + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter( + ColoredFormatter( + "%(asctime)s|[%(name)s]-%(levelname)s: %(message)s", "%Y-%m-%d %H:%M:%S" + ) + ) + logger.addHandler(handler) + + +@cache +def info_once(msg): + logger.info(msg) + + +@cache +def warning_once(msg): + logger.warning(msg) + + +@cache +def error_once(msg): + logger.error(msg) diff --git a/lycoris/modules/__init__.py b/lycoris/modules/__init__.py new file mode 100644 index 0000000..9d0c8bd --- /dev/null +++ b/lycoris/modules/__init__.py @@ -0,0 +1,46 @@ +import torch +import torch.nn as nn + +from .base import LycorisBaseModule +from .locon import LoConModule +from .loha import LohaModule +from .lokr import LokrModule +from .full import FullModule +from .norms import NormModule +from .diag_oft import DiagOFTModule +from .boft import ButterflyOFTModule +from .glora import GLoRAModule +from .dylora import DyLoraModule +from .ia3 import IA3Module + +from ..functional.general import factorization + + +MODULE_LIST = [ + LoConModule, + LohaModule, + IA3Module, + LokrModule, + FullModule, + NormModule, + DiagOFTModule, + ButterflyOFTModule, + GLoRAModule, + DyLoraModule, +] + + +def get_module(lyco_state_dict, lora_name): + for module in MODULE_LIST: + if module.algo_check(lyco_state_dict, lora_name): + return module, tuple(module.extract_state_dict(lyco_state_dict, lora_name)) + return None, None + + +@torch.no_grad() +def make_module(lyco_type: LycorisBaseModule, params, lora_name, orig_module): + try: + module = lyco_type.make_module_from_state_dict(lora_name, orig_module, *params) + except NotImplementedError: + module = None + return module diff --git a/lycoris/modules/base.py b/lycoris/modules/base.py new file mode 100644 index 0000000..e8dec35 --- /dev/null +++ b/lycoris/modules/base.py @@ -0,0 +1,315 @@ +from collections import OrderedDict + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.nn.utils.parametrize as parametrize + +from ..utils.quant import QuantLinears, log_bypass, log_suspect + + +class ModuleCustomSD(nn.Module): + def __init__(self): + super().__init__() + self._register_load_state_dict_pre_hook(self.load_weight_prehook) + self.register_load_state_dict_post_hook(self.load_weight_hook) + + def load_weight_prehook( + self, + state_dict, + prefix, + local_metadata, + strict, + missing_keys, + unexpected_keys, + error_msgs, + ): + pass + + def load_weight_hook(self, module, incompatible_keys): + pass + + def custom_state_dict(self): + return None + + def state_dict(self, *args, destination=None, prefix="", keep_vars=False): + # TODO: Remove `args` and the parsing logic when BC allows. + if len(args) > 0: + if destination is None: + destination = args[0] + if len(args) > 1 and prefix == "": + prefix = args[1] + if len(args) > 2 and keep_vars is False: + keep_vars = args[2] + # DeprecationWarning is ignored by default + + if destination is None: + destination = OrderedDict() + destination._metadata = OrderedDict() + + local_metadata = dict(version=self._version) + if hasattr(destination, "_metadata"): + destination._metadata[prefix[:-1]] = local_metadata + + if (custom_sd := self.custom_state_dict()) is not None: + for k, v in custom_sd.items(): + destination[f"{prefix}{k}"] = v + return destination + else: + return super().state_dict( + *args, destination=destination, prefix=prefix, keep_vars=keep_vars + ) + + +class LycorisBaseModule(ModuleCustomSD): + name: str + dtype_tensor: torch.Tensor + support_module = {} + weight_list = [] + weight_list_det = [] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + rank_dropout_scale=False, + bypass_mode=None, + **kwargs, + ): + """if alpha == 0 or None, alpha is rank (no scaling).""" + super().__init__() + self.lora_name = lora_name + self.not_supported = False + + self.module = type(org_module) + if isinstance(org_module, nn.Linear): + self.module_type = "linear" + self.shape = (org_module.out_features, org_module.in_features) + self.op = F.linear + self.dim = org_module.out_features + self.kw_dict = {} + elif isinstance(org_module, nn.Conv1d): + self.module_type = "conv1d" + self.shape = ( + org_module.out_channels, + org_module.in_channels, + *org_module.kernel_size, + ) + self.op = F.conv1d + self.dim = org_module.out_channels + self.kw_dict = { + "stride": org_module.stride, + "padding": org_module.padding, + "dilation": org_module.dilation, + "groups": org_module.groups, + } + elif isinstance(org_module, nn.Conv2d): + self.module_type = "conv2d" + self.shape = ( + org_module.out_channels, + org_module.in_channels, + *org_module.kernel_size, + ) + self.op = F.conv2d + self.dim = org_module.out_channels + self.kw_dict = { + "stride": org_module.stride, + "padding": org_module.padding, + "dilation": org_module.dilation, + "groups": org_module.groups, + } + elif isinstance(org_module, nn.Conv3d): + self.module_type = "conv3d" + self.shape = ( + org_module.out_channels, + org_module.in_channels, + *org_module.kernel_size, + ) + self.op = F.conv3d + self.dim = org_module.out_channels + self.kw_dict = { + "stride": org_module.stride, + "padding": org_module.padding, + "dilation": org_module.dilation, + "groups": org_module.groups, + } + elif isinstance(org_module, nn.LayerNorm): + self.module_type = "layernorm" + self.shape = tuple(org_module.normalized_shape) + self.op = F.layer_norm + self.dim = org_module.normalized_shape[0] + self.kw_dict = { + "normalized_shape": org_module.normalized_shape, + "eps": org_module.eps, + } + elif isinstance(org_module, nn.GroupNorm): + self.module_type = "groupnorm" + self.shape = (org_module.num_channels,) + self.op = F.group_norm + self.group_num = org_module.num_groups + self.dim = org_module.num_channels + self.kw_dict = {"num_groups": org_module.num_groups, "eps": org_module.eps} + else: + self.not_supported = True + self.module_type = "unknown" + + self.register_buffer("dtype_tensor", torch.tensor(0.0), persistent=False) + + self.is_quant = False + if isinstance(org_module, QuantLinears): + if not bypass_mode: + log_bypass() + self.is_quant = True + bypass_mode = True + if ( + isinstance(org_module, nn.Linear) + and org_module.__class__.__name__ != "Linear" + ): + if bypass_mode is None: + log_suspect() + bypass_mode = True + if bypass_mode == True: + self.is_quant = True + self.bypass_mode = bypass_mode + self.dropout = dropout + self.rank_dropout = rank_dropout + self.rank_dropout_scale = rank_dropout_scale + self.module_dropout = module_dropout + + ## Dropout things + # Since LoKr/LoHa/OFT/BOFT are hard to follow the rank_dropout definition from kohya + # We redefine the dropout procedure here. + # g(x) = WX + drop(Brank_drop(AX)) for LoCon(lora), bypass + # g(x) = WX + drop(ΔWX) for any algo except LoCon(lora), bypass + # g(x) = (W + Brank_drop(A))X for LoCon(lora), rebuid + # g(x) = (W + rank_drop(ΔW))X for any algo except LoCon(lora), rebuild + self.drop = nn.Identity() if dropout == 0 else nn.Dropout(dropout) + self.rank_drop = ( + nn.Identity() if rank_dropout == 0 else nn.Dropout(rank_dropout) + ) + + self.multiplier = multiplier + self.org_forward = org_module.forward + self.org_module = [org_module] + + @classmethod + def parametrize(cls, org_module, attr, *args, **kwargs): + from .full import FullModule + + if cls is FullModule: + raise RuntimeError("FullModule cannot be used for parametrize.") + target_param = getattr(org_module, attr) + kwargs["bypass_mode"] = False + if target_param.dim() == 2: + proxy_module = nn.Linear( + target_param.shape[0], target_param.shape[1], bias=False + ) + proxy_module.weight = target_param + elif target_param.dim() > 2: + module_type = [ + None, + None, + None, + nn.Conv1d, + nn.Conv2d, + nn.Conv3d, + None, + None, + ][target_param.dim()] + proxy_module = module_type( + target_param.shape[0], + target_param.shape[1], + *target_param.shape[2:], + bias=False, + ) + proxy_module.weight = target_param + module_obj = cls("", proxy_module, *args, **kwargs) + module_obj.forward = module_obj.parametrize_forward + module_obj.to(target_param) + parametrize.register_parametrization(org_module, attr, module_obj) + return module_obj + + @classmethod + def algo_check(cls, state_dict, lora_name): + return any(f"{lora_name}.{k}" in state_dict for k in cls.weight_list_det) + + @classmethod + def extract_state_dict(cls, state_dict, lora_name): + return [state_dict.get(f"{lora_name}.{k}", None) for k in cls.weight_list] + + @classmethod + def make_module_from_state_dict(cls, lora_name, orig_module, *weights): + raise NotImplementedError + + @property + def dtype(self): + return self.dtype_tensor.dtype + + @property + def device(self): + return self.dtype_tensor.device + + @property + def org_weight(self): + return self.org_module[0].weight + + @org_weight.setter + def org_weight(self, value): + self.org_module[0].weight.data.copy_(value) + + def apply_to(self, **kwargs): + if self.not_supported: + return + self.org_forward = self.org_module[0].forward + self.org_module[0].forward = self.forward + + def restore(self): + if self.not_supported: + return + self.org_module[0].forward = self.org_forward + + def merge_to(self, multiplier=1.0): + if self.not_supported: + return + self_device = next(self.parameters()).device + self_dtype = next(self.parameters()).dtype + self.to(self.org_weight) + weight, bias = self.get_merged_weight( + multiplier, self.org_weight.shape, self.org_weight.device + ) + self.org_weight = weight.to(self.org_weight) + if bias is not None: + bias = bias.to(self.org_weight) + if self.org_module[0].bias is not None: + self.org_module[0].bias.data.copy_(bias) + else: + self.org_module[0].bias = nn.Parameter(bias) + self.to(self_device, self_dtype) + + def get_diff_weight(self, multiplier=1.0, shape=None, device=None): + raise NotImplementedError + + def get_merged_weight(self, multiplier=1.0, shape=None, device=None): + raise NotImplementedError + + @torch.no_grad() + def apply_max_norm(self, max_norm, device=None): + return None, None + + def bypass_forward_diff(self, x, scale=1): + raise NotImplementedError + + def bypass_forward(self, x, scale=1): + raise NotImplementedError + + def parametrize_forward(self, x: torch.Tensor, *args, **kwargs): + return self.get_merged_weight( + multiplier=self.multiplier, shape=x.shape, device=x.device + )[0].to(x.dtype) + + def forward(self, *args, **kwargs): + raise NotImplementedError diff --git a/lycoris/modules/boft.py b/lycoris/modules/boft.py new file mode 100644 index 0000000..c547527 --- /dev/null +++ b/lycoris/modules/boft.py @@ -0,0 +1,255 @@ +from functools import cache +from math import log2 + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from einops import rearrange + +from .base import LycorisBaseModule +from ..functional import power2factorization +from ..logging import logger + + +@cache +def log_butterfly_factorize(dim, factor, result): + logger.info( + f"Use BOFT({int(log2(result[1]))}, {result[0]//2})" + f" (equivalent to factor={result[0]}) " + f"for {dim=} and {factor=}" + ) + + +def butterfly_factor(dimension: int, factor: int = -1) -> tuple[int, int]: + m, n = power2factorization(dimension, factor) + + if n == 0: + raise ValueError( + f"It is impossible to decompose {dimension} with factor {factor} under BOFT constraints." + ) + + log_butterfly_factorize(dimension, factor, (m, n)) + return m, n + + +class ButterflyOFTModule(LycorisBaseModule): + name = "boft" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = [ + "oft_blocks", + "rescale", + "alpha", + ] + weight_list_det = ["oft_blocks"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + rank_dropout_scale=False, + constraint=0, + rescaled=False, + bypass_mode=None, + **kwargs, + ): + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in BOFT algo.") + + out_dim = self.dim + b, m_exp = butterfly_factor(out_dim, lora_dim) + self.block_size = b + self.block_num = m_exp + # BOFT(m, b) + self.boft_b = b + self.boft_m = sum(int(i) for i in f"{m_exp-1:b}") + 1 + # block_num > block_size + self.rescaled = rescaled + self.constraint = constraint * out_dim + self.register_buffer("alpha", torch.tensor(constraint)) + self.oft_blocks = nn.Parameter( + torch.zeros(self.boft_m, self.block_num, self.block_size, self.block_size) + ) + if rescaled: + self.rescale = nn.Parameter( + torch.ones(out_dim, *(1 for _ in range(org_module.weight.dim() - 1))) + ) + + @classmethod + def algo_check(cls, state_dict, lora_name): + if f"{lora_name}.oft_blocks" in state_dict: + oft_blocks = state_dict[f"{lora_name}.oft_blocks"] + if oft_blocks.ndim == 4: + return True + return False + + @classmethod + def make_module_from_state_dict( + cls, lora_name, orig_module, oft_blocks, rescale, alpha + ): + m, n, s, _ = oft_blocks.shape + module = cls( + lora_name, + orig_module, + 1, + lora_dim=s, + constraint=float(alpha), + rescaled=rescale is not None, + ) + module.oft_blocks.copy_(oft_blocks) + if rescale is not None: + module.rescale.copy_(rescale) + return module + + @property + def I(self): + return torch.eye(self.block_size, device=self.device) + + def get_r(self): + I = self.I + # for Q = -Q^T + q = self.oft_blocks - self.oft_blocks.transpose(-1, -2) + normed_q = q + # Diag OFT style constrain + if self.constraint > 0: + q_norm = torch.norm(q) + 1e-8 + if q_norm > self.constraint: + normed_q = q * self.constraint / q_norm + # use float() to prevent unsupported type + r = (I + normed_q) @ (I - normed_q).float().inverse() + return r + + def make_weight(self, scale=1, device=None, diff=False): + m = self.boft_m + b = self.boft_b + r_b = b // 2 + r = self.get_r() + inp = org = self.org_weight.to(device, dtype=r.dtype) + + for i in range(m): + bi = r[i] # b_num, b_size, b_size + g = 2 + k = 2**i * r_b + if scale != 1: + bi = bi * scale + (1 - scale) * self.I + inp = ( + inp.unflatten(-1, (-1, g, k)) + .transpose(-2, -1) + .flatten(-3) + .unflatten(-1, (-1, b)) + ) + inp = torch.einsum("b i j, b j ... -> b i ...", bi, inp) + inp = ( + inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3) + ) + + if self.rescaled: + inp = inp * self.rescale + + if diff: + inp = inp - org + + return inp.to(self.oft_blocks.dtype) + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + diff = self.make_weight(scale=multiplier, device=device, diff=True) + if shape is not None: + diff = diff.view(shape) + return diff, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + diff = self.make_weight(scale=multiplier, device=device) + if shape is not None: + diff = diff.view(shape) + return diff, None + + @torch.no_grad() + def apply_max_norm(self, max_norm, device=None): + orig_norm = self.oft_blocks.to(device).norm() + norm = torch.clamp(orig_norm, max_norm / 2) + desired = torch.clamp(norm, max=max_norm) + ratio = desired / norm + + scaled = norm != desired + if scaled: + self.oft_blocks *= ratio + + return scaled, orig_norm * ratio + + def _bypass_forward(self, x, scale=1, diff=False): + m = self.boft_m + b = self.boft_b + r_b = b // 2 + r = self.get_r() + inp = org = self.org_forward(x) + if self.op in {F.conv2d, F.conv1d, F.conv3d}: + inp = inp.transpose(1, -1) + + for i in range(m): + bi = r[i] # b_num, b_size, b_size + g = 2 + k = 2**i * r_b + if scale != 1: + bi = bi * scale + (1 - scale) * self.I + inp = ( + inp.unflatten(-1, (-1, g, k)) + .transpose(-2, -1) + .flatten(-3) + .unflatten(-1, (-1, b)) + ) + inp = torch.einsum("b i j, ... b j -> ... b i", bi, inp) + inp = ( + inp.flatten(-2).unflatten(-1, (-1, k, g)).transpose(-2, -1).flatten(-3) + ) + + if self.rescaled: + inp = inp * self.rescale.transpose(0, -1) + + if self.op in {F.conv2d, F.conv1d, F.conv3d}: + inp = inp.transpose(1, -1) + + if diff: + inp = inp - org + return inp + + def bypass_forward_diff(self, x, scale=1): + return self._bypass_forward(x, scale, diff=True) + + def bypass_forward(self, x, scale=1): + return self._bypass_forward(x, scale, diff=False) + + def forward(self, x, *args, **kwargs): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.org_forward(x) + scale = self.multiplier + + if self.bypass_mode: + return self.bypass_forward(x, scale) + else: + w = self.make_weight(scale, x.device) + kw_dict = self.kw_dict | {"weight": w, "bias": self.org_module[0].bias} + return self.op(x, **kw_dict) diff --git a/lycoris/modules/diag_oft.py b/lycoris/modules/diag_oft.py new file mode 100644 index 0000000..3805995 --- /dev/null +++ b/lycoris/modules/diag_oft.py @@ -0,0 +1,217 @@ +from functools import cache + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .base import LycorisBaseModule +from ..functional import factorization +from ..logging import logger + + +@cache +def log_oft_factorize(dim, factor, num, bdim): + logger.info( + f"Use OFT(block num: {num}, block dim: {bdim})" + f" (equivalent to lora_dim={num}) " + f"for {dim=} and lora_dim={factor=}" + ) + + +class DiagOFTModule(LycorisBaseModule): + name = "diag-oft" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = [ + "oft_blocks", + "rescale", + "alpha", + ] + weight_list_det = ["oft_blocks"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + rank_dropout_scale=False, + constraint=0, + rescaled=False, + bypass_mode=None, + **kwargs, + ): + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in Diag-OFT algo.") + + out_dim = self.dim + self.block_size, self.block_num = factorization(out_dim, lora_dim) + # block_num > block_size + self.rescaled = rescaled + self.constraint = constraint * out_dim + self.register_buffer("alpha", torch.tensor(constraint)) + self.oft_blocks = nn.Parameter( + torch.zeros(self.block_num, self.block_size, self.block_size) + ) + if rescaled: + self.rescale = nn.Parameter( + torch.ones(out_dim, *(1 for _ in range(org_module.weight.dim() - 1))) + ) + + log_oft_factorize( + dim=out_dim, + factor=lora_dim, + num=self.block_num, + bdim=self.block_size, + ) + + @classmethod + def algo_check(cls, state_dict, lora_name): + if f"{lora_name}.oft_blocks" in state_dict: + oft_blocks = state_dict[f"{lora_name}.oft_blocks"] + if oft_blocks.ndim == 3: + return True + return False + + @classmethod + def make_module_from_state_dict( + cls, lora_name, orig_module, oft_blocks, rescale, alpha + ): + n, s, _ = oft_blocks.shape + module = cls( + lora_name, + orig_module, + 1, + lora_dim=s, + constraint=float(alpha), + rescaled=rescale is not None, + ) + module.oft_blocks.copy_(oft_blocks) + if rescale is not None: + module.rescale.copy_(rescale) + return module + + @property + def I(self): + return torch.eye(self.block_size, device=self.device) + + def get_r(self): + I = self.I + # for Q = -Q^T + q = self.oft_blocks - self.oft_blocks.transpose(1, 2) + normed_q = q + if self.constraint > 0: + q_norm = torch.norm(q) + 1e-8 + if q_norm > self.constraint: + normed_q = q * self.constraint / q_norm + # use float() to prevent unsupported type + r = (I + normed_q) @ (I - normed_q).float().inverse() + return r + + def make_weight(self, scale=1, device=None, diff=False): + r = self.get_r() + _, *shape = self.org_weight.shape + org_weight = self.org_weight.to(device, dtype=r.dtype) + org_weight = org_weight.view(self.block_num, self.block_size, *shape) + # Init R=0, so add I on it to ensure the output of step0 is original model output + weight = torch.einsum( + "k n m, k n ... -> k m ...", + self.rank_drop(r * scale) - scale * self.I + (0 if diff else self.I), + org_weight, + ).view(-1, *shape) + if self.rescaled: + weight = self.rescale * weight + if diff: + weight = weight + (self.rescale - 1) * org_weight + return weight.to(self.oft_blocks.dtype) + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + diff = self.make_weight(scale=multiplier, device=device, diff=True) + if shape is not None: + diff = diff.view(shape) + return diff, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + diff = self.make_weight(scale=multiplier, device=device) + if shape is not None: + diff = diff.view(shape) + return diff, None + + @torch.no_grad() + def apply_max_norm(self, max_norm, device=None): + orig_norm = self.oft_blocks.to(device).norm() + norm = torch.clamp(orig_norm, max_norm / 2) + desired = torch.clamp(norm, max=max_norm) + ratio = desired / norm + + scaled = norm != desired + if scaled: + self.oft_blocks *= ratio + + return scaled, orig_norm * ratio + + def _bypass_forward(self, x, scale=1, diff=False): + r = self.get_r() + org_out = self.org_forward(x) + if self.op in {F.conv2d, F.conv1d, F.conv3d}: + org_out = org_out.transpose(1, -1) + *shape, _ = org_out.shape + org_out = org_out.view(*shape, self.block_num, self.block_size) + mask = neg_mask = 1 + if self.dropout != 0 and self.training: + mask = torch.ones_like(org_out) + mask = self.drop(mask) + neg_mask = torch.max(mask) - mask + oft_out = torch.einsum( + "k n m, ... k n -> ... k m", + r * scale * mask + (1 - scale) * self.I * neg_mask, + org_out, + ) + if diff: + out = out - org_out + out = oft_out.view(*shape, -1) + if self.rescaled: + out = self.rescale.transpose(-1, 0) * out + out = out + (self.rescale.transpose(-1, 0) - 1) * org_out + if self.op in {F.conv2d, F.conv1d, F.conv3d}: + out = out.transpose(1, -1) + return out + + def bypass_forward_diff(self, x, scale=1): + return self._bypass_forward(x, scale, diff=True) + + def bypass_forward(self, x, scale=1): + return self._bypass_forward(x, scale, diff=False) + + def forward(self, x: torch.Tensor, *args, **kwargs): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.org_forward(x) + scale = self.multiplier + + if self.bypass_mode: + return self.bypass_forward(x, scale) + else: + w = self.make_weight(scale, x.device) + kw_dict = self.kw_dict | {"weight": w, "bias": self.org_module[0].bias} + return self.op(x, **kw_dict) diff --git a/lycoris/modules/dylora.py b/lycoris/modules/dylora.py new file mode 100644 index 0000000..4138142 --- /dev/null +++ b/lycoris/modules/dylora.py @@ -0,0 +1,156 @@ +import math +import random + +import torch +import torch.nn as nn + +from .base import LycorisBaseModule +from ..utils import product + + +class DyLoraModule(LycorisBaseModule): + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + block_size=4, + use_scalar=False, + rank_dropout_scale=False, + weight_decompose=False, + bypass_mode=None, + rs_lora=False, + train_on_input=False, + **kwargs, + ): + """if alpha == 0 or None, alpha is rank (no scaling).""" + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in IA^3 algo.") + assert lora_dim % block_size == 0, "lora_dim must be a multiple of block_size" + self.block_count = lora_dim // block_size + self.block_size = block_size + + shape = ( + self.shape[0], + product(self.shape[1:]), + ) + + self.lora_dim = lora_dim + self.up_list = nn.ParameterList( + [torch.empty(shape[0], self.block_size) for i in range(self.block_count)] + ) + self.down_list = nn.ParameterList( + [torch.empty(self.block_size, shape[1]) for i in range(self.block_count)] + ) + + if type(alpha) == torch.Tensor: + alpha = alpha.detach().float().numpy() # without casting, bf16 causes error + alpha = lora_dim if alpha is None or alpha == 0 else alpha + self.scale = alpha / self.lora_dim + self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える + + # Need more experiences on init method + for v in self.down_list: + torch.nn.init.kaiming_uniform_(v, a=math.sqrt(5)) + for v in self.up_list: + torch.nn.init.zeros_(v) + + def load_state_dict(self, state_dict, strict: bool = True, assign: bool = False): + return + + def custom_state_dict(self): + destination = {} + destination["alpha"] = self.alpha + destination["lora_up.weight"] = nn.Parameter( + torch.concat(list(self.up_list), dim=1) + ) + destination["lora_down.weight"] = nn.Parameter( + torch.concat(list(self.down_list)).reshape( + self.lora_dim, -1, *self.shape[2:] + ) + ) + return destination + + def get_weight(self, rank): + b = math.ceil(rank / self.block_size) + down = torch.concat( + list(i.data for i in self.down_list[:b]) + list(self.down_list[b : (b + 1)]) + ) + up = torch.concat( + list(i.data for i in self.up_list[:b]) + list(self.up_list[b : (b + 1)]), + dim=1, + ) + return down, up, self.alpha / (b + 1) + + def get_random_rank_weight(self): + b = random.randint(0, self.block_count - 1) + return self.get_weight(b * self.block_size) + + def get_diff_weight(self, multiplier=1, shape=None, device=None, rank=None): + if rank is None: + down, up, scale = self.get_random_rank_weight() + else: + down, up, scale = self.get_weight(rank) + w = up @ (down * (scale * multiplier)) + if device is not None: + w = w.to(device) + if shape is not None: + w = w.view(shape) + else: + w = w.view(self.shape) + return w, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None, rank=None): + diff, _ = self.get_diff_weight(multiplier, shape, device, rank) + return diff + self.org_weight, None + + def bypass_forward_diff(self, x, scale=1, rank=None): + if rank is None: + down, up, gamma = self.get_random_rank_weight() + else: + down, up, scale = self.get_weight(rank) + down = down.view(self.lora_dim, -1, *self.shape[2:]) + up = up.view(-1, self.lora_dim, *(1 for _ in self.shape[2:])) + scale = scale * gamma + return self.op(self.op(x, down, **self.kw_dict), up) + + def bypass_forward(self, x, scale=1, rank=None): + return self.org_forward(x) + self.bypass_forward_diff(x, scale, rank) + + def forward(self, x, *args, **kwargs): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.org_forward(x) + if self.bypass_mode: + return self.bypass_forward(x, self.multiplier) + else: + weight = self.get_merged_weight(multiplier=self.multiplier)[0] + bias = ( + None + if self.org_module[0].bias is None + else self.org_module[0].bias.data + ) + return self.op(x, weight, bias, **self.kw_dict) diff --git a/lycoris/modules/full.py b/lycoris/modules/full.py new file mode 100644 index 0000000..d4b0dc2 --- /dev/null +++ b/lycoris/modules/full.py @@ -0,0 +1,214 @@ +from functools import cache + +import torch +import torch.nn as nn + +from .base import LycorisBaseModule +from ..logging import logger + + +@cache +def log_bypass_override(): + return logger.warning( + "Automatic Bypass-Mode detected in algo=full, " + "override with bypass_mode=False since algo=full not support bypass mode. " + "If you are using quantized model which require bypass mode, please don't use algo=full. " + ) + + +class FullModule(LycorisBaseModule): + name = "full" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = ["diff", "diff_b"] + weight_list_det = ["diff"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + rank_dropout_scale=False, + bypass_mode=None, + **kwargs, + ): + org_bypass = bypass_mode + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if bypass_mode and org_bypass is None: + self.bypass_mode = False + log_bypass_override() + + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in Full algo.") + + if self.is_quant: + raise ValueError( + "Quant Linear is not supported and meaningless in Full algo." + ) + + if self.bypass_mode: + raise ValueError("bypass mode is not supported in Full algo.") + + self.weight = nn.Parameter(torch.zeros_like(org_module.weight)) + if org_module.bias is not None: + self.bias = nn.Parameter(torch.zeros_like(org_module.bias)) + else: + self.bias = None + self.is_diff = True + self._org_weight = [self.org_module[0].weight.data.cpu().clone()] + if self.org_module[0].bias is not None: + self.org_bias = [self.org_module[0].bias.data.cpu().clone()] + else: + self.org_bias = None + + @classmethod + def make_module_from_state_dict(cls, lora_name, orig_module, diff, diff_b): + module = cls( + lora_name, + orig_module, + 1, + ) + module.weight.copy_(diff) + if diff_b is not None: + if orig_module.bias is not None: + module.bias.copy_(diff_b) + else: + module.bias = nn.Parameter(diff_b) + module.is_diff = True + return module + + @property + def org_weight(self): + return self._org_weight[0] + + @org_weight.setter + def org_weight(self, value): + self.org_module[0].weight.data.copy_(value) + + def apply_to(self, **kwargs): + self.org_forward = self.org_module[0].forward + self.org_module[0].forward = self.forward + self.weight.data.add_(self.org_module[0].weight.data) + self._org_weight = [self.org_module[0].weight.data.cpu().clone()] + delattr(self.org_module[0], "weight") + if self.org_module[0].bias is not None: + self.bias.data.add_(self.org_module[0].bias.data) + self.org_bias = [self.org_module[0].bias.data.cpu().clone()] + delattr(self.org_module[0], "bias") + else: + self.org_bias = None + self.is_diff = False + + def restore(self): + self.org_module[0].forward = self.org_forward + self.org_module[0].weight = nn.Parameter(self._org_weight[0]) + if self.org_bias is not None: + self.org_module[0].bias = nn.Parameter(self.org_bias[0]) + + def custom_state_dict(self): + sd = {"diff": self.weight.data.cpu() - self._org_weight[0]} + if self.bias is not None: + sd["diff_b"] = self.bias.data.cpu() - self.org_bias[0] + return sd + + def load_weight_prehook( + self, + state_dict, + prefix, + local_metadata, + strict, + missing_keys, + unexpected_keys, + error_msgs, + ): + diff_weight = state_dict.pop(f"{prefix}diff") + state_dict[f"{prefix}weight"] = diff_weight + self.weight.data.to(diff_weight) + if f"{prefix}diff_b" in state_dict: + diff_bias = state_dict.pop(f"{prefix}diff_b") + state_dict[f"{prefix}bias"] = diff_bias + self.bias.data.to(diff_bias) + + def make_weight(self, scale=1, device=None): + drop = ( + torch.rand(self.dim, device=device) > self.rank_dropout + if self.rank_dropout and self.training + else 1 + ) + if drop != 1 or scale != 1 or self.is_diff: + diff_w, diff_b = self.get_diff_weight(scale, device=device) + weight = self.org_weight + diff_w * drop + if self.org_bias is not None: + bias = self.org_bias + diff_b * drop + else: + bias = None + else: + weight = self.weight + bias = self.bias + return weight, bias + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + if self.is_diff: + diff_b = None + if self.bias is not None: + diff_b = self.bias * multiplier + return self.weight * multiplier, diff_b + org_weight = self.org_module[0].weight.to(device, dtype=self.weight.dtype) + diff = self.weight.to(device) - org_weight + diff_b = None + if shape: + diff = diff.view(shape) + if self.bias is not None: + org_bias = self.org_module[0].bias.to(device, dtype=self.bias.dtype) + diff_b = self.bias.to(device) - org_bias + if device is not None: + diff = diff.to(device) + if self.bias is not None: + diff_b = diff_b.to(device) + if multiplier != 1: + diff = diff * multiplier + if diff_b is not None: + diff_b = diff_b * multiplier + return diff * multiplier, diff_b + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + weight, bias = self.make_weight(multiplier, device) + if shape is not None: + weight = weight.view(shape) + if bias is not None: + bias = bias.view(shape[0]) + return weight, bias + + def forward(self, x: torch.Tensor, *args, **kwargs): + if ( + self.module_dropout + and self.training + and torch.rand(1) < self.module_dropout + ): + original = True + else: + original = False + if original: + return self.org_forward(x) + scale = self.multiplier + weight, bias = self.make_weight(scale, x.device) + kw_dict = self.kw_dict | {"weight": weight, "bias": bias} + return self.op(x, **kw_dict) diff --git a/lycoris/modules/glora.py b/lycoris/modules/glora.py new file mode 100644 index 0000000..45a47d8 --- /dev/null +++ b/lycoris/modules/glora.py @@ -0,0 +1,262 @@ +import math + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .base import LycorisBaseModule +from ..functional import tucker_weight_from_conv + + +class GLoRAModule(LycorisBaseModule): + name = "glora" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = [ + "a1.weight", + "a2.weight", + "b1.weight", + "b2.weight", + "bm.weight", + "alpha", + ] + weight_list_det = ["a1.weight"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + rank_dropout_scale=False, + weight_decompose=False, + bypass_mode=None, + rs_lora=False, + **kwargs, + ): + """ + f(x) = WX + WAX + BX, where A and B are low-rank matrices + bypass_forward(x) = W(X+A(X)) + B(X) + bypass_forward_diff(x) = W(A(X)) + B(X) + get_merged_weight() = W + WA + B + get_diff_weight() = WA + B + """ + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in GLoRA algo.") + self.lora_dim = lora_dim + self.tucker = False + self.rs_lora = rs_lora + + if self.module_type.startswith("conv"): + self.isconv = True + # For general LoCon + in_dim = org_module.in_channels + k_size = org_module.kernel_size + stride = org_module.stride + padding = org_module.padding + out_dim = org_module.out_channels + use_tucker = use_tucker and all(i == 1 for i in k_size) + self.down_op = self.op + self.up_op = self.op + + # A + self.a2 = self.module(in_dim, lora_dim, 1, bias=False) + self.a1 = self.module(lora_dim, in_dim, 1, bias=False) + + # B + if use_tucker and any(i != 1 for i in k_size): + self.b2 = self.module(in_dim, lora_dim, 1, bias=False) + self.bm = self.module( + lora_dim, lora_dim, k_size, stride, padding, bias=False + ) + self.tucker = True + else: + self.b2 = self.module( + in_dim, lora_dim, k_size, stride, padding, bias=False + ) + self.b1 = self.module(lora_dim, out_dim, 1, bias=False) + else: + self.isconv = False + self.down_op = F.linear + self.up_op = F.linear + in_dim = org_module.in_features + out_dim = org_module.out_features + self.a2 = nn.Linear(in_dim, lora_dim, bias=False) + self.a1 = nn.Linear(lora_dim, in_dim, bias=False) + self.b2 = nn.Linear(in_dim, lora_dim, bias=False) + self.b1 = nn.Linear(lora_dim, out_dim, bias=False) + + if type(alpha) == torch.Tensor: + alpha = alpha.detach().float().numpy() # without casting, bf16 causes error + alpha = lora_dim if alpha is None or alpha == 0 else alpha + + r_factor = lora_dim + if self.rs_lora: + r_factor = math.sqrt(r_factor) + + self.scale = alpha / r_factor + + self.register_buffer("alpha", torch.tensor(alpha)) # 定数として扱える + + if use_scalar: + self.scalar = nn.Parameter(torch.tensor(0.0)) + else: + self.register_buffer("scalar", torch.tensor(1.0), persistent=False) + + # same as microsoft's + torch.nn.init.kaiming_uniform_(self.a1.weight, a=math.sqrt(5)) + torch.nn.init.kaiming_uniform_(self.b1.weight, a=math.sqrt(5)) + if use_scalar: + torch.nn.init.kaiming_uniform_(self.a2.weight, a=math.sqrt(5)) + torch.nn.init.kaiming_uniform_(self.b2.weight, a=math.sqrt(5)) + else: + torch.nn.init.zeros_(self.a2.weight) + torch.nn.init.zeros_(self.b2.weight) + + @classmethod + def make_module_from_state_dict( + cls, lora_name, orig_module, a1, a2, b1, b2, bm, alpha + ): + module = cls( + lora_name, + orig_module, + 1, + a2.size(0), + float(alpha), + use_tucker=bm is not None, + ) + module.a1.weight.data.copy_(a1) + module.a2.weight.data.copy_(a2) + module.b1.weight.data.copy_(b1) + module.b2.weight.data.copy_(b2) + if bm is not None: + module.bm.weight.data.copy_(bm) + return module + + def custom_state_dict(self): + destination = {} + destination["alpha"] = self.alpha + destination["a1.weight"] = self.a1.weight + destination["a2.weight"] = self.a2.weight * self.scalar + destination["b1.weight"] = self.b1.weight + destination["b2.weight"] = self.b2.weight * self.scalar + if self.tucker: + destination["bm.weight"] = self.bm.weight + return destination + + def load_weight_hook(self, module: nn.Module, incompatible_keys): + missing_keys = incompatible_keys.missing_keys + for key in missing_keys: + if "scalar" in key: + del missing_keys[missing_keys.index(key)] + if isinstance(self.scalar, nn.Parameter): + self.scalar.data.copy_(torch.ones_like(self.scalar)) + elif getattr(self, "scalar", None) is not None: + self.scalar.copy_(torch.ones_like(self.scalar)) + else: + self.register_buffer( + "scalar", torch.ones_like(self.scalar), persistent=False + ) + + def make_weight(self, device=None): + wa1 = self.a1.weight.view(self.a1.weight.size(0), -1) + wa2 = self.a2.weight.view(self.a2.weight.size(0), -1) + orig = self.org_weight + + if self.tucker: + wb = tucker_weight_from_conv(self.b1.weight, self.b2.weight, self.bm.weight) + else: + wb1 = self.b1.weight.view(self.b1.weight.size(0), -1) + wb2 = self.b2.weight.view(self.b2.weight.size(0), -1) + wb = wb1 @ wb2 + wb = wb.view(*orig.shape) + if orig.dim() > 2: + w_wa1 = torch.einsum("o i ..., i j -> o j ...", orig, wa1) + w_wa2 = torch.einsum("o i ..., i j -> o j ...", w_wa1, wa2) + else: + w_wa2 = (orig @ wa1) @ wa2 + return (wb + w_wa2) * self.scale * self.scalar + + def get_diff_weight(self, multiplier=1.0, shape=None, device=None): + weight = self.make_weight(device) * multiplier + if shape is not None: + weight = weight.view(shape) + return weight, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + diff_w, _ = self.get_diff_weight(multiplier, shape, device) + return self.org_weight + diff_w, None + + def _bypass_forward(self, x, scale=1, diff=False): + scale = self.scale * scale + ax_mid = self.a2(x) * scale + bx_mid = self.b2(x) * scale + + if self.rank_dropout and self.training: + drop_a = ( + torch.rand(self.lora_dim, device=ax_mid.device) < self.rank_dropout + ).to(ax_mid.dtype) + drop_b = ( + torch.rand(self.lora_dim, device=bx_mid.device) < self.rank_dropout + ).to(bx_mid.dtype) + if self.rank_dropout_scale: + drop_a /= drop_a.mean() + drop_b /= drop_b.mean() + if (dims := len(x.shape)) == 4: + drop_a = drop_a.view(1, -1, 1, 1) + drop_b = drop_b.view(1, -1, 1, 1) + else: + drop_a = drop_a.view(*[1] * (dims - 1), -1) + drop_b = drop_b.view(*[1] * (dims - 1), -1) + ax_mid = ax_mid * drop_a + bx_mid = bx_mid * drop_b + return ( + self.org_forward( + (0 if diff else x) + self.drop(self.a1(ax_mid)) * self.scale + ) + + self.drop(self.b1(bx_mid)) * self.scale + ) + + def bypass_forward_diff(self, x, scale=1): + return self._bypass_forward(x, scale=scale, diff=True) + + def bypass_forward(self, x, scale=1): + return self._bypass_forward(x, scale=scale, diff=False) + + def forward(self, x, *args, **kwargs): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.org_forward(x) + if self.bypass_mode: + return self.bypass_forward(x, self.multiplier) + else: + weight = ( + self.org_module[0].weight.data.to(self.dtype) + + self.get_diff_weight(multiplier=self.multiplier)[0] + ) + bias = ( + None + if self.org_module[0].bias is None + else self.org_module[0].bias.data + ) + return self.op(x, weight, bias, **self.kw_dict) diff --git a/lycoris/modules/ia3.py b/lycoris/modules/ia3.py new file mode 100644 index 0000000..eeeaa1b --- /dev/null +++ b/lycoris/modules/ia3.py @@ -0,0 +1,142 @@ +import torch +import torch.nn as nn + +from .base import LycorisBaseModule + + +class IA3Module(LycorisBaseModule): + name = "ia3" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = ["weight", "on_input"] + weight_list_det = ["on_input"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + rank_dropout_scale=False, + weight_decompose=False, + bypass_mode=None, + rs_lora=False, + train_on_input=False, + **kwargs, + ): + """if alpha == 0 or None, alpha is rank (no scaling).""" + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in IA^3 algo.") + + if self.module_type.startswith("conv"): + self.isconv = True + in_dim = org_module.in_channels + out_dim = org_module.out_channels + if train_on_input: + train_dim = in_dim + else: + train_dim = out_dim + self.weight = nn.Parameter( + torch.empty(1, train_dim, *(1 for _ in self.shape[2:])) + ) + else: + in_dim = org_module.in_features + out_dim = org_module.out_features + if train_on_input: + train_dim = in_dim + else: + train_dim = out_dim + + self.weight = nn.Parameter(torch.empty(train_dim)) + + # Need more experiences on init method + torch.nn.init.constant_(self.weight, 0) + self.train_input = train_on_input + self.register_buffer("on_input", torch.tensor(int(train_on_input))) + + @classmethod + def make_module_from_state_dict(cls, lora_name, orig_module, weight): + module = cls( + lora_name, + orig_module, + 1, + ) + module.weight.data.copy_(weight) + return module + + def apply_to(self): + self.org_forward = self.org_module[0].forward + self.org_module[0].forward = self.forward + + def make_weight(self, multiplier=1, shape=None, device=None, diff=False): + weight = self.weight * multiplier + int(not diff) + if self.train_input: + diff = self.org_weight * weight + else: + diff = self.org_weight.transpose(0, 1) * weight + diff = diff.transpose(0, 1) + if shape is not None: + diff = diff.view(shape) + if device is not None: + diff = diff.to(device) + return diff + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + diff = self.make_weight( + multiplier=multiplier, shape=shape, device=device, diff=True + ) + return diff, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + diff = self.make_weight(multiplier=multiplier, shape=shape, device=device) + return diff, None + + def _bypass_forward(self, x, scale=1, diff=False): + weight = self.weight * scale + int(not diff) + if self.train_input: + x = x * weight + out = self.org_forward(x) + if not self.train_input: + out = out * weight + return out + + def bypass_forward_diff(self, x, scale=1): + return self._bypass_forward(x, scale, diff=True) + + def bypass_forward(self, x, scale=1): + return self._bypass_forward(x, scale, diff=False) + + def forward(self, x, *args, **kwargs): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.org_forward(x) + if self.bypass_mode: + return self.bypass_forward(x, self.multiplier) + else: + weight = self.get_merged_weight(multiplier=self.multiplier)[0] + bias = ( + None + if self.org_module[0].bias is None + else self.org_module[0].bias.data + ) + return self.op(x, weight, bias, **self.kw_dict) diff --git a/lycoris/modules/locon.py b/lycoris/modules/locon.py new file mode 100644 index 0000000..0338684 --- /dev/null +++ b/lycoris/modules/locon.py @@ -0,0 +1,332 @@ +import math +from functools import cache + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .base import LycorisBaseModule +from ..functional.general import rebuild_tucker +from ..logging import logger + + +@cache +def log_wd(): + return logger.warning( + "Using weight_decompose=True with LoRA (DoRA) will ignore network_dropout." + "Only rank dropout and module dropout will be applied" + ) + + +class LoConModule(LycorisBaseModule): + name = "locon" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = [ + "lora_up.weight", + "lora_down.weight", + "lora_mid.weight", + "alpha", + "dora_scale", + ] + weight_list_det = ["lora_up.weight"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + rank_dropout_scale=False, + weight_decompose=False, + wd_on_out=False, + bypass_mode=None, + rs_lora=False, + **kwargs, + ): + """if alpha == 0 or None, alpha is rank (no scaling).""" + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in LoRA/LoCon algo.") + self.lora_dim = lora_dim + self.tucker = False + self.rs_lora = rs_lora + + if self.module_type.startswith("conv"): + self.isconv = True + # For general LoCon + in_dim = org_module.in_channels + k_size = org_module.kernel_size + stride = org_module.stride + padding = org_module.padding + out_dim = org_module.out_channels + use_tucker = use_tucker and any(i != 1 for i in k_size) + self.down_op = self.op + self.up_op = self.op + if use_tucker and any(i != 1 for i in k_size): + self.lora_down = self.module(in_dim, lora_dim, 1, bias=False) + self.lora_mid = self.module( + lora_dim, lora_dim, k_size, stride, padding, bias=False + ) + self.tucker = True + else: + self.lora_down = self.module( + in_dim, lora_dim, k_size, stride, padding, bias=False + ) + self.lora_up = self.module(lora_dim, out_dim, 1, bias=False) + elif isinstance(org_module, nn.Linear): + self.isconv = False + self.down_op = F.linear + self.up_op = F.linear + in_dim = org_module.in_features + out_dim = org_module.out_features + self.lora_down = nn.Linear(in_dim, lora_dim, bias=False) + self.lora_up = nn.Linear(lora_dim, out_dim, bias=False) + else: + raise NotImplementedError + + self.wd = weight_decompose + self.wd_on_out = wd_on_out + if self.wd: + org_weight = org_module.weight.cpu().clone().float() + self.dora_norm_dims = org_weight.dim() - 1 + if self.wd_on_out: + self.dora_scale = nn.Parameter( + torch.norm( + org_weight.reshape(org_weight.shape[0], -1), + dim=1, + keepdim=True, + ).reshape(org_weight.shape[0], *[1] * self.dora_norm_dims) + ).float() + else: + self.dora_scale = nn.Parameter( + torch.norm( + org_weight.transpose(1, 0).reshape(org_weight.shape[1], -1), + dim=1, + keepdim=True, + ) + .reshape(org_weight.shape[1], *[1] * self.dora_norm_dims) + .transpose(1, 0) + ).float() + + if dropout: + self.dropout = nn.Dropout(dropout) + if self.wd: + log_wd() + else: + self.dropout = nn.Identity() + + if type(alpha) == torch.Tensor: + alpha = alpha.detach().float().numpy() # without casting, bf16 causes error + alpha = lora_dim if alpha is None or alpha == 0 else alpha + + r_factor = lora_dim + if self.rs_lora: + r_factor = math.sqrt(r_factor) + + self.scale = alpha / r_factor + + self.register_buffer("alpha", torch.tensor(alpha * (lora_dim / r_factor))) + + if use_scalar: + self.scalar = nn.Parameter(torch.tensor(0.0)) + else: + self.register_buffer("scalar", torch.tensor(1.0), persistent=False) + # same as microsoft's + torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5)) + if use_scalar: + torch.nn.init.kaiming_uniform_(self.lora_up.weight, a=math.sqrt(5)) + else: + torch.nn.init.constant_(self.lora_up.weight, 0) + if self.tucker: + torch.nn.init.kaiming_uniform_(self.lora_mid.weight, a=math.sqrt(5)) + + @classmethod + def make_module_from_state_dict( + cls, lora_name, orig_module, up, down, mid, alpha, dora_scale + ): + module = cls( + lora_name, + orig_module, + 1, + down.size(0), + float(alpha), + use_tucker=mid is not None, + weight_decompose=dora_scale is not None, + ) + module.lora_up.weight.data.copy_(up) + module.lora_down.weight.data.copy_(down) + if mid is not None: + module.lora_mid.weight.data.copy_(mid) + if dora_scale is not None: + module.dora_scale.copy_(dora_scale) + return module + + def load_weight_hook(self, module: nn.Module, incompatible_keys): + missing_keys = incompatible_keys.missing_keys + for key in missing_keys: + if "scalar" in key: + del missing_keys[missing_keys.index(key)] + if isinstance(self.scalar, nn.Parameter): + self.scalar.data.copy_(torch.ones_like(self.scalar)) + elif getattr(self, "scalar", None) is not None: + self.scalar.copy_(torch.ones_like(self.scalar)) + else: + self.register_buffer( + "scalar", torch.ones_like(self.scalar), persistent=False + ) + + def make_weight(self, device=None): + wa = self.lora_up.weight.to(device) + wb = self.lora_down.weight.to(device) + if self.tucker: + t = self.lora_mid.weight + wa = wa.view(wa.size(0), -1).transpose(0, 1) + wb = wb.view(wb.size(0), -1) + weight = rebuild_tucker(t, wa, wb) + else: + weight = wa.view(wa.size(0), -1) @ wb.view(wb.size(0), -1) + + weight = weight.view(self.shape) + if self.training and self.rank_dropout: + drop = (torch.rand(weight.size(0), device=device) > self.rank_dropout).to( + weight.dtype + ) + drop = drop.view(-1, *[1] * len(weight.shape[1:])) + if self.rank_dropout_scale: + drop /= drop.mean() + weight *= drop + + return weight * self.scalar.to(device) + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + scale = self.scale * multiplier + diff = self.make_weight(device=device) * scale + if shape is not None: + diff = diff.view(shape) + if device is not None: + diff = diff.to(device) + return diff, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + diff = self.get_diff_weight(multiplier=1, shape=shape, device=device)[0] + weight = self.org_weight + if self.wd: + merged = self.apply_weight_decompose(weight + diff, multiplier) + else: + merged = weight + diff * multiplier + return merged, None + + def apply_weight_decompose(self, weight, multiplier=1): + weight = weight.to(self.dora_scale.dtype) + if self.wd_on_out: + weight_norm = ( + weight.reshape(weight.shape[0], -1) + .norm(dim=1) + .reshape(weight.shape[0], *[1] * self.dora_norm_dims) + ) + torch.finfo(weight.dtype).eps + else: + weight_norm = ( + weight.transpose(0, 1) + .reshape(weight.shape[1], -1) + .norm(dim=1, keepdim=True) + .reshape(weight.shape[1], *[1] * self.dora_norm_dims) + .transpose(0, 1) + ) + torch.finfo(weight.dtype).eps + + scale = self.dora_scale.to(weight.device) / weight_norm + if multiplier != 1: + scale = multiplier * (scale - 1) + 1 + + return weight * scale + + def custom_state_dict(self): + destination = {} + if self.wd: + destination["dora_scale"] = self.dora_scale + destination["alpha"] = self.alpha + destination["lora_up.weight"] = self.lora_up.weight * self.scalar + destination["lora_down.weight"] = self.lora_down.weight + if self.tucker: + destination["lora_mid.weight"] = self.lora_mid.weight + return destination + + @torch.no_grad() + def apply_max_norm(self, max_norm, device=None): + orig_norm = self.make_weight(device).norm() * self.scale + norm = torch.clamp(orig_norm, max_norm / 2) + desired = torch.clamp(norm, max=max_norm) + ratio = desired.cpu() / norm.cpu() + + scaled = norm != desired + if scaled: + self.scalar *= ratio + + return scaled, orig_norm * ratio + + def bypass_forward_diff(self, x, scale=1): + if self.tucker: + mid = self.lora_mid(self.lora_down(x)) + else: + mid = self.lora_down(x) + + if self.rank_dropout and self.training: + drop = ( + torch.rand(self.lora_dim, device=mid.device) > self.rank_dropout + ).to(mid.dtype) + if self.rank_dropout_scale: + drop /= drop.mean() + if (dims := len(x.shape)) == 4: + drop = drop.view(1, -1, 1, 1) + else: + drop = drop.view(*[1] * (dims - 1), -1) + mid = mid * drop + + return self.dropout(self.lora_up(mid) * self.scalar * self.scale * scale) + + def bypass_forward(self, x, scale=1): + return self.org_forward(x) + self.bypass_forward_diff(x, scale=scale) + + def forward(self, x): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.org_forward(x) + scale = self.scale + + dtype = self.dtype + if not self.bypass_mode: + diff_weight = self.make_weight(x.device).to(dtype) * scale + weight = self.org_module[0].weight.data.to(dtype) + if self.wd: + weight = self.apply_weight_decompose( + weight + diff_weight, self.multiplier + ) + else: + weight = weight + diff_weight * self.multiplier + bias = ( + None + if self.org_module[0].bias is None + else self.org_module[0].bias.data + ) + return self.op(x, weight, bias, **self.kw_dict) + else: + return self.bypass_forward(x, scale=self.multiplier) diff --git a/lycoris/modules/loha.py b/lycoris/modules/loha.py new file mode 100644 index 0000000..54f8e52 --- /dev/null +++ b/lycoris/modules/loha.py @@ -0,0 +1,329 @@ +import math + +import torch +import torch.nn as nn + +from .base import LycorisBaseModule +from ..functional.loha import diff_weight as loha_diff_weight + + +class LohaModule(LycorisBaseModule): + name = "loha" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = [ + "hada_w1_a", + "hada_w1_b", + "hada_w2_a", + "hada_w2_b", + "hada_t1", + "hada_t2", + "alpha", + "dora_scale", + ] + weight_list_det = ["hada_w1_a"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + rank_dropout_scale=False, + weight_decompose=False, + wd_on_out=False, + bypass_mode=None, + rs_lora=False, + **kwargs, + ): + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in LoHa algo.") + self.lora_name = lora_name + self.lora_dim = lora_dim + self.tucker = False + self.rs_lora = rs_lora + + w_shape = self.shape + if self.module_type.startswith("conv"): + in_dim = org_module.in_channels + k_size = org_module.kernel_size + out_dim = org_module.out_channels + self.shape = (out_dim, in_dim, *k_size) + self.tucker = use_tucker and any(i != 1 for i in k_size) + if self.tucker: + w_shape = (out_dim, in_dim, *k_size) + else: + w_shape = (out_dim, in_dim * torch.tensor(k_size).prod().item()) + + if self.tucker: + self.hada_t1 = nn.Parameter(torch.empty(lora_dim, lora_dim, *w_shape[2:])) + self.hada_w1_a = nn.Parameter( + torch.empty(lora_dim, w_shape[0]) + ) # out_dim, 1-mode + self.hada_w1_b = nn.Parameter( + torch.empty(lora_dim, w_shape[1]) + ) # in_dim , 2-mode + + self.hada_t2 = nn.Parameter(torch.empty(lora_dim, lora_dim, *w_shape[2:])) + self.hada_w2_a = nn.Parameter( + torch.empty(lora_dim, w_shape[0]) + ) # out_dim, 1-mode + self.hada_w2_b = nn.Parameter( + torch.empty(lora_dim, w_shape[1]) + ) # in_dim , 2-mode + else: + self.hada_w1_a = nn.Parameter(torch.empty(w_shape[0], lora_dim)) + self.hada_w1_b = nn.Parameter(torch.empty(lora_dim, w_shape[1])) + + self.hada_w2_a = nn.Parameter(torch.empty(w_shape[0], lora_dim)) + self.hada_w2_b = nn.Parameter(torch.empty(lora_dim, w_shape[1])) + + self.wd = weight_decompose + self.wd_on_out = wd_on_out + if self.wd: + org_weight = org_module.weight.cpu().clone().float() + self.dora_norm_dims = org_weight.dim() - 1 + if self.wd_on_out: + self.dora_scale = nn.Parameter( + torch.norm( + org_weight.reshape(org_weight.shape[0], -1), + dim=1, + keepdim=True, + ).reshape(org_weight.shape[0], *[1] * self.dora_norm_dims) + ).float() + else: + self.dora_scale = nn.Parameter( + torch.norm( + org_weight.transpose(1, 0).reshape(org_weight.shape[1], -1), + dim=1, + keepdim=True, + ) + .reshape(org_weight.shape[1], *[1] * self.dora_norm_dims) + .transpose(1, 0) + ).float() + + if self.dropout: + print("[WARN]LoHa/LoKr haven't implemented normal dropout yet.") + + if type(alpha) == torch.Tensor: + alpha = alpha.detach().float().numpy() # without casting, bf16 causes error + alpha = lora_dim if alpha is None or alpha == 0 else alpha + + r_factor = lora_dim + if self.rs_lora: + r_factor = math.sqrt(r_factor) + + self.scale = alpha / r_factor + + self.register_buffer("alpha", torch.tensor(alpha * (lora_dim / r_factor))) + + if use_scalar: + self.scalar = nn.Parameter(torch.tensor(0.0)) + else: + self.register_buffer("scalar", torch.tensor(1.0), persistent=False) + # Need more experiments on init method + if self.tucker: + torch.nn.init.normal_(self.hada_t1, std=0.1) + torch.nn.init.normal_(self.hada_t2, std=0.1) + torch.nn.init.normal_(self.hada_w1_b, std=1) + torch.nn.init.normal_(self.hada_w1_a, std=0.1) + torch.nn.init.normal_(self.hada_w2_b, std=1) + if use_scalar: + torch.nn.init.normal_(self.hada_w2_a, std=0.1) + else: + torch.nn.init.constant_(self.hada_w2_a, 0) + + @classmethod + def make_module_from_state_dict( + cls, lora_name, orig_module, w1a, w1b, w2a, w2b, t1, t2, alpha, dora_scale + ): + module = cls( + lora_name, + orig_module, + 1, + w1b.size(0), + float(alpha), + use_tucker=t1 is not None, + weight_decompose=dora_scale is not None, + ) + module.hada_w1_a.copy_(w1a) + module.hada_w1_b.copy_(w1b) + module.hada_w2_a.copy_(w2a) + module.hada_w2_b.copy_(w2b) + if t1 is not None: + module.hada_t1.copy_(t1) + module.hada_t2.copy_(t2) + if dora_scale is not None: + module.dora_scale.copy_(dora_scale) + return module + + def load_weight_hook(self, module: nn.Module, incompatible_keys): + missing_keys = incompatible_keys.missing_keys + for key in missing_keys: + if "scalar" in key: + del missing_keys[missing_keys.index(key)] + if isinstance(self.scalar, nn.Parameter): + self.scalar.data.copy_(torch.ones_like(self.scalar)) + elif getattr(self, "scalar", None) is not None: + self.scalar.copy_(torch.ones_like(self.scalar)) + else: + self.register_buffer( + "scalar", torch.ones_like(self.scalar), persistent=False + ) + + def get_weight(self, shape): + scale = torch.tensor( + self.scale, dtype=self.hada_w1_b.dtype, device=self.hada_w1_b.device + ) + if self.tucker: + weight = loha_diff_weight( + self.hada_w1_b, + self.hada_w1_a, + self.hada_w2_b, + self.hada_w2_a, + self.hada_t1, + self.hada_t2, + gamma=scale, + ) + else: + weight = loha_diff_weight( + self.hada_w1_b, + self.hada_w1_a, + self.hada_w2_b, + self.hada_w2_a, + None, + None, + gamma=scale, + ) + if shape is not None: + weight = weight.reshape(shape) + if self.training and self.rank_dropout: + drop = (torch.rand(weight.size(0)) > self.rank_dropout).to(weight.dtype) + drop = drop.view(-1, *[1] * len(weight.shape[1:])).to(weight.device) + if self.rank_dropout_scale: + drop /= drop.mean() + weight *= drop + return weight + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + scale = self.scale * multiplier + diff = self.get_weight(shape) * scale + if device is not None: + diff = diff.to(device) + return diff, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + diff = self.get_diff_weight(multiplier=1, shape=shape, device=device)[0] + weight = self.org_weight + if self.wd: + merged = self.apply_weight_decompose(weight + diff, multiplier) + else: + merged = weight + diff * multiplier + return merged, None + + def apply_weight_decompose(self, weight, multiplier=1): + weight = weight.to(self.dora_scale.dtype) + if self.wd_on_out: + weight_norm = ( + weight.reshape(weight.shape[0], -1) + .norm(dim=1) + .reshape(weight.shape[0], *[1] * self.dora_norm_dims) + ) + torch.finfo(weight.dtype).eps + else: + weight_norm = ( + weight.transpose(0, 1) + .reshape(weight.shape[1], -1) + .norm(dim=1, keepdim=True) + .reshape(weight.shape[1], *[1] * self.dora_norm_dims) + .transpose(0, 1) + ) + torch.finfo(weight.dtype).eps + + scale = self.dora_scale.to(weight.device) / weight_norm + if multiplier != 1: + scale = multiplier * (scale - 1) + 1 + + return weight * scale + + def custom_state_dict(self): + destination = {} + destination["alpha"] = self.alpha + if self.wd: + destination["dora_scale"] = self.dora_scale + destination["hada_w1_a"] = self.hada_w1_a * self.scalar + destination["hada_w1_b"] = self.hada_w1_b + destination["hada_w2_a"] = self.hada_w2_a + destination["hada_w2_b"] = self.hada_w2_b + if self.tucker: + destination["hada_t1"] = self.hada_t1 + destination["hada_t2"] = self.hada_t2 + return destination + + @torch.no_grad() + def apply_max_norm(self, max_norm, device=None): + orig_norm = (self.get_weight(self.shape) * self.scalar).norm() + norm = torch.clamp(orig_norm, max_norm / 2) + desired = torch.clamp(norm, max=max_norm) + ratio = desired.cpu() / norm.cpu() + + scaled = norm != desired + if scaled: + self.scalar *= ratio + + return scaled, orig_norm * ratio + + def bypass_forward_diff(self, x, scale=1): + diff_weight = self.get_weight(self.shape) * self.scalar * scale + return self.drop(self.op(x, diff_weight, **self.kw_dict)) + + def bypass_forward(self, x, scale=1): + return self.org_forward(x) + self.bypass_forward_diff(x, scale=scale) + + def forward(self, x: torch.Tensor, *args, **kwargs): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.op( + x, + self.org_module[0].weight.data, + ( + None + if self.org_module[0].bias is None + else self.org_module[0].bias.data + ), + ) + if self.bypass_mode: + return self.bypass_forward(x, scale=self.multiplier) + else: + diff_weight = self.get_weight(self.shape).to(self.dtype) * self.scalar + weight = self.org_module[0].weight.data.to(self.dtype) + if self.wd: + weight = self.apply_weight_decompose( + weight + diff_weight, self.multiplier + ) + else: + weight = weight + diff_weight * self.multiplier + bias = ( + None + if self.org_module[0].bias is None + else self.org_module[0].bias.data + ) + return self.op(x, weight, bias, **self.kw_dict) diff --git a/lycoris/modules/lokr.py b/lycoris/modules/lokr.py new file mode 100644 index 0000000..12ecf42 --- /dev/null +++ b/lycoris/modules/lokr.py @@ -0,0 +1,609 @@ +import math +from functools import cache + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from .base import LycorisBaseModule +from ..functional import factorization, rebuild_tucker +from ..functional.lokr import make_kron +from ..logging import logger + + +@cache +def logging_force_full_matrix(lora_dim, dim, factor): + logger.warning( + f"lora_dim {lora_dim} is too large for" + f" dim={dim} and {factor=}" + ", using full matrix mode." + ) + + +class LokrModule(LycorisBaseModule): + name = "kron" + support_module = { + "linear", + "conv1d", + "conv2d", + "conv3d", + } + weight_list = [ + "lokr_w1", + "lokr_w1_a", + "lokr_w1_b", + "lokr_w2", + "lokr_w2_a", + "lokr_w2_b", + "lokr_t1", + "lokr_t2", + "alpha", + "dora_scale", + ] + weight_list_det = ["lokr_w1", "lokr_w1_a"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + lora_dim=4, + alpha=1, + dropout=0.0, + rank_dropout=0.0, + module_dropout=0.0, + use_tucker=False, + use_scalar=False, + decompose_both=False, + factor: int = -1, # factorization factor + rank_dropout_scale=False, + weight_decompose=False, + wd_on_out=False, + full_matrix=False, + bypass_mode=None, + rs_lora=False, + unbalanced_factorization=False, + **kwargs, + ): + super().__init__( + lora_name, + org_module, + multiplier, + dropout, + rank_dropout, + module_dropout, + rank_dropout_scale, + bypass_mode, + ) + if self.module_type not in self.support_module: + raise ValueError(f"{self.module_type} is not supported in LoKr algo.") + + factor = int(factor) + self.lora_dim = lora_dim + self.tucker = False + self.use_w1 = False + self.use_w2 = False + self.full_matrix = full_matrix + self.rs_lora = rs_lora + + if self.module_type.startswith("conv"): + in_dim = org_module.in_channels + k_size = org_module.kernel_size + out_dim = org_module.out_channels + self.shape = (out_dim, in_dim, *k_size) + + in_m, in_n = factorization(in_dim, factor) + out_l, out_k = factorization(out_dim, factor) + if unbalanced_factorization: + out_l, out_k = out_k, out_l + shape = ((out_l, out_k), (in_m, in_n), *k_size) # ((a, b), (c, d), *k_size) + self.tucker = use_tucker and any(i != 1 for i in k_size) + if ( + decompose_both + and lora_dim < max(shape[0][0], shape[1][0]) / 2 + and not self.full_matrix + ): + self.lokr_w1_a = nn.Parameter(torch.empty(shape[0][0], lora_dim)) + self.lokr_w1_b = nn.Parameter(torch.empty(lora_dim, shape[1][0])) + else: + self.use_w1 = True + self.lokr_w1 = nn.Parameter( + torch.empty(shape[0][0], shape[1][0]) + ) # a*c, 1-mode + + if lora_dim >= max(shape[0][1], shape[1][1]) / 2 or self.full_matrix: + if not self.full_matrix: + logging_force_full_matrix(lora_dim, max(in_dim, out_dim), factor) + self.use_w2 = True + self.lokr_w2 = nn.Parameter( + torch.empty(shape[0][1], shape[1][1], *k_size) + ) + elif self.tucker: + self.lokr_t2 = nn.Parameter(torch.empty(lora_dim, lora_dim, *shape[2:])) + self.lokr_w2_a = nn.Parameter( + torch.empty(lora_dim, shape[0][1]) + ) # b, 1-mode + self.lokr_w2_b = nn.Parameter( + torch.empty(lora_dim, shape[1][1]) + ) # d, 2-mode + else: # Conv2d not tucker + # bigger part. weight and LoRA. [b, dim] x [dim, d*k1*k2] + self.lokr_w2_a = nn.Parameter(torch.empty(shape[0][1], lora_dim)) + self.lokr_w2_b = nn.Parameter( + torch.empty( + lora_dim, shape[1][1] * torch.tensor(shape[2:]).prod().item() + ) + ) + # w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d*k1*k2)) = (a, b)⊗(c, d*k1*k2) = (ac, bd*k1*k2) + else: # Linear + in_dim = org_module.in_features + out_dim = org_module.out_features + self.shape = (out_dim, in_dim) + + in_m, in_n = factorization(in_dim, factor) + out_l, out_k = factorization(out_dim, factor) + if unbalanced_factorization: + out_l, out_k = out_k, out_l + shape = ( + (out_l, out_k), + (in_m, in_n), + ) # ((a, b), (c, d)), out_dim = a*c, in_dim = b*d + # smaller part. weight scale + if ( + decompose_both + and lora_dim < max(shape[0][0], shape[1][0]) / 2 + and not self.full_matrix + ): + self.lokr_w1_a = nn.Parameter(torch.empty(shape[0][0], lora_dim)) + self.lokr_w1_b = nn.Parameter(torch.empty(lora_dim, shape[1][0])) + else: + self.use_w1 = True + self.lokr_w1 = nn.Parameter( + torch.empty(shape[0][0], shape[1][0]) + ) # a*c, 1-mode + if lora_dim < max(shape[0][1], shape[1][1]) / 2 and not self.full_matrix: + # bigger part. weight and LoRA. [b, dim] x [dim, d] + self.lokr_w2_a = nn.Parameter(torch.empty(shape[0][1], lora_dim)) + self.lokr_w2_b = nn.Parameter(torch.empty(lora_dim, shape[1][1])) + # w1 ⊗ (w2_a x w2_b) = (a, b)⊗((c, dim)x(dim, d)) = (a, b)⊗(c, d) = (ac, bd) + else: + if not self.full_matrix: + logging_force_full_matrix(lora_dim, max(in_dim, out_dim), factor) + self.use_w2 = True + self.lokr_w2 = nn.Parameter(torch.empty(shape[0][1], shape[1][1])) + + self.wd = weight_decompose + self.wd_on_out = wd_on_out + if self.wd: + org_weight = org_module.weight.cpu().clone().float() + self.dora_norm_dims = org_weight.dim() - 1 + if self.wd_on_out: + self.dora_scale = nn.Parameter( + torch.norm( + org_weight.reshape(org_weight.shape[0], -1), + dim=1, + keepdim=True, + ).reshape(org_weight.shape[0], *[1] * self.dora_norm_dims) + ).float() + else: + self.dora_scale = nn.Parameter( + torch.norm( + org_weight.transpose(1, 0).reshape(org_weight.shape[1], -1), + dim=1, + keepdim=True, + ) + .reshape(org_weight.shape[1], *[1] * self.dora_norm_dims) + .transpose(1, 0) + ).float() + + self.dropout = dropout + if dropout: + print("[WARN]LoHa/LoKr haven't implemented normal dropout yet.") + self.rank_dropout = rank_dropout + self.rank_dropout_scale = rank_dropout_scale + self.module_dropout = module_dropout + + if isinstance(alpha, torch.Tensor): + alpha = alpha.detach().float().numpy() # without casting, bf16 causes error + alpha = lora_dim if alpha is None or alpha == 0 else alpha + if self.use_w2 and self.use_w1: + # use scale = 1 + alpha = lora_dim + + r_factor = lora_dim + if self.rs_lora: + r_factor = math.sqrt(r_factor) + + self.scale = alpha / r_factor + + self.register_buffer("alpha", torch.tensor(alpha * (lora_dim / r_factor))) + + if use_scalar: + self.scalar = nn.Parameter(torch.tensor(0.0)) + else: + self.register_buffer("scalar", torch.tensor(1.0), persistent=False) + + if self.use_w2: + if use_scalar: + torch.nn.init.kaiming_uniform_(self.lokr_w2, a=math.sqrt(5)) + else: + torch.nn.init.constant_(self.lokr_w2, 0) + else: + if self.tucker: + torch.nn.init.kaiming_uniform_(self.lokr_t2, a=math.sqrt(5)) + torch.nn.init.kaiming_uniform_(self.lokr_w2_a, a=math.sqrt(5)) + if use_scalar: + torch.nn.init.kaiming_uniform_(self.lokr_w2_b, a=math.sqrt(5)) + else: + torch.nn.init.constant_(self.lokr_w2_b, 0) + + if self.use_w1: + torch.nn.init.kaiming_uniform_(self.lokr_w1, a=math.sqrt(5)) + else: + torch.nn.init.kaiming_uniform_(self.lokr_w1_a, a=math.sqrt(5)) + torch.nn.init.kaiming_uniform_(self.lokr_w1_b, a=math.sqrt(5)) + + @classmethod + def make_module_from_state_dict( + cls, + lora_name, + orig_module, + w1, + w1a, + w1b, + w2, + w2a, + w2b, + _, + t2, + alpha, + dora_scale, + ): + full_matrix = False + if w1a is not None: + lora_dim = w1a.size(1) + elif w2a is not None: + lora_dim = w2a.size(1) + else: + full_matrix = True + lora_dim = 1 + + if w1 is None: + out_dim = w1a.size(0) + in_dim = w1b.size(1) + else: + out_dim, in_dim = w1.shape + + shape_s = [out_dim, in_dim] + + if w2 is None: + out_dim *= w2a.size(0) + in_dim *= w2b.size(1) + else: + out_dim *= w2.size(0) + in_dim *= w2.size(1) + + if ( + shape_s[0] == factorization(out_dim, -1)[0] + and shape_s[1] == factorization(in_dim, -1)[0] + ): + factor = -1 + else: + w1_shape = w1.shape if w1 is not None else (w1a.size(0), w1b.size(1)) + w2_shape = w2.shape if w2 is not None else (w2a.size(0), w2b.size(1)) + shape_group_1 = (w1_shape[0], w2_shape[0]) + shape_group_2 = (w1_shape[1], w2_shape[1]) + w_shape = (w1_shape[0] * w2_shape[0], w1_shape[1] * w2_shape[1]) + factor1 = max(w1.shape) if w1 is not None else max(w1a.size(0), w1b.size(1)) + factor2 = max(w2.shape) if w2 is not None else max(w2a.size(0), w2b.size(1)) + if ( + w_shape[0] % factor1 == 0 + and w_shape[1] % factor1 == 0 + and factor1 in shape_group_1 + and factor1 in shape_group_2 + ): + factor = factor1 + elif ( + w_shape[0] % factor2 == 0 + and w_shape[1] % factor2 == 0 + and factor2 in shape_group_1 + and factor2 in shape_group_2 + ): + factor = factor2 + else: + factor = min(factor1, factor2) + + module = cls( + lora_name, + orig_module, + 1, + lora_dim, + float(alpha), + use_tucker=t2 is not None, + decompose_both=w1 is None and w2 is None, + factor=factor, + weight_decompose=dora_scale is not None, + full_matrix=full_matrix, + ) + if w1 is not None: + module.lokr_w1.copy_(w1) + else: + module.lokr_w1_a.copy_(w1a) + module.lokr_w1_b.copy_(w1b) + if w2 is not None: + module.lokr_w2.copy_(w2) + else: + module.lokr_w2_a.copy_(w2a) + module.lokr_w2_b.copy_(w2b) + if t2 is not None: + module.lokr_t2.copy_(t2) + if dora_scale is not None: + module.dora_scale.copy_(dora_scale) + return module + + def load_weight_hook(self, module: nn.Module, incompatible_keys): + missing_keys = incompatible_keys.missing_keys + for key in missing_keys: + if "scalar" in key: + del missing_keys[missing_keys.index(key)] + if isinstance(self.scalar, nn.Parameter): + self.scalar.data.copy_(torch.ones_like(self.scalar)) + elif getattr(self, "scalar", None) is not None: + self.scalar.copy_(torch.ones_like(self.scalar)) + else: + self.register_buffer( + "scalar", torch.ones_like(self.scalar), persistent=False + ) + + def get_weight(self, shape): + weight = make_kron( + self.lokr_w1 if self.use_w1 else self.lokr_w1_a @ self.lokr_w1_b, + ( + self.lokr_w2 + if self.use_w2 + else ( + rebuild_tucker(self.lokr_t2, self.lokr_w2_a, self.lokr_w2_b) + if self.tucker + else self.lokr_w2_a @ self.lokr_w2_b + ) + ), + self.scale, + ) + dtype = weight.dtype + if shape is not None: + weight = weight.view(shape) + if self.training and self.rank_dropout: + drop = (torch.rand(weight.size(0)) > self.rank_dropout).to(dtype) + drop = drop.view(-1, *[1] * len(weight.shape[1:])) + if self.rank_dropout_scale: + drop /= drop.mean() + weight *= drop + return weight + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + scale = self.scale * multiplier + diff = self.get_weight(shape) * scale + if device is not None: + diff = diff.to(device) + return diff, None + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + diff = self.get_diff_weight(multiplier=1, shape=shape, device=device)[0] + weight = self.org_weight + if self.wd: + merged = self.apply_weight_decompose(weight + diff, multiplier) + else: + merged = weight + diff * multiplier + return merged, None + + def apply_weight_decompose(self, weight, multiplier=1): + weight = weight.to(self.dora_scale.dtype) + if self.wd_on_out: + weight_norm = ( + weight.reshape(weight.shape[0], -1) + .norm(dim=1) + .reshape(weight.shape[0], *[1] * self.dora_norm_dims) + ) + torch.finfo(weight.dtype).eps + else: + weight_norm = ( + weight.transpose(0, 1) + .reshape(weight.shape[1], -1) + .norm(dim=1, keepdim=True) + .reshape(weight.shape[1], *[1] * self.dora_norm_dims) + .transpose(0, 1) + ) + torch.finfo(weight.dtype).eps + + scale = self.dora_scale.to(weight.device) / weight_norm + if multiplier != 1: + scale = multiplier * (scale - 1) + 1 + + return weight * scale + + def custom_state_dict(self): + destination = {} + destination["alpha"] = self.alpha + if self.wd: + destination["dora_scale"] = self.dora_scale + if self.use_w1: + destination["lokr_w1"] = self.lokr_w1 * self.scalar + else: + destination["lokr_w1_a"] = self.lokr_w1_a * self.scalar + destination["lokr_w1_b"] = self.lokr_w1_b + + if self.use_w2: + destination["lokr_w2"] = self.lokr_w2 + else: + destination["lokr_w2_a"] = self.lokr_w2_a + destination["lokr_w2_b"] = self.lokr_w2_b + if self.tucker: + destination["lokr_t2"] = self.lokr_t2 + return destination + + @torch.no_grad() + def apply_max_norm(self, max_norm, device=None): + orig_norm = self.get_weight(self.shape).norm() + norm = torch.clamp(orig_norm, max_norm / 2) + desired = torch.clamp(norm, max=max_norm) + ratio = desired.cpu() / norm.cpu() + + scaled = norm != desired + if scaled: + modules = 4 - self.use_w1 - self.use_w2 + (not self.use_w2 and self.tucker) + if self.use_w1: + self.lokr_w1 *= ratio ** (1 / modules) + else: + self.lokr_w1_a *= ratio ** (1 / modules) + self.lokr_w1_b *= ratio ** (1 / modules) + + if self.use_w2: + self.lokr_w2 *= ratio ** (1 / modules) + else: + if self.tucker: + self.lokr_t2 *= ratio ** (1 / modules) + self.lokr_w2_a *= ratio ** (1 / modules) + self.lokr_w2_b *= ratio ** (1 / modules) + + return scaled, orig_norm * ratio + + def bypass_forward_diff(self, h, scale=1): + is_conv = self.module_type.startswith("conv") + if self.use_w2: + ba = self.lokr_w2 + else: + a = self.lokr_w2_b + b = self.lokr_w2_a + + if self.tucker: + t = self.lokr_t2 + a = a.view(*a.shape, *[1] * (len(t.shape) - 2)) + b = b.view(*b.shape, *[1] * (len(t.shape) - 2)) + elif is_conv: + a = a.view(*a.shape, *self.shape[2:]) + b = b.view(*b.shape, *[1] * (len(self.shape) - 2)) + + if self.use_w1: + c = self.lokr_w1 + else: + c = self.lokr_w1_a @ self.lokr_w1_b + uq = c.size(1) + + if is_conv: + # (b, uq), vq, ... + b, _, *rest = h.shape + h_in_group = h.reshape(b * uq, -1, *rest) + else: + # b, ..., uq, vq + h_in_group = h.reshape(*h.shape[:-1], uq, -1) + + if self.use_w2: + hb = self.op(h_in_group, ba, **self.kw_dict) + else: + if is_conv: + if self.tucker: + ha = self.op(h_in_group, a) + ht = self.op(ha, t, **self.kw_dict) + hb = self.op(ht, b) + else: + ha = self.op(h_in_group, a, **self.kw_dict) + hb = self.op(ha, b) + else: + ha = self.op(h_in_group, a, **self.kw_dict) + hb = self.op(ha, b) + + if is_conv: + # (b, uq), vp, ..., f + # -> b, uq, vp, ..., f + # -> b, f, vp, ..., uq + hb = hb.view(b, -1, *hb.shape[1:]) + h_cross_group = hb.transpose(1, -1) + else: + # b, ..., uq, vq + # -> b, ..., vq, uq + h_cross_group = hb.transpose(-1, -2) + + hc = F.linear(h_cross_group, c) + if is_conv: + # b, f, vp, ..., up + # -> b, up, vp, ... ,f + # -> b, c, ..., f + hc = hc.transpose(1, -1) + h = hc.reshape(b, -1, *hc.shape[3:]) + else: + # b, ..., vp, up + # -> b, ..., up, vp + # -> b, ..., c + hc = hc.transpose(-1, -2) + h = hc.reshape(*hc.shape[:-2], -1) + + return self.drop(h * scale * self.scalar) + + def bypass_forward(self, x, scale=1): + return self.org_forward(x) + self.bypass_forward_diff(x, scale=scale) + + def forward(self, x: torch.Tensor, *args, **kwargs): + if self.module_dropout and self.training: + if torch.rand(1) < self.module_dropout: + return self.org_forward(x) + if self.bypass_mode: + return self.bypass_forward(x, self.multiplier) + else: + diff_weight = self.get_weight(self.shape).to(self.dtype) * self.scalar + weight = self.org_module[0].weight.data.to(self.dtype) + if self.wd: + weight = self.apply_weight_decompose( + weight + diff_weight, self.multiplier + ) + elif self.multiplier == 1: + weight = weight + diff_weight + else: + weight = weight + diff_weight * self.multiplier + bias = ( + None + if self.org_module[0].bias is None + else self.org_module[0].bias.data + ) + return self.op(x, weight, bias, **self.kw_dict) + + +if __name__ == "__main__": + base = nn.Conv2d(128, 128, 3, 1, 1) + net = LokrModule( + "", + base, + multiplier=1, + lora_dim=4, + alpha=1, + weight_decompose=False, + use_tucker=False, + use_scalar=False, + decompose_both=True, + ) + net.apply_to() + sd = net.state_dict() + for key in sd: + if key != "alpha": + sd[key] = torch.randn_like(sd[key]) + net.load_state_dict(sd) + + test_input = torch.randn(1, 128, 16, 16) + test_output = net(test_input) + print(test_output.shape) + + net2 = LokrModule( + "", + base, + multiplier=1, + lora_dim=4, + alpha=1, + weight_decompose=False, + use_tucker=False, + use_scalar=False, + bypass_mode=True, + decompose_both=True, + ) + net2.apply_to() + net2.load_state_dict(sd) + print(net2) + + test_output2 = net(test_input) + print(F.mse_loss(test_output, test_output2)) diff --git a/lycoris/modules/norms.py b/lycoris/modules/norms.py new file mode 100644 index 0000000..6ab627a --- /dev/null +++ b/lycoris/modules/norms.py @@ -0,0 +1,161 @@ +import torch +import torch.nn as nn + +from .base import LycorisBaseModule +from ..logging import warning_once + + +class NormModule(LycorisBaseModule): + name = "norm" + support_module = { + "layernorm", + "groupnorm", + } + weight_list = ["w_norm", "b_norm"] + weight_list_det = ["w_norm"] + + def __init__( + self, + lora_name, + org_module: nn.Module, + multiplier=1.0, + rank_dropout=0.0, + module_dropout=0.0, + rank_dropout_scale=False, + **kwargs, + ): + """if alpha == 0 or None, alpha is rank (no scaling).""" + super().__init__( + lora_name=lora_name, + org_module=org_module, + multiplier=multiplier, + rank_dropout=rank_dropout, + module_dropout=module_dropout, + rank_dropout_scale=rank_dropout_scale, + **kwargs, + ) + if self.module_type == "unknown": + if not hasattr(org_module, "weight") or not hasattr(org_module, "_norm"): + warning_once(f"{type(org_module)} is not supported in Norm algo.") + self.not_supported = True + return + else: + self.dim = org_module.weight.numel() + self.not_supported = False + elif self.module_type not in self.support_module: + warning_once(f"{self.module_type} is not supported in Norm algo.") + self.not_supported = True + return + + self.w_norm = nn.Parameter(torch.zeros(self.dim)) + if hasattr(org_module, "bias"): + self.b_norm = nn.Parameter(torch.zeros(self.dim)) + if hasattr(org_module, "_norm"): + self.org_norm = org_module._norm + else: + self.org_norm = None + + @classmethod + def make_module_from_state_dict(cls, lora_name, orig_module, w_norm, b_norm): + module = cls( + lora_name, + orig_module, + 1, + ) + module.w_norm.copy_(w_norm) + if b_norm is not None: + module.b_norm.copy_(b_norm) + return module + + def make_weight(self, scale=1, device=None): + org_weight = self.org_module[0].weight.to(device, dtype=self.w_norm.dtype) + if hasattr(self.org_module[0], "bias"): + org_bias = self.org_module[0].bias.to(device, dtype=self.b_norm.dtype) + else: + org_bias = None + if self.rank_dropout and self.training: + drop = (torch.rand(self.dim, device=device) < self.rank_dropout).to( + self.w_norm.device + ) + if self.rank_dropout_scale: + drop /= drop.mean() + else: + drop = 1 + drop = ( + torch.rand(self.dim, device=device) < self.rank_dropout + if self.rank_dropout and self.training + else 1 + ) + weight = self.w_norm.to(device) * drop * scale + if org_bias is not None: + bias = self.b_norm.to(device) * drop * scale + return org_weight + weight, org_bias + bias if org_bias is not None else None + + def get_diff_weight(self, multiplier=1, shape=None, device=None): + if self.not_supported: + return 0, 0 + w = self.w_norm * multiplier + if device is not None: + w = w.to(device) + if shape is not None: + w = w.view(shape) + if self.b_norm is not None: + b = self.b_norm * multiplier + if device is not None: + b = b.to(device) + if shape is not None: + b = b.view(shape) + else: + b = None + return w, b + + def get_merged_weight(self, multiplier=1, shape=None, device=None): + if self.not_supported: + return None, None + diff_w, diff_b = self.get_diff_weight(multiplier, shape, device) + org_w = self.org_module[0].weight.to(device, dtype=self.w_norm.dtype) + weight = org_w + diff_w + if diff_b is not None: + org_b = self.org_module[0].bias.to(device, dtype=self.b_norm.dtype) + bias = org_b + diff_b + else: + bias = None + return weight, bias + + def forward(self, x): + if self.not_supported or ( + self.module_dropout + and self.training + and torch.rand(1) < self.module_dropout + ): + return self.org_forward(x) + scale = self.multiplier + + w, b = self.make_weight(scale, x.device) + if self.org_norm is not None: + normed = self.org_norm(x) + scaled = normed * w + if b is not None: + scaled += b + return scaled + + kw_dict = self.kw_dict | {"weight": w, "bias": b} + return self.op(x, **kw_dict) + + +if __name__ == "__main__": + base = nn.LayerNorm(128).cuda() + norm = NormModule("test", base, 1).cuda() + print(norm) + test_input = torch.randn(1, 128).cuda() + test_output = norm(test_input) + torch.sum(test_output).backward() + print(test_output.shape) + + base = nn.GroupNorm(4, 128).cuda() + norm = NormModule("test", base, 1).cuda() + print(norm) + test_input = torch.randn(1, 128, 3, 3).cuda() + test_output = norm(test_input) + torch.sum(test_output).backward() + print(test_output.shape) diff --git a/lycoris/utils/__init__.py b/lycoris/utils/__init__.py new file mode 100644 index 0000000..f1acf0e --- /dev/null +++ b/lycoris/utils/__init__.py @@ -0,0 +1,483 @@ +import re +import hashlib +from io import BytesIO +from typing import Dict, Tuple, Union + +import numpy as np + +import torch +import torch.nn as nn +import torch.nn.functional as F +import torch.linalg as linalg + +import safetensors.torch + +from tqdm import tqdm +from .general import * + + +def load_bytes_in_safetensors(tensors): + bytes = safetensors.torch.save(tensors) + b = BytesIO(bytes) + + b.seek(0) + header = b.read(8) + n = int.from_bytes(header, "little") + + offset = n + 8 + b.seek(offset) + + return b.read() + + +def precalculate_safetensors_hashes(state_dict): + # calculate each tensor one by one to reduce memory usage + hash_sha256 = hashlib.sha256() + for tensor in state_dict.values(): + single_tensor_sd = {"tensor": tensor} + bytes_for_tensor = load_bytes_in_safetensors(single_tensor_sd) + hash_sha256.update(bytes_for_tensor) + + return f"0x{hash_sha256.hexdigest()}" + + +def str_bool(val): + return str(val).lower() != "false" + + +def default(val, d): + return val if val is not None else d + + +def make_sparse(t: torch.Tensor, sparsity=0.95): + abs_t = torch.abs(t) + np_array = abs_t.detach().cpu().numpy() + quan = float(np.quantile(np_array, sparsity)) + sparse_t = t.masked_fill(abs_t < quan, 0) + return sparse_t + + +def extract_conv( + weight: Union[torch.Tensor, nn.Parameter], + mode="fixed", + mode_param=0, + device="cpu", + is_cp=False, +) -> Tuple[nn.Parameter, nn.Parameter]: + weight = weight.to(device) + out_ch, in_ch, kernel_size, _ = weight.shape + + U, S, Vh = linalg.svd(weight.reshape(out_ch, -1)) + + if mode == "full": + return weight, "full" + elif mode == "fixed": + lora_rank = mode_param + elif mode == "threshold": + assert mode_param >= 0 + lora_rank = torch.sum(S > mode_param) + elif mode == "ratio": + assert 1 >= mode_param >= 0 + min_s = torch.max(S) * mode_param + lora_rank = torch.sum(S > min_s) + elif mode == "quantile" or mode == "percentile": + assert 1 >= mode_param >= 0 + s_cum = torch.cumsum(S, dim=0) + min_cum_sum = mode_param * torch.sum(S) + lora_rank = torch.sum(s_cum < min_cum_sum) + else: + raise NotImplementedError( + 'Extract mode should be "fixed", "threshold", "ratio" or "quantile"' + ) + lora_rank = max(1, lora_rank) + lora_rank = min(out_ch, in_ch, lora_rank) + if lora_rank >= out_ch / 2 and not is_cp: + return weight, "full" + + U = U[:, :lora_rank] + S = S[:lora_rank] + U = U @ torch.diag(S).to(device) + Vh = Vh[:lora_rank, :] + + diff = (weight - (U @ Vh).reshape(out_ch, in_ch, kernel_size, kernel_size)).detach() + extract_weight_A = Vh.reshape(lora_rank, in_ch, kernel_size, kernel_size).detach() + extract_weight_B = U.reshape(out_ch, lora_rank, 1, 1).detach() + del U, S, Vh, weight + return (extract_weight_A, extract_weight_B, diff), "low rank" + + +def extract_linear( + weight: Union[torch.Tensor, nn.Parameter], + mode="fixed", + mode_param=0, + device="cpu", +) -> Tuple[nn.Parameter, nn.Parameter]: + weight = weight.to(device) + out_ch, in_ch = weight.shape + + U, S, Vh = linalg.svd(weight) + + if mode == "full": + return weight, "full" + elif mode == "fixed": + lora_rank = mode_param + elif mode == "threshold": + assert mode_param >= 0 + lora_rank = torch.sum(S > mode_param) + elif mode == "ratio": + assert 1 >= mode_param >= 0 + min_s = torch.max(S) * mode_param + lora_rank = torch.sum(S > min_s) + elif mode == "quantile" or mode == "percentile": + assert 1 >= mode_param >= 0 + s_cum = torch.cumsum(S, dim=0) + min_cum_sum = mode_param * torch.sum(S) + lora_rank = torch.sum(s_cum < min_cum_sum) + else: + raise NotImplementedError( + 'Extract mode should be "fixed", "threshold", "ratio" or "quantile"' + ) + lora_rank = max(1, lora_rank) + lora_rank = min(out_ch, in_ch, lora_rank) + if lora_rank >= out_ch / 2: + return weight, "full" + + U = U[:, :lora_rank] + S = S[:lora_rank] + U = U @ torch.diag(S).to(device) + Vh = Vh[:lora_rank, :] + + diff = (weight - U @ Vh).detach() + extract_weight_A = Vh.reshape(lora_rank, in_ch).detach() + extract_weight_B = U.reshape(out_ch, lora_rank).detach() + del U, S, Vh, weight + return (extract_weight_A, extract_weight_B, diff), "low rank" + + +@torch.no_grad() +def extract_diff( + base_tes, + db_tes, + base_unet, + db_unet, + mode="fixed", + linear_mode_param=0, + conv_mode_param=0, + extract_device="cpu", + use_bias=False, + sparsity=0.98, + small_conv=True, +): + UNET_TARGET_REPLACE_MODULE = [ + "Linear", + "Conv2d", + "LayerNorm", + "GroupNorm", + "GroupNorm32", + ] + TEXT_ENCODER_TARGET_REPLACE_MODULE = [ + "Embedding", + "Linear", + "Conv2d", + "LayerNorm", + "GroupNorm", + "GroupNorm32", + ] + LORA_PREFIX_UNET = "lora_unet" + LORA_PREFIX_TEXT_ENCODER = "lora_te" + + def make_state_dict( + prefix, + root_module: torch.nn.Module, + target_module: torch.nn.Module, + target_replace_modules, + ): + loras = {} + temp = {} + + for name, module in root_module.named_modules(): + if module.__class__.__name__ in target_replace_modules: + temp[name] = module + + for name, module in tqdm( + list((n, m) for n, m in target_module.named_modules() if n in temp) + ): + weights = temp[name] + lora_name = prefix + "." + name + lora_name = lora_name.replace(".", "_") + layer = module.__class__.__name__ + + if layer in { + "Linear", + "Conv2d", + "LayerNorm", + "GroupNorm", + "GroupNorm32", + "Embedding", + }: + root_weight = module.weight + if torch.allclose(root_weight, weights.weight): + continue + else: + continue + module = module.to(extract_device) + weights = weights.to(extract_device) + + if mode == "full": + decompose_mode = "full" + elif layer == "Linear": + weight, decompose_mode = extract_linear( + (root_weight - weights.weight), + mode, + linear_mode_param, + device=extract_device, + ) + if decompose_mode == "low rank": + extract_a, extract_b, diff = weight + elif layer == "Conv2d": + is_linear = root_weight.shape[2] == 1 and root_weight.shape[3] == 1 + weight, decompose_mode = extract_conv( + (root_weight - weights.weight), + mode, + linear_mode_param if is_linear else conv_mode_param, + device=extract_device, + ) + if decompose_mode == "low rank": + extract_a, extract_b, diff = weight + if small_conv and not is_linear and decompose_mode == "low rank": + dim = extract_a.size(0) + (extract_c, extract_a, _), _ = extract_conv( + extract_a.transpose(0, 1), + "fixed", + dim, + extract_device, + True, + ) + extract_a = extract_a.transpose(0, 1) + extract_c = extract_c.transpose(0, 1) + loras[f"{lora_name}.lora_mid.weight"] = ( + extract_c.detach().cpu().contiguous().half() + ) + diff = ( + ( + root_weight + - torch.einsum( + "i j k l, j r, p i -> p r k l", + extract_c, + extract_a.flatten(1, -1), + extract_b.flatten(1, -1), + ) + ) + .detach() + .cpu() + .contiguous() + ) + del extract_c + else: + module = module.to("cpu") + weights = weights.to("cpu") + continue + + if decompose_mode == "low rank": + loras[f"{lora_name}.lora_down.weight"] = ( + extract_a.detach().cpu().contiguous().half() + ) + loras[f"{lora_name}.lora_up.weight"] = ( + extract_b.detach().cpu().contiguous().half() + ) + loras[f"{lora_name}.alpha"] = torch.Tensor([extract_a.shape[0]]).half() + if use_bias: + diff = diff.detach().cpu().reshape(extract_b.size(0), -1) + sparse_diff = make_sparse(diff, sparsity).to_sparse().coalesce() + + indices = sparse_diff.indices().to(torch.int16) + values = sparse_diff.values().half() + loras[f"{lora_name}.bias_indices"] = indices + loras[f"{lora_name}.bias_values"] = values + loras[f"{lora_name}.bias_size"] = torch.tensor(diff.shape).to( + torch.int16 + ) + del extract_a, extract_b, diff + elif decompose_mode == "full": + if "Norm" in layer: + w_key = "w_norm" + b_key = "b_norm" + else: + w_key = "diff" + b_key = "diff_b" + weight_diff = module.weight - weights.weight + loras[f"{lora_name}.{w_key}"] = ( + weight_diff.detach().cpu().contiguous().half() + ) + if getattr(weights, "bias", None) is not None: + bias_diff = module.bias - weights.bias + loras[f"{lora_name}.{b_key}"] = ( + bias_diff.detach().cpu().contiguous().half() + ) + else: + raise NotImplementedError + module = module.to("cpu") + weights = weights.to("cpu") + return loras + + all_loras = {} + + all_loras |= make_state_dict( + LORA_PREFIX_UNET, + base_unet, + db_unet, + UNET_TARGET_REPLACE_MODULE, + ) + del base_unet, db_unet + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + for idx, (te1, te2) in enumerate(zip(base_tes, db_tes)): + if len(base_tes) > 1: + prefix = f"{LORA_PREFIX_TEXT_ENCODER}{idx+1}" + else: + prefix = LORA_PREFIX_TEXT_ENCODER + all_loras |= make_state_dict( + prefix, + te1, + te2, + TEXT_ENCODER_TARGET_REPLACE_MODULE, + ) + del te1, te2 + + all_lora_name = set() + for k in all_loras: + lora_name, weight = k.rsplit(".", 1) + all_lora_name.add(lora_name) + print(len(all_lora_name)) + return all_loras + + +re_digits = re.compile(r"\d+") +re_compiled = {} + +suffix_conversion = { + "attentions": {}, + "resnets": { + "conv1": "in_layers_2", + "conv2": "out_layers_3", + "norm1": "in_layers_0", + "norm2": "out_layers_0", + "time_emb_proj": "emb_layers_1", + "conv_shortcut": "skip_connection", + }, +} + + +def convert_diffusers_name_to_compvis(key): + def match(match_list, regex_text): + regex = re_compiled.get(regex_text) + if regex is None: + regex = re.compile(regex_text) + re_compiled[regex_text] = regex + + r = re.match(regex, key) + if not r: + return False + + match_list.clear() + match_list.extend([int(x) if re.match(re_digits, x) else x for x in r.groups()]) + return True + + m = [] + + if match(m, r"lora_unet_conv_in(.*)"): + return f"lora_unet_input_blocks_0_0{m[0]}" + + if match(m, r"lora_unet_conv_out(.*)"): + return f"lora_unet_out_2{m[0]}" + + if match(m, r"lora_unet_time_embedding_linear_(\d+)(.*)"): + return f"lora_unet_time_embed_{m[0] * 2 - 2}{m[1]}" + + if match(m, r"lora_unet_down_blocks_(\d+)_(attentions|resnets)_(\d+)_(.+)"): + suffix = suffix_conversion.get(m[1], {}).get(m[3], m[3]) + return f"lora_unet_input_blocks_{1 + m[0] * 3 + m[2]}_{1 if m[1] == 'attentions' else 0}_{suffix}" + + if match(m, r"lora_unet_mid_block_(attentions|resnets)_(\d+)_(.+)"): + suffix = suffix_conversion.get(m[0], {}).get(m[2], m[2]) + return ( + f"lora_unet_middle_block_{1 if m[0] == 'attentions' else m[1] * 2}_{suffix}" + ) + + if match(m, r"lora_unet_up_blocks_(\d+)_(attentions|resnets)_(\d+)_(.+)"): + suffix = suffix_conversion.get(m[1], {}).get(m[3], m[3]) + return f"lora_unet_output_blocks_{m[0] * 3 + m[2]}_{1 if m[1] == 'attentions' else 0}_{suffix}" + + if match(m, r"lora_unet_down_blocks_(\d+)_downsamplers_0_conv"): + return f"lora_unet_input_blocks_{3 + m[0] * 3}_0_op" + + if match(m, r"lora_unet_up_blocks_(\d+)_upsamplers_0_conv"): + return f"lora_unet_output_blocks_{2 + m[0] * 3}_2_conv" + return key + + +@torch.no_grad() +def merge(tes, unet, lyco_state_dict, scale: float = 1.0, device="cpu"): + from ..modules import make_module, get_module + + LORA_PREFIX_UNET = "lora_unet" + LORA_PREFIX_TEXT_ENCODER = "lora_te" + merged = 0 + + def merge_state_dict( + prefix, + root_module: torch.nn.Module, + lyco_state_dict: Dict[str, torch.Tensor], + ): + nonlocal merged + for child_name, child_module in tqdm( + list(root_module.named_modules()), desc=f"Merging {prefix}" + ): + lora_name = prefix + "." + child_name + lora_name = lora_name.replace(".", "_") + + lyco_type, params = get_module(lyco_state_dict, lora_name) + if lyco_type is None: + continue + module = make_module(lyco_type, params, lora_name, child_module) + if module is None: + continue + module.to(device) + module.merge_to(scale) + key_dict.pop(convert_diffusers_name_to_compvis(lora_name), None) + key_dict.pop(lora_name, None) + merged += 1 + + key_dict = {} + for k, v in tqdm(list(lyco_state_dict.items()), desc="Converting Dtype and Device"): + module, weight_key = k.split(".", 1) + convert_key = convert_diffusers_name_to_compvis(module) + if convert_key != module and len(tes) > 1: + # kohya's format for sdxl is as same as SGM, not diffusers + del lyco_state_dict[k] + key_dict[convert_key] = key_dict.get(convert_key, []) + [k] + k = f"{convert_key}.{weight_key}" + else: + key_dict[module] = key_dict.get(module, []) + [k] + lyco_state_dict[k] = v.float().cpu() + + for idx, te in enumerate(tes): + if len(tes) > 1: + prefix = LORA_PREFIX_TEXT_ENCODER + str(idx + 1) + else: + prefix = LORA_PREFIX_TEXT_ENCODER + merge_state_dict( + prefix, + te, + lyco_state_dict, + ) + torch.cuda.empty_cache() + merge_state_dict( + LORA_PREFIX_UNET, + unet, + lyco_state_dict, + ) + torch.cuda.empty_cache() + print(f"Unused state dict key: {key_dict}") + print(f"{merged} Modules been merged") diff --git a/lycoris/utils/general.py b/lycoris/utils/general.py new file mode 100644 index 0000000..ccfaaa1 --- /dev/null +++ b/lycoris/utils/general.py @@ -0,0 +1,5 @@ +def product(xs: list[int | float]): + res = 1 + for x in xs: + res *= x + return res diff --git a/lycoris/utils/logger.py b/lycoris/utils/logger.py new file mode 100644 index 0000000..e3f51e5 --- /dev/null +++ b/lycoris/utils/logger.py @@ -0,0 +1,35 @@ +import logging +import copy +import sys + + +class ColoredFormatter(logging.Formatter): + COLORS = { + "DEBUG": "\033[0;36m", # CYAN + "INFO": "\033[0;32m", # GREEN + "WARNING": "\033[0;33m", # YELLOW + "ERROR": "\033[0;31m", # RED + "CRITICAL": "\033[0;37;41m", # WHITE ON RED + "RESET": "\033[0m", # RESET COLOR + } + + def format(self, record): + colored_record = copy.copy(record) + levelname = colored_record.levelname + seq = self.COLORS.get(levelname, self.COLORS["RESET"]) + colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}" + return super().format(colored_record) + + +# Create a new logger +logger = logging.getLogger("LyCORIS") +logger.propagate = False + +# Add handler if we don't have one. +if not logger.handlers: + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter(ColoredFormatter("[%(name)s]-%(levelname)s: %(message)s")) + logger.addHandler(handler) + +logger.setLevel(logging.DEBUG) +logger.debug("Logger initialized.") diff --git a/lycoris/utils/preset.py b/lycoris/utils/preset.py new file mode 100644 index 0000000..f6ec844 --- /dev/null +++ b/lycoris/utils/preset.py @@ -0,0 +1,9 @@ +import toml + + +def read_preset(preset): + try: + return toml.load(preset) + except Exception as e: + print("Error: cannot read preset file. ", e) + return None diff --git a/lycoris/utils/quant.py b/lycoris/utils/quant.py new file mode 100644 index 0000000..aa4660e --- /dev/null +++ b/lycoris/utils/quant.py @@ -0,0 +1,88 @@ +from functools import cache + +SUPPORT_QUANT = False +try: + from bitsandbytes.nn import LinearNF4, Linear8bitLt, LinearFP4 + + SUPPORT_QUANT = True +except Exception: + import torch.nn as nn + + class LinearNF4(nn.Linear): + pass + + class Linear8bitLt(nn.Linear): + pass + + class LinearFP4(nn.Linear): + pass + + +try: + from quanto.nn import QLinear, QConv2d, QLayerNorm + + SUPPORT_QUANT = True +except Exception: + import torch.nn as nn + + class QLinear(nn.Linear): + pass + + class QConv2d(nn.Conv2d): + pass + + class QLayerNorm(nn.LayerNorm): + pass + + +try: + from optimum.quanto.nn import ( + QLinear as QLinearOpt, + QConv2d as QConv2dOpt, + QLayerNorm as QLayerNormOpt, + ) + + SUPPORT_QUANT = True +except Exception: + import torch.nn as nn + + class QLinearOpt(nn.Linear): + pass + + class QConv2dOpt(nn.Conv2d): + pass + + class QLayerNormOpt(nn.LayerNorm): + pass + + +from ..logging import logger + + +QuantLinears = ( + Linear8bitLt, + LinearFP4, + LinearNF4, + QLinear, + QConv2d, + QLayerNorm, + QLinearOpt, + QConv2dOpt, + QLayerNormOpt, +) + + +@cache +def log_bypass(): + return logger.warning( + "Using bnb/quanto/optimum-quanto with LyCORIS will enable force-bypass mode." + ) + + +@cache +def log_suspect(): + return logger.warning( + "Non-native Linear detected but bypass_mode is not set. " + "Automatically using force-bypass mode to avoid possible issues. " + "Please set bypass_mode=False explicitly if there are no quantized layers." + ) diff --git a/lycoris/utils/xformers_utils.py b/lycoris/utils/xformers_utils.py new file mode 100644 index 0000000..31f737c --- /dev/null +++ b/lycoris/utils/xformers_utils.py @@ -0,0 +1,13 @@ +memory_efficient_attention = None +try: + import xformers +except Exception: + pass + +try: + from xformers.ops import memory_efficient_attention + + XFORMERS_AVAIL = True +except Exception: + memory_efficient_attention = None + XFORMERS_AVAIL = False diff --git a/lycoris/wrapper.py b/lycoris/wrapper.py new file mode 100644 index 0000000..3d53a11 --- /dev/null +++ b/lycoris/wrapper.py @@ -0,0 +1,640 @@ +# General LyCORIS wrapper based on kohya-ss/sd-scripts' style +import os +import fnmatch +import re +import logging + +from typing import Any, List + +import torch +import torch.nn as nn + +from .modules.locon import LoConModule +from .modules.loha import LohaModule +from .modules.lokr import LokrModule +from .modules.dylora import DyLoraModule +from .modules.glora import GLoRAModule +from .modules.norms import NormModule +from .modules.full import FullModule +from .modules.diag_oft import DiagOFTModule +from .modules.boft import ButterflyOFTModule +from .modules import get_module, make_module + +from .config import PRESET +from .utils.preset import read_preset +from .utils import str_bool +from .logging import logger + + +VALID_PRESET_KEYS = [ + "enable_conv", + "target_module", + "target_name", + "module_algo_map", + "name_algo_map", + "lora_prefix", + "use_fnmatch", + "unet_target_module", + "unet_target_name", + "text_encoder_target_module", + "text_encoder_target_name", + "exclude_name", +] + + +network_module_dict = { + "lora": LoConModule, + "locon": LoConModule, + "loha": LohaModule, + "lokr": LokrModule, + "dylora": DyLoraModule, + "glora": GLoRAModule, + "full": FullModule, + "diag-oft": DiagOFTModule, + "boft": ButterflyOFTModule, +} +deprecated_arg_dict = { + "disable_conv_cp": "use_tucker", + "use_cp": "use_tucker", + "use_conv_cp": "use_tucker", + "constrain": "constraint", +} + + +def create_lycoris(module, multiplier=1.0, linear_dim=4, linear_alpha=1, **kwargs): + for key, value in list(kwargs.items()): + if key in deprecated_arg_dict: + logger.warning( + f"{key} is deprecated. Please use {deprecated_arg_dict[key]} instead.", + stacklevel=2, + ) + kwargs[deprecated_arg_dict[key]] = value + if linear_dim is None: + linear_dim = 4 # default + conv_dim = int(kwargs.get("conv_dim", linear_dim) or linear_dim) + conv_alpha = float(kwargs.get("conv_alpha", linear_alpha) or linear_alpha) + dropout = float(kwargs.get("dropout", 0.0) or 0.0) + rank_dropout = float(kwargs.get("rank_dropout", 0.0) or 0.0) + module_dropout = float(kwargs.get("module_dropout", 0.0) or 0.0) + algo = (kwargs.get("algo", "lora") or "lora").lower() + use_tucker = str_bool( + not kwargs.get("disable_conv_cp", True) + or kwargs.get("use_conv_cp", False) + or kwargs.get("use_cp", False) + or kwargs.get("use_tucker", False) + ) + use_scalar = str_bool(kwargs.get("use_scalar", False)) + block_size = int(kwargs.get("block_size", 4) or 4) + train_norm = str_bool(kwargs.get("train_norm", False)) + constraint = float(kwargs.get("constraint", 0) or 0) + rescaled = str_bool(kwargs.get("rescaled", False)) + weight_decompose = str_bool(kwargs.get("dora_wd", False)) + wd_on_output = str_bool(kwargs.get("wd_on_output", False)) + full_matrix = str_bool(kwargs.get("full_matrix", False)) + bypass_mode = str_bool(kwargs.get("bypass_mode", None)) + unbalanced_factorization = str_bool(kwargs.get("unbalanced_factorization", False)) + + if unbalanced_factorization: + logger.info("Unbalanced factorization for LoKr is enabled") + + if bypass_mode: + logger.info("Bypass mode is enabled") + + if weight_decompose: + logger.info("Weight decomposition is enabled") + + if full_matrix: + logger.info("Full matrix mode for LoKr is enabled") + + preset = kwargs.get("preset", "full") + if preset not in PRESET: + preset = read_preset(preset) + else: + preset = PRESET[preset] + assert preset is not None + LycorisNetwork.apply_preset(preset) + + logger.info(f"Using rank adaptation algo: {algo}") + + network = LycorisNetwork( + module, + multiplier=multiplier, + lora_dim=linear_dim, + conv_lora_dim=conv_dim, + alpha=linear_alpha, + conv_alpha=conv_alpha, + dropout=dropout, + rank_dropout=rank_dropout, + module_dropout=module_dropout, + use_tucker=use_tucker, + use_scalar=use_scalar, + network_module=algo, + train_norm=train_norm, + decompose_both=kwargs.get("decompose_both", False), + factor=kwargs.get("factor", -1), + block_size=block_size, + constraint=constraint, + rescaled=rescaled, + weight_decompose=weight_decompose, + wd_on_out=wd_on_output, + full_matrix=full_matrix, + bypass_mode=bypass_mode, + unbalanced_factorization=unbalanced_factorization, + ) + + return network + + +def create_lycoris_from_weights(multiplier, file, module, weights_sd=None, **kwargs): + if weights_sd is None: + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import load_file + + weights_sd = load_file(file) + else: + weights_sd = torch.load(file, map_location="cpu") + + # get dim/alpha mapping + loras = {} + for key in weights_sd: + if "." not in key: + continue + + lora_name = key.split(".")[0] + loras[lora_name] = None + + for name, modules in module.named_modules(): + lora_name = f"{LycorisNetwork.LORA_PREFIX}_{name}".replace(".", "_") + if lora_name in loras: + loras[lora_name] = modules + + original_level = logger.level + logger.setLevel(logging.ERROR) + network = LycorisNetwork(module, init_only=True) + network.multiplier = multiplier + network.loras = [] + logger.setLevel(original_level) + + logger.info("Loading Modules from state dict...") + for lora_name, orig_modules in loras.items(): + if orig_modules is None: + continue + lyco_type, params = get_module(weights_sd, lora_name) + module = make_module(lyco_type, params, lora_name, orig_modules) + if module is not None: + network.loras.append(module) + network.algo_table[module.__class__.__name__] = ( + network.algo_table.get(module.__class__.__name__, 0) + 1 + ) + logger.info(f"{len(network.loras)} Modules Loaded") + + for lora in network.loras: + lora.multiplier = multiplier + + return network, weights_sd + + +class LycorisNetwork(torch.nn.Module): + ENABLE_CONV = True + TARGET_REPLACE_MODULE = [ + "Linear", + "Conv1d", + "Conv2d", + "Conv3d", + "GroupNorm", + "LayerNorm", + ] + TARGET_REPLACE_NAME = [] + LORA_PREFIX = "lycoris" + MODULE_ALGO_MAP = {} + NAME_ALGO_MAP = {} + USE_FNMATCH = False + TARGET_EXCLUDE_NAME = [] + + @classmethod + def apply_preset(cls, preset): + for preset_key in preset.keys(): + if preset_key not in VALID_PRESET_KEYS: + raise KeyError( + f'Unknown preset key "{preset_key}". Valid keys: {VALID_PRESET_KEYS}' + ) + + if "enable_conv" in preset: + cls.ENABLE_CONV = preset["enable_conv"] + if "target_module" in preset: + cls.TARGET_REPLACE_MODULE = preset["target_module"] + if "target_name" in preset: + cls.TARGET_REPLACE_NAME = preset["target_name"] + if "module_algo_map" in preset: + cls.MODULE_ALGO_MAP = preset["module_algo_map"] + if "name_algo_map" in preset: + cls.NAME_ALGO_MAP = preset["name_algo_map"] + if "lora_prefix" in preset: + cls.LORA_PREFIX = preset["lora_prefix"] + if "use_fnmatch" in preset: + cls.USE_FNMATCH = preset["use_fnmatch"] + if "exclude_name" in preset: + cls.TARGET_EXCLUDE_NAME = preset["exclude_name"] + return cls + + def __init__( + self, + module: nn.Module, + multiplier=1.0, + lora_dim=4, + conv_lora_dim=4, + alpha=1, + conv_alpha=1, + use_tucker=False, + dropout=0, + rank_dropout=0, + module_dropout=0, + network_module: str = "locon", + norm_modules=NormModule, + train_norm=False, + init_only=False, + **kwargs, + ) -> None: + super().__init__() + root_kwargs = kwargs + self.weights_sd = None + if init_only: + self.multiplier = 1 + self.lora_dim = 0 + self.alpha = 1 + self.conv_lora_dim = 0 + self.conv_alpha = 1 + self.dropout = 0 + self.rank_dropout = 0 + self.module_dropout = 0 + self.use_tucker = False + self.loras = [] + self.algo_table = {} + return + self.multiplier = multiplier + self.lora_dim = lora_dim + + if not self.ENABLE_CONV: + conv_lora_dim = 0 + + self.conv_lora_dim = int(conv_lora_dim) + if self.conv_lora_dim and self.conv_lora_dim != self.lora_dim: + logger.info("Apply different lora dim for conv layer") + logger.info(f"Conv Dim: {conv_lora_dim}, Linear Dim: {lora_dim}") + elif self.conv_lora_dim == 0: + logger.info("Disable conv layer") + + self.alpha = alpha + self.conv_alpha = float(conv_alpha) + if self.conv_lora_dim and self.alpha != self.conv_alpha: + logger.info("Apply different alpha value for conv layer") + logger.info(f"Conv alpha: {conv_alpha}, Linear alpha: {alpha}") + + if 1 >= dropout >= 0: + logger.info(f"Use Dropout value: {dropout}") + self.dropout = dropout + self.rank_dropout = rank_dropout + self.module_dropout = module_dropout + + self.use_tucker = use_tucker + + def create_single_module( + lora_name: str, + module: torch.nn.Module, + algo_name, + dim=None, + alpha=None, + use_tucker=self.use_tucker, + **kwargs, + ): + for k, v in root_kwargs.items(): + if k in kwargs: + continue + kwargs[k] = v + + if train_norm and "Norm" in module.__class__.__name__: + return norm_modules( + lora_name, + module, + self.multiplier, + self.rank_dropout, + self.module_dropout, + **kwargs, + ) + lora = None + if isinstance(module, torch.nn.Linear) and lora_dim > 0: + dim = dim or lora_dim + alpha = alpha or self.alpha + elif isinstance( + module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d) + ): + k_size, *_ = module.kernel_size + if k_size == 1 and lora_dim > 0: + dim = dim or lora_dim + alpha = alpha or self.alpha + elif conv_lora_dim > 0 or dim: + dim = dim or conv_lora_dim + alpha = alpha or self.conv_alpha + else: + return None + else: + return None + lora = network_module_dict[algo_name]( + lora_name, + module, + self.multiplier, + dim, + alpha, + self.dropout, + self.rank_dropout, + self.module_dropout, + use_tucker, + **kwargs, + ) + return lora + + def create_modules_( + prefix: str, + root_module: torch.nn.Module, + algo, + current_lora_map: dict[str, Any], + configs={}, + ): + assert current_lora_map is not None, "No mapping supplied" + loras = current_lora_map + lora_names = [] + for name, module in root_module.named_modules(): + module_name = module.__class__.__name__ + if module_name in self.MODULE_ALGO_MAP and module is not root_module: + next_config = self.MODULE_ALGO_MAP[module_name] + next_algo = next_config.get("algo", algo) + new_loras, new_lora_names, new_lora_map = create_modules_( + f"{prefix}_{name}" if name else prefix, + module, + next_algo, + loras, + configs=next_config, + ) + loras = {**loras, **new_lora_map} + for lora_name, lora in zip(new_lora_names, new_loras): + if lora_name not in loras and lora_name not in current_lora_map: + loras[lora_name] = lora + if lora_name not in lora_names: + lora_names.append(lora_name) + continue + + if name: + lora_name = prefix + "." + name + else: + lora_name = prefix + + if f"{self.LORA_PREFIX}_." in lora_name: + lora_name = lora_name.replace( + f"{self.LORA_PREFIX}_.", + f"{self.LORA_PREFIX}.", + ) + + lora_name = lora_name.replace(".", "_") + if lora_name in loras: + continue + + lora = create_single_module(lora_name, module, algo, **configs) + if lora is not None: + loras[lora_name] = lora + lora_names.append(lora_name) + return [loras[lora_name] for lora_name in lora_names], lora_names, loras + + # create module instances + def create_modules( + prefix, + root_module: torch.nn.Module, + target_replace_modules, + target_replace_names=[], + target_exclude_names=[], + ) -> List: + logger.info("Create LyCORIS Module") + loras = [] + lora_map = {} + next_config = {} + for name, module in root_module.named_modules(): + if name in target_exclude_names or any( + self.match_fn(t, name) for t in target_exclude_names + ): + continue + + module_name = module.__class__.__name__ + if module_name in target_replace_modules and not any( + self.match_fn(t, name) for t in target_replace_names + ): + if module_name in self.MODULE_ALGO_MAP: + next_config = self.MODULE_ALGO_MAP[module_name] + algo = next_config.get("algo", network_module) + else: + algo = network_module + + lora_lst, _, _lora_map = create_modules_( + f"{prefix}_{name}", + module, + algo, + lora_map, + configs=next_config, + ) + lora_map = {**lora_map, **_lora_map} + loras.extend(lora_lst) + next_config = {} + elif name in target_replace_names or any( + self.match_fn(t, name) for t in target_replace_names + ): + conf_from_name = self.find_conf_for_name(name) + if conf_from_name is not None: + next_config = conf_from_name + algo = next_config.get("algo", network_module) + elif module_name in self.MODULE_ALGO_MAP: + next_config = self.MODULE_ALGO_MAP[module_name] + algo = next_config.get("algo", network_module) + else: + algo = network_module + lora_name = prefix + "." + name + lora_name = lora_name.replace(".", "_") + + if lora_name in lora_map: + continue + + lora = create_single_module(lora_name, module, algo, **next_config) + next_config = {} + if lora is not None: + lora_map[lora.lora_name] = lora + loras.append(lora) + return loras + + self.loras = create_modules( + LycorisNetwork.LORA_PREFIX, + module, + list( + set( + [ + *LycorisNetwork.TARGET_REPLACE_MODULE, + *LycorisNetwork.MODULE_ALGO_MAP.keys(), + ] + ) + ), + list( + set( + [ + *LycorisNetwork.TARGET_REPLACE_NAME, + *LycorisNetwork.NAME_ALGO_MAP.keys(), + ] + ) + ), + target_exclude_names=LycorisNetwork.TARGET_EXCLUDE_NAME, + ) + logger.info(f"create LyCORIS: {len(self.loras)} modules.") + + algo_table = {} + for lora in self.loras: + algo_table[lora.__class__.__name__] = ( + algo_table.get(lora.__class__.__name__, 0) + 1 + ) + logger.info(f"module type table: {algo_table}") + + # Assertion to ensure we have not accidentally wrapped some layers + # multiple times. + names = set() + for lora in self.loras: + assert ( + lora.lora_name not in names + ), f"duplicated lora name: {lora.lora_name}" + names.add(lora.lora_name) + + def match_fn(self, pattern: str, name: str) -> bool: + if self.USE_FNMATCH: + return fnmatch.fnmatch(name, pattern) + return bool(re.match(pattern, name)) + + def find_conf_for_name( + self, + name: str, + ) -> dict[str, Any]: + if name in self.NAME_ALGO_MAP.keys(): + return self.NAME_ALGO_MAP[name] + + for key, value in self.NAME_ALGO_MAP.items(): + if self.match_fn(key, name): + return value + + return None + + def set_multiplier(self, multiplier): + self.multiplier = multiplier + for lora in self.loras: + lora.multiplier = self.multiplier + + def load_weights(self, file): + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import load_file, safe_open + + self.weights_sd = load_file(file) + else: + self.weights_sd = torch.load(file, map_location="cpu") + missing, unexpected = self.load_state_dict(self.weights_sd, strict=False) + state = {} + if missing: + state["missing keys"] = missing + if unexpected: + state["unexpected keys"] = unexpected + return state + + def apply_to(self): + """ + Register to modules to the subclass so that torch sees them. + """ + for lora in self.loras: + lora.apply_to() + self.add_module(lora.lora_name, lora) + + if self.weights_sd: + # if some weights are not in state dict, it is ok because initial LoRA does nothing (lora_up is initialized by zeros) + info = self.load_state_dict(self.weights_sd, False) + logger.info(f"weights are loaded: {info}") + + def is_mergeable(self): + return True + + def restore(self): + for lora in self.loras: + lora.restore() + + def merge_to(self, weight=1.0): + for lora in self.loras: + lora.merge_to(weight) + + def apply_max_norm_regularization(self, max_norm_value, device): + key_scaled = 0 + norms = [] + for module in self.loras: + scaled, norm = module.apply_max_norm(max_norm_value, device) + if scaled is None: + continue + norms.append(norm) + key_scaled += scaled + + if key_scaled == 0: + return key_scaled, 0, 0 + + return key_scaled, sum(norms) / len(norms), max(norms) + + def enable_gradient_checkpointing(self): + # not supported + def make_ckpt(module): + if isinstance(module, torch.nn.Module): + module.grad_ckpt = True + + self.apply(make_ckpt) + pass + + def prepare_optimizer_params(self, lr): + def enumerate_params(loras): + params = [] + for lora in loras: + params.extend(lora.parameters()) + return params + + self.requires_grad_(True) + all_params = [] + + param_data = {"params": enumerate_params(self.loras)} + if lr is not None: + param_data["lr"] = lr + all_params.append(param_data) + return all_params + + def prepare_grad_etc(self, *args): + self.requires_grad_(True) + + def on_epoch_start(self, *args): + self.train() + + def get_trainable_params(self, *args): + return self.parameters() + + def save_weights(self, file, dtype, metadata): + if metadata is not None and len(metadata) == 0: + metadata = None + + state_dict = self.state_dict() + + if dtype is not None: + for key in list(state_dict.keys()): + v = state_dict[key] + v = v.detach().clone().to("cpu").to(dtype) + state_dict[key] = v + + if os.path.splitext(file)[1] == ".safetensors": + from safetensors.torch import save_file + + # Precalculate model hashes to save time on indexing + if metadata is None: + metadata = {} + save_file(state_dict, file, metadata) + else: + torch.save(state_dict, file) diff --git a/nodes.py b/nodes.py index 9f5258d..c43b8cb 100644 --- a/nodes.py +++ b/nodes.py @@ -339,6 +339,50 @@ class OptimizerConfigProdigy: kwargs["min_snr_gamma"] = min_snr_gamma if min_snr_gamma != 0.0 else None return (kwargs,) + +class TrainNetworkConfig: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "network_type": (["lora", "LyCORIS/LoKr", "LyCORIS/Locon", "LyCORIS/LoHa"], {"default": "lora", "tooltip": "network type"}), + "lycoris_preset": (["full", "full-lin", "attn-mlp", "attn-only"], {"default": "attn-mlp"}), + "factor": ("INT",{"default": -1, "min": -1, "max": 16, "step": 1, "tooltip": "LoKr factor"}), + "extra_network_args": ("STRING",{"multiline": True, "default": "", "tooltip": "additional network args"}), + }, + } + + RETURN_TYPES = ("NETWORK_CONFIG",) + RETURN_NAMES = ("network_config",) + FUNCTION = "create_config" + CATEGORY = "FluxTrainer" + + def create_config(self, network_type, extra_network_args, lycoris_preset, factor): + + extra_args = [arg.strip() for arg in extra_network_args.strip().split('|') if arg.strip()] + + if network_type == "lora": + network_module = ".networks.lora" + elif network_type == "LyCORIS/LoKr": + network_module = ".lycoris.kohya" + algo = "lokr" + elif network_type == "LyCORIS/Locon": + network_module = ".lycoris.kohya" + algo = "locon" + elif network_type == "LyCORIS/LoHa": + network_module = ".lycoris.kohya" + algo = "loha" + + network_args = [ + f"algo={algo}", + f"factor={factor}", + f"preset={lycoris_preset}" + ] + network_config = { + "network_module": network_module, + "network_args": network_args + extra_args + } + + return (network_config,) class OptimizerConfigProdigyPlusScheduleFree: @classmethod @@ -390,7 +434,7 @@ class InitFluxLoRATraining: "optimizer_settings": ("ARGS",), "output_name": ("STRING", {"default": "flux_lora", "multiline": False}), "output_dir": ("STRING", {"default": "flux_trainer_output", "multiline": False, "tooltip": "path to dataset, root is the 'ComfyUI' folder, with windows portable 'ComfyUI_windows_portable'"}), - "network_dim": ("INT", {"default": 4, "min": 1, "max": 2048, "step": 1, "tooltip": "network dim"}), + "network_dim": ("INT", {"default": 4, "min": 1, "max": 100000, "step": 1, "tooltip": "network dim"}), "network_alpha": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2048.0, "step": 0.01, "tooltip": "network alpha"}), "learning_rate": ("FLOAT", {"default": 4e-4, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "learning rate"}), "max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 100000, "step": 1, "tooltip": "max number of training steps"}), @@ -423,6 +467,7 @@ class InitFluxLoRATraining: "block_args": ("ARGS", {"default": "", "tooltip": "limit the blocks used in the LoRA"}), "gradient_checkpointing": (["enabled", "enabled_with_cpu_offloading", "disabled"], {"default": "enabled", "tooltip": "use gradient checkpointing"}), "loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}), + "network_config": ("NETWORK_CONFIG", {"tooltip": "additional network config"}), }, "hidden": { "prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO" @@ -436,7 +481,7 @@ class InitFluxLoRATraining: def init_training(self, flux_models, dataset, optimizer_settings, sample_prompts, output_name, attention_mode, gradient_dtype, save_dtype, additional_args=None, resume_args=None, train_text_encoder='disabled', - block_args=None, gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, clip_l_lr=0, T5_lr=0, loss_args=None, **kwargs): + block_args=None, gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, clip_l_lr=0, T5_lr=0, loss_args=None, network_config=None, **kwargs): mm.soft_empty_cache() output_dir = os.path.abspath(kwargs.get("output_dir")) @@ -500,7 +545,7 @@ class InitFluxLoRATraining: "persistent_data_loader_workers": False, "max_data_loader_n_workers": 0, "seed": 42, - "network_module": ".networks.lora_flux", + "network_module": ".networks.lora_flux" if network_config is None else network_config["network_module"], "dataset_config": dataset_toml, "output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}", "loss_type": "l2", @@ -509,6 +554,7 @@ class InitFluxLoRATraining: "network_train_unet_only": True if train_text_encoder == 'disabled' else False, "fp8_base_unet": True if "fp8" in train_text_encoder else False, "disable_mmap_load_safetensors": False, + "network_args": None if network_config is None else network_config["network_args"], } attention_settings = { "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True}, @@ -1722,6 +1768,7 @@ NODE_CLASS_MAPPINGS = { "FluxTrainAndValidateLoop": FluxTrainAndValidateLoop, "OptimizerConfigProdigyPlusScheduleFree": OptimizerConfigProdigyPlusScheduleFree, "FluxTrainerLossConfig": FluxTrainerLossConfig, + "TrainNetworkConfig": TrainNetworkConfig, } NODE_DISPLAY_NAME_MAPPINGS = { "InitFluxLoRATraining": "Init Flux LoRA Training", @@ -1748,4 +1795,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "FluxTrainAndValidateLoop": "Flux Train And Validate Loop", "OptimizerConfigProdigyPlusScheduleFree": "Optimizer Config ProdigyPlusScheduleFree", "FluxTrainerLossConfig": "Flux Trainer Loss Config", + "TrainNetworkConfig": "Train Network Config", } diff --git a/nodes_sdxl.py b/nodes_sdxl.py index 7badc9c..753fc9f 100644 --- a/nodes_sdxl.py +++ b/nodes_sdxl.py @@ -60,7 +60,7 @@ class InitSDXLLoRATraining: "optimizer_settings": ("ARGS",), "output_name": ("STRING", {"default": "SDXL_lora", "multiline": False}), "output_dir": ("STRING", {"default": "SDXL_trainer_output", "multiline": False, "tooltip": "path to dataset, root is the 'ComfyUI' folder, with windows portable 'ComfyUI_windows_portable'"}), - "network_dim": ("INT", {"default": 16, "min": 1, "max": 2048, "step": 1, "tooltip": "network dim"}), + "network_dim": ("INT", {"default": 16, "min": 1, "max": 100000, "step": 1, "tooltip": "network dim"}), "network_alpha": ("FLOAT", {"default": 16, "min": 0.0, "max": 2048.0, "step": 0.01, "tooltip": "network alpha"}), "learning_rate": ("FLOAT", {"default": 1e-6, "min": 0.0, "max": 10.0, "step": 0.0000001, "tooltip": "learning rate"}), "max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 100000, "step": 1, "tooltip": "max number of training steps"}), @@ -70,7 +70,7 @@ class InitSDXLLoRATraining: "blocks_to_swap": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "option for memory use reduction. The maximum number of blocks that can be swapped is 36 for SDXL.5L and 22 for SDXL.5M"}), "fp8_base": ("BOOLEAN", {"default": False, "tooltip": "use fp8 for base model"}), "gradient_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "the actual dtype training uses"}), - "save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "bf16", "tooltip": "the dtype to save checkpoints as"}), + "save_dtype": (["fp32", "fp16", "bf16", "fp8_e4m3fn", "fp8_e5m2"], {"default": "fp16", "tooltip": "the dtype to save checkpoints as"}), "attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}), "train_text_encoder": (['disabled', 'clip_l',], {"default": 'disabled', "tooltip": "also train the selected text encoders using specified dtype, T5 can not be trained without clip_l"}), "clip_l_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.000001, "tooltip": "text encoder learning rate"}), @@ -84,6 +84,7 @@ class InitSDXLLoRATraining: "resume_args": ("ARGS", {"default": "", "tooltip": "resume args to pass to the training command"}), "block_args": ("ARGS", {"default": "", "tooltip": "limit the blocks used in the LoRA"}), "loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}), + "network_config": ("NETWORK_CONFIG", {"tooltip": "additional network config"}), }, "hidden": { "prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO" @@ -97,7 +98,7 @@ class InitSDXLLoRATraining: def init_training(self, SDXL_models, dataset, optimizer_settings, sample_prompts_pos, sample_prompts_neg, output_name, attention_mode, gradient_dtype, save_dtype, additional_args=None, resume_args=None, train_text_encoder='disabled', - gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, clip_l_lr=0, clip_g_lr=0, loss_args=None, **kwargs): + gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, clip_l_lr=0, clip_g_lr=0, loss_args=None, network_config=None, **kwargs): mm.soft_empty_cache() output_dir = os.path.abspath(kwargs.get("output_dir")) @@ -158,20 +159,21 @@ class InitSDXLLoRATraining: "sample_prompts": positive_prompts, "negative_prompts": negative_prompts, "save_precision": save_dtype, - "mixed_precision": "fp16", + "mixed_precision": "bf16", "num_cpu_threads_per_process": 1, "pretrained_model_name_or_path": SDXL_models["checkpoint"], "save_model_as": "safetensors", "persistent_data_loader_workers": False, "max_data_loader_n_workers": 0, "seed": 42, - "network_module": ".networks.lora", + "network_module": ".networks.lora" if network_config is None else network_config["network_module"], "dataset_config": dataset_toml, "output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}", "loss_type": "l2", "alpha_mask": dataset["alpha_mask"], "network_train_unet_only": True if train_text_encoder == 'disabled' else False, "disable_mmap_load_safetensors": False, + "network_args": None if network_config is None else network_config["network_args"], } attention_settings = { "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},