@@ -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
|
||||||
@@ -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},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
@@ -0,0 +1,9 @@
|
|||||||
|
from .general import (
|
||||||
|
rebuild_tucker,
|
||||||
|
factorization,
|
||||||
|
power2factorization,
|
||||||
|
FUNC_LIST,
|
||||||
|
tucker_weight,
|
||||||
|
tucker_weight_from_conv,
|
||||||
|
apply_dora_scale,
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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))
|
||||||
@@ -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)
|
||||||
@@ -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")
|
||||||
@@ -0,0 +1,5 @@
|
|||||||
|
def product(xs: list[int | float]):
|
||||||
|
res = 1
|
||||||
|
for x in xs:
|
||||||
|
res *= x
|
||||||
|
return res
|
||||||
@@ -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.")
|
||||||
@@ -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
|
||||||
@@ -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."
|
||||||
|
)
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -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
@@ -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},
|
||||||
|
|||||||
Reference in New Issue
Block a user