lora_merger: diffusers形式のLoRAとsvd_fastモードに対応
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
d7376518a0
commit
f80628bfc0
@@ -139,7 +139,7 @@ class LoraLoaderWeightOnly:
|
||||
|
||||
strength_clip = strength_clip * weight_list[0]
|
||||
|
||||
up_keys = [key for key in lora.keys() if "lora_up" in key and not "lora_te" in key]
|
||||
up_keys = [key for key in lora.keys() if ("lora_up" in key or "lora_B" in key) and not "lora_te" in key]
|
||||
|
||||
for key in up_keys:
|
||||
ids = extract_numbers(key)
|
||||
@@ -166,9 +166,16 @@ class LoraLoaderWeightOnly:
|
||||
if weight != 0.0:
|
||||
lora[key] = lora[key] * weight
|
||||
else:
|
||||
if "lora_up" in key:
|
||||
down_key = key.replace("lora_up", "lora_down")
|
||||
alpha_key = key.replace("lora_up.weight", "alpha")
|
||||
else:
|
||||
down_key = key.replace("lora_B", "lora_A")
|
||||
alpha_key = key.replace("lora_B.weight", "alpha")
|
||||
del lora[key]
|
||||
del lora[key.replace("lora_up", "lora_down")]
|
||||
del lora[key.replace("lora_up.weight", "alpha")]
|
||||
del lora[down_key]
|
||||
if alpha_key in lora:
|
||||
del lora[alpha_key]
|
||||
|
||||
self.loaded_lora = (lora_path, lora)
|
||||
self.lbw = lbw
|
||||
|
||||
+112
-30
@@ -5,6 +5,8 @@ from ... import ROOT_NAME
|
||||
|
||||
CATEGORY_NAME = ROOT_NAME + "lora_merger"
|
||||
CLAMP_QUANTILE = 0.99
|
||||
REGULAR_LORA = "regular"
|
||||
DIFFUSERS_LORA = "diffusers"
|
||||
|
||||
class LoraMerge:
|
||||
def __init__(self):
|
||||
@@ -15,7 +17,7 @@ class LoraMerge:
|
||||
return {
|
||||
"required": {
|
||||
"lora_1": ("LoRA",),
|
||||
"mode": (["add", "concat", "svd"], ),
|
||||
"mode": (["add", "concat", "svd", "svd_fast"], ),
|
||||
"rank": ("INT", {
|
||||
"default": 16,
|
||||
"min": 1, #Minimum value
|
||||
@@ -57,32 +59,36 @@ class LoraMerge:
|
||||
if lora_2 is None:
|
||||
lora_2 = {"lora":{}, "strength_model":0, "strength_clip":0}
|
||||
|
||||
keys_1 = [key[: key.rfind(".lora_down")] for key in lora_1["lora"].keys() if ".lora_down" in key]
|
||||
keys_2 = [key[: key.rfind(".lora_down")] for key in lora_2["lora"].keys() if ".lora_down" in key]
|
||||
keys_1 = lora_module_keys(lora_1)
|
||||
keys_2 = lora_module_keys(lora_2)
|
||||
keys = list(set(keys_1 + keys_2))
|
||||
print(f"Merging {len(keys)} modules")
|
||||
print(f"{len(keys)-len(keys_1)} modules only in lora_1")
|
||||
print(f"{len(keys)-len(keys_2)} modules only in lora_2")
|
||||
print(f"{len(keys)-len(keys_2)} modules only in lora_1")
|
||||
print(f"{len(keys)-len(keys_1)} modules only in lora_2")
|
||||
pber = comfy.utils.ProgressBar(len(keys))
|
||||
|
||||
for key in keys:
|
||||
output_format = lora_key_format(key, lora_1) or lora_key_format(key, lora_2) or REGULAR_LORA
|
||||
|
||||
if key not in keys_1:
|
||||
up, down, alpha = calc_up_down_alpha(key, lora_2)
|
||||
if mode == "svd":
|
||||
up, down = svd_merge(up, down, None, None, rank, threshold, device)
|
||||
if mode in ("svd", "svd_fast"):
|
||||
up, down = svd_merge(up, down, None, None, rank, threshold, device, fast=mode=="svd_fast")
|
||||
elif key not in keys_2:
|
||||
up, down, alpha = calc_up_down_alpha(key, lora_1)
|
||||
if mode == "svd":
|
||||
up, down = svd_merge(up, down, None, None, rank, threshold, device)
|
||||
if mode in ("svd", "svd_fast"):
|
||||
up, down = svd_merge(up, down, None, None, rank, threshold, device, fast=mode=="svd_fast")
|
||||
else:
|
||||
up_1, down_1, alpha_1 = calc_up_down_alpha(key, lora_1, add=mode!="add")
|
||||
up_2, down_2, alpha_2 = calc_up_down_alpha(key, lora_2, add=mode!="add")
|
||||
|
||||
alpha = alpha_1
|
||||
alpha_1_value = alpha_to_float(alpha_1)
|
||||
alpha_2_value = alpha_to_float(alpha_2)
|
||||
|
||||
# Scale to match alpha_1
|
||||
up_2 = up_2 * math.sqrt(alpha_2/alpha)
|
||||
down_2 = down_2 * math.sqrt(alpha_2/alpha)
|
||||
up_2 = up_2 * math.sqrt(alpha_2_value/alpha_1_value)
|
||||
down_2 = down_2 * math.sqrt(alpha_2_value/alpha_1_value)
|
||||
|
||||
up_1 = up_1.to(dtype=dtype)
|
||||
down_1 = down_1.to(dtype=dtype)
|
||||
@@ -104,12 +110,10 @@ class LoraMerge:
|
||||
scale_2 = math.sqrt((r_1+r_2)/r_2)
|
||||
up = torch.cat([up_1*scale_1, up_2*scale_2], dim=1)
|
||||
down = torch.cat([down_1*scale_1, down_2*scale_2], dim=0)
|
||||
elif mode == "svd":
|
||||
up, down = svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device)
|
||||
elif mode in ("svd", "svd_fast"):
|
||||
up, down = svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device, fast=mode=="svd_fast")
|
||||
|
||||
weight[key + ".lora_up.weight"] = up
|
||||
weight[key + ".lora_down.weight"] = down
|
||||
weight[key + ".alpha"] = alpha
|
||||
set_up_down_alpha(weight, key, up, down, alpha, output_format)
|
||||
|
||||
pber.update(1)
|
||||
|
||||
@@ -132,7 +136,7 @@ class LoraSVDRank:
|
||||
"default": 1.0,
|
||||
"min": 0,
|
||||
"max": 1,
|
||||
"step": 0.01,
|
||||
"step": 0.001,
|
||||
}),
|
||||
"device": (["cuda", "cpu"], ),
|
||||
},
|
||||
@@ -145,7 +149,7 @@ class LoraSVDRank:
|
||||
@torch.no_grad()
|
||||
def show(self, lora, threshold, device):
|
||||
|
||||
keys = [key[: key.rfind(".lora_down")] for key in lora["lora"].keys() if ".lora_down" in key]
|
||||
keys = lora_module_keys(lora)
|
||||
pber = comfy.utils.ProgressBar(len(keys))
|
||||
|
||||
content = ""
|
||||
@@ -159,8 +163,13 @@ class LoraSVDRank:
|
||||
|
||||
@torch.no_grad()
|
||||
def calc_up_down_alpha(key, lora, add=True):
|
||||
up_key = key + ".lora_up.weight"
|
||||
down_key = key + ".lora_down.weight"
|
||||
lora_format = lora_key_format(key, lora)
|
||||
if lora_format == DIFFUSERS_LORA:
|
||||
up_key = key + ".lora_B.weight"
|
||||
down_key = key + ".lora_A.weight"
|
||||
else:
|
||||
up_key = key + ".lora_up.weight"
|
||||
down_key = key + ".lora_down.weight"
|
||||
alpha_key = key + ".alpha"
|
||||
|
||||
is_te = "lora_te" in key
|
||||
@@ -171,14 +180,47 @@ def calc_up_down_alpha(key, lora, add=True):
|
||||
|
||||
up = lora["lora"][up_key] * sqrt_scale * sign_scale
|
||||
down = lora["lora"][down_key] * sqrt_scale
|
||||
alpha = lora["lora"][alpha_key]
|
||||
alpha = lora["lora"].get(alpha_key)
|
||||
if alpha is None:
|
||||
alpha = torch.tensor(down.shape[0], dtype=torch.float32, device=down.device)
|
||||
|
||||
return up, down, alpha
|
||||
|
||||
def lora_module_keys(lora):
|
||||
keys = set()
|
||||
for key in lora["lora"].keys():
|
||||
if key.endswith(".lora_down.weight"):
|
||||
keys.add(key[: key.rfind(".lora_down.weight")])
|
||||
elif key.endswith(".lora_A.weight"):
|
||||
keys.add(key[: key.rfind(".lora_A.weight")])
|
||||
return list(keys)
|
||||
|
||||
def lora_key_format(key, lora):
|
||||
state_dict = lora["lora"]
|
||||
if key + ".lora_up.weight" in state_dict and key + ".lora_down.weight" in state_dict:
|
||||
return REGULAR_LORA
|
||||
if key + ".lora_B.weight" in state_dict and key + ".lora_A.weight" in state_dict:
|
||||
return DIFFUSERS_LORA
|
||||
return None
|
||||
|
||||
def set_up_down_alpha(weight, key, up, down, alpha, lora_format):
|
||||
if lora_format == DIFFUSERS_LORA:
|
||||
weight[key + ".lora_B.weight"] = up
|
||||
weight[key + ".lora_A.weight"] = down
|
||||
else:
|
||||
weight[key + ".lora_up.weight"] = up
|
||||
weight[key + ".lora_down.weight"] = down
|
||||
weight[key + ".alpha"] = alpha
|
||||
|
||||
def alpha_to_float(alpha):
|
||||
if torch.is_tensor(alpha):
|
||||
return float(alpha.detach().cpu())
|
||||
return float(alpha)
|
||||
|
||||
# frovenius normによるrankの計算
|
||||
def index_sv_fro(S, target):
|
||||
def index_sv_fro(S, target, total_sq=None):
|
||||
S_squared = S.pow(2)
|
||||
s_fro_sq = float(torch.sum(S_squared))
|
||||
s_fro_sq = float(torch.sum(S_squared) if total_sq is None else total_sq)
|
||||
sum_S_squared = torch.cumsum(S_squared, dim=0)/s_fro_sq
|
||||
index = int(torch.searchsorted(sum_S_squared, target**2)) + 1
|
||||
index = max(1, min(index, len(S)-1))
|
||||
@@ -186,10 +228,13 @@ def index_sv_fro(S, target):
|
||||
return index
|
||||
|
||||
@torch.no_grad()
|
||||
def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None):
|
||||
def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None, fast=False):
|
||||
org_device = up_1.device
|
||||
org_dtype = up_1.dtype
|
||||
|
||||
if up_2 is None and threshold >= 1 and rank == up_1.shape[1]:
|
||||
return up_1.contiguous(), down_1.contiguous()
|
||||
|
||||
up_1 = up_1.to(device)
|
||||
down_1 = down_1.to(device)
|
||||
r_1 = up_1.shape[1]
|
||||
@@ -206,11 +251,17 @@ def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None):
|
||||
|
||||
weight = weight.to(dtype=torch.float32) # SVD only supports float32
|
||||
|
||||
U, S, Vh = torch.linalg.svd(weight)
|
||||
total_sq = torch.sum(weight.pow(2)) if fast and threshold < 1 else None
|
||||
|
||||
if fast:
|
||||
U, S, Vh = svd_lowrank(weight, rank, threshold, total_sq=total_sq)
|
||||
else:
|
||||
U, S, Vh = torch.linalg.svd(weight, full_matrices=False)
|
||||
|
||||
if threshold < 1:
|
||||
rank = index_sv_fro(S, threshold) + 1
|
||||
rank = index_sv_fro(S, threshold, total_sq=total_sq) + 1
|
||||
|
||||
rank = min(rank, len(S))
|
||||
U = U[:, :rank]
|
||||
S = S[:rank]
|
||||
U = U @ torch.diag(S)
|
||||
@@ -228,11 +279,40 @@ def svd_merge(up_1, down_1, up_2, down_2, rank, threshold, device=None):
|
||||
U = U.reshape(up_1.shape[0], rank, 1, 1)
|
||||
Vh = Vh.reshape(rank, down_1.shape[1], down_1.shape[2], down_1.shape[3])
|
||||
|
||||
up = U.to(org_device, dtype=org_dtype) * math.sqrt(rank)
|
||||
down = Vh.to(org_device, dtype=org_dtype) * math.sqrt(rank)
|
||||
up = (U.to(org_device, dtype=org_dtype) * math.sqrt(rank)).contiguous()
|
||||
down = (Vh.to(org_device, dtype=org_dtype) * math.sqrt(rank)).contiguous()
|
||||
|
||||
return up, down
|
||||
|
||||
@torch.no_grad()
|
||||
def svd_lowrank(weight, rank, threshold, total_sq=None, oversample=8, niter=2):
|
||||
max_rank = min(weight.shape)
|
||||
q = min(max_rank, max(1, rank + oversample))
|
||||
|
||||
if threshold < 1:
|
||||
q = min(max_rank, max(q, 32))
|
||||
|
||||
while True:
|
||||
U, S, V = torch.svd_lowrank(weight, q=q, niter=niter)
|
||||
order = torch.argsort(S, descending=True)
|
||||
U = U[:, order]
|
||||
S = S[order]
|
||||
V = V[:, order]
|
||||
|
||||
if threshold >= 1 or q >= max_rank:
|
||||
break
|
||||
|
||||
estimated_rank = index_sv_fro(S, threshold, total_sq=total_sq) + 1
|
||||
if estimated_rank < len(S) - 1:
|
||||
break
|
||||
|
||||
next_q = min(max_rank, q * 2)
|
||||
if next_q == q:
|
||||
break
|
||||
q = next_q
|
||||
|
||||
return U, S, V.T
|
||||
|
||||
@torch.no_grad()
|
||||
def svd_show(up, down, threshold, device):
|
||||
up = up.to(device)
|
||||
@@ -241,8 +321,10 @@ def svd_show(up, down, threshold, device):
|
||||
weight = up.view(-1, rank) @ down.view(rank, -1)
|
||||
weight = weight.to(dtype=torch.float32) # SVD only supports float32
|
||||
|
||||
U, S, Vh = torch.linalg.svd(weight)
|
||||
U, S, Vh = torch.linalg.svd(weight, full_matrices=False)
|
||||
if threshold < 1:
|
||||
index = index_sv_fro(S, threshold)
|
||||
else:
|
||||
index = rank
|
||||
|
||||
return index
|
||||
return index
|
||||
|
||||
@@ -27,20 +27,24 @@ class LoraSave:
|
||||
save_path = os.path.join(folder_paths.folder_names_and_paths["loras"][0][0], file_name + "." + extension)
|
||||
|
||||
if lora["strength_model"] == 1 and lora["strength_clip"] == 1:
|
||||
new_state_dict = lora["lora"]
|
||||
new_state_dict = make_contiguous(lora["lora"])
|
||||
else:
|
||||
new_state_dict = {}
|
||||
for key in lora["lora"].keys():
|
||||
scale = lora["strength_clip"] if "lora_te" in key else lora["strength_model"]
|
||||
sqrt_scale = math.sqrt(abs(scale))
|
||||
sign_scale = 1 if scale >= 0 else -1
|
||||
if "lora_up" in key:
|
||||
if "lora_up" in key or "lora_B" in key:
|
||||
new_state_dict[key] = lora["lora"][key] * sqrt_scale * sign_scale
|
||||
elif "lora_down" in key:
|
||||
elif "lora_down" in key or "lora_A" in key:
|
||||
new_state_dict[key] = lora["lora"][key] * sqrt_scale
|
||||
else:
|
||||
new_state_dict[key] = lora["lora"][key]
|
||||
new_state_dict = make_contiguous(new_state_dict)
|
||||
print(f"Saving LoRA to {save_path}")
|
||||
comfy.utils.save_torch_file(new_state_dict, save_path)
|
||||
|
||||
return {}
|
||||
return {}
|
||||
|
||||
def make_contiguous(state_dict):
|
||||
return {key: value.contiguous() if hasattr(value, "contiguous") else value for key, value in state_dict.items()}
|
||||
|
||||
Reference in New Issue
Block a user