lora_merger: diffusers形式のLoRAとsvd_fastモードに対応

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
laksjdjf
2026-07-04 08:06:38 +09:00
co-authored by Claude Fable 5
parent d7376518a0
commit f80628bfc0
3 changed files with 130 additions and 37 deletions
+10 -3
View File
@@ -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
View File
@@ -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
+8 -4
View File
@@ -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()}