Support LyCORIS

https://github.com/KohakuBlueleaf/Lycoris
This commit is contained in:
kijai
2025-01-16 19:36:18 +02:00
parent 30cea9e372
commit 998968f5ff
32 changed files with 6124 additions and 8 deletions
+28
View File
@@ -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
+151
View File
@@ -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},
},
},
}
+9
View File
@@ -0,0 +1,9 @@
from .general import (
rebuild_tucker,
factorization,
power2factorization,
FUNC_LIST,
tucker_weight,
tucker_weight_from_conv,
apply_dora_scale,
)
+122
View File
@@ -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
+112
View File
@@ -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
+108
View File
@@ -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
+85
View File
@@ -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
+165
View File
@@ -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)
+247
View File
@@ -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
+676
View File
@@ -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)
+52
View File
@@ -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)
+46
View File
@@ -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
+315
View File
@@ -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
+255
View File
@@ -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)
+217
View File
@@ -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)
+156
View File
@@ -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)
+214
View File
@@ -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)
+262
View File
@@ -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)
+142
View File
@@ -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)
+332
View File
@@ -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)
+329
View File
@@ -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)
+609
View File
@@ -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))
+161
View File
@@ -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)
+483
View File
@@ -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")
+5
View File
@@ -0,0 +1,5 @@
def product(xs: list[int | float]):
res = 1
for x in xs:
res *= x
return res
+35
View File
@@ -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.")
+9
View File
@@ -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
+88
View File
@@ -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."
)
+13
View File
@@ -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
+640
View File
@@ -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)
+51 -3
View File
@@ -340,6 +340,50 @@ class OptimizerConfigProdigy:
return (kwargs,) 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: class OptimizerConfigProdigyPlusScheduleFree:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -390,7 +434,7 @@ class InitFluxLoRATraining:
"optimizer_settings": ("ARGS",), "optimizer_settings": ("ARGS",),
"output_name": ("STRING", {"default": "flux_lora", "multiline": False}), "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'"}), "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"}), "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"}), "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"}), "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"}), "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"}), "gradient_checkpointing": (["enabled", "enabled_with_cpu_offloading", "disabled"], {"default": "enabled", "tooltip": "use gradient checkpointing"}),
"loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}), "loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}),
"network_config": ("NETWORK_CONFIG", {"tooltip": "additional network config"}),
}, },
"hidden": { "hidden": {
"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO" "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, 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', 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() mm.soft_empty_cache()
output_dir = os.path.abspath(kwargs.get("output_dir")) output_dir = os.path.abspath(kwargs.get("output_dir"))
@@ -500,7 +545,7 @@ class InitFluxLoRATraining:
"persistent_data_loader_workers": False, "persistent_data_loader_workers": False,
"max_data_loader_n_workers": 0, "max_data_loader_n_workers": 0,
"seed": 42, "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, "dataset_config": dataset_toml,
"output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}", "output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}",
"loss_type": "l2", "loss_type": "l2",
@@ -509,6 +554,7 @@ class InitFluxLoRATraining:
"network_train_unet_only": True if train_text_encoder == 'disabled' else False, "network_train_unet_only": True if train_text_encoder == 'disabled' else False,
"fp8_base_unet": True if "fp8" in train_text_encoder else False, "fp8_base_unet": True if "fp8" in train_text_encoder else False,
"disable_mmap_load_safetensors": False, "disable_mmap_load_safetensors": False,
"network_args": None if network_config is None else network_config["network_args"],
} }
attention_settings = { attention_settings = {
"sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True}, "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},
@@ -1722,6 +1768,7 @@ NODE_CLASS_MAPPINGS = {
"FluxTrainAndValidateLoop": FluxTrainAndValidateLoop, "FluxTrainAndValidateLoop": FluxTrainAndValidateLoop,
"OptimizerConfigProdigyPlusScheduleFree": OptimizerConfigProdigyPlusScheduleFree, "OptimizerConfigProdigyPlusScheduleFree": OptimizerConfigProdigyPlusScheduleFree,
"FluxTrainerLossConfig": FluxTrainerLossConfig, "FluxTrainerLossConfig": FluxTrainerLossConfig,
"TrainNetworkConfig": TrainNetworkConfig,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"InitFluxLoRATraining": "Init Flux LoRA Training", "InitFluxLoRATraining": "Init Flux LoRA Training",
@@ -1748,4 +1795,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"FluxTrainAndValidateLoop": "Flux Train And Validate Loop", "FluxTrainAndValidateLoop": "Flux Train And Validate Loop",
"OptimizerConfigProdigyPlusScheduleFree": "Optimizer Config ProdigyPlusScheduleFree", "OptimizerConfigProdigyPlusScheduleFree": "Optimizer Config ProdigyPlusScheduleFree",
"FluxTrainerLossConfig": "Flux Trainer Loss Config", "FluxTrainerLossConfig": "Flux Trainer Loss Config",
"TrainNetworkConfig": "Train Network Config",
} }
+7 -5
View File
@@ -60,7 +60,7 @@ class InitSDXLLoRATraining:
"optimizer_settings": ("ARGS",), "optimizer_settings": ("ARGS",),
"output_name": ("STRING", {"default": "SDXL_lora", "multiline": False}), "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'"}), "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"}), "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"}), "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"}), "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"}), "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"}), "fp8_base": ("BOOLEAN", {"default": False, "tooltip": "use fp8 for base model"}),
"gradient_dtype": (["fp32", "fp16", "bf16"], {"default": "fp32", "tooltip": "the actual dtype training uses"}), "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"}), "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"}), "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"}), "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"}), "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"}), "block_args": ("ARGS", {"default": "", "tooltip": "limit the blocks used in the LoRA"}),
"loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}), "loss_args": ("ARGS", {"default": "", "tooltip": "loss args"}),
"network_config": ("NETWORK_CONFIG", {"tooltip": "additional network config"}),
}, },
"hidden": { "hidden": {
"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO" "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, 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_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() mm.soft_empty_cache()
output_dir = os.path.abspath(kwargs.get("output_dir")) output_dir = os.path.abspath(kwargs.get("output_dir"))
@@ -158,20 +159,21 @@ class InitSDXLLoRATraining:
"sample_prompts": positive_prompts, "sample_prompts": positive_prompts,
"negative_prompts": negative_prompts, "negative_prompts": negative_prompts,
"save_precision": save_dtype, "save_precision": save_dtype,
"mixed_precision": "fp16", "mixed_precision": "bf16",
"num_cpu_threads_per_process": 1, "num_cpu_threads_per_process": 1,
"pretrained_model_name_or_path": SDXL_models["checkpoint"], "pretrained_model_name_or_path": SDXL_models["checkpoint"],
"save_model_as": "safetensors", "save_model_as": "safetensors",
"persistent_data_loader_workers": False, "persistent_data_loader_workers": False,
"max_data_loader_n_workers": 0, "max_data_loader_n_workers": 0,
"seed": 42, "seed": 42,
"network_module": ".networks.lora", "network_module": ".networks.lora" if network_config is None else network_config["network_module"],
"dataset_config": dataset_toml, "dataset_config": dataset_toml,
"output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}", "output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}",
"loss_type": "l2", "loss_type": "l2",
"alpha_mask": dataset["alpha_mask"], "alpha_mask": dataset["alpha_mask"],
"network_train_unet_only": True if train_text_encoder == 'disabled' else False, "network_train_unet_only": True if train_text_encoder == 'disabled' else False,
"disable_mmap_load_safetensors": False, "disable_mmap_load_safetensors": False,
"network_args": None if network_config is None else network_config["network_args"],
} }
attention_settings = { attention_settings = {
"sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True}, "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},