lora_merger: Node V3スキーマに移行
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
0cbd5045ae
commit
eff60ffa4f
+53
-53
@@ -2,7 +2,8 @@ import comfy
|
||||
import folder_paths
|
||||
import os
|
||||
import re
|
||||
from ... import ROOT_NAME
|
||||
from comfy_api.v0_0_2 import io
|
||||
from ... import ROOT_NAME, NODE_SURFIX, SYMBOL
|
||||
|
||||
CATEGORY_NAME = ROOT_NAME + "lora_merger"
|
||||
|
||||
@@ -63,73 +64,72 @@ LBW12TO20 = [1, 2, 3, 4, 7, 17, 18, 19]
|
||||
|
||||
MID_ID = {26:13, 20:10}
|
||||
|
||||
class LoraLoaderFromWeight:
|
||||
def __init__(self):
|
||||
self.loaded_lora = None
|
||||
class LoraLoaderFromWeight(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
return io.Schema(
|
||||
node_id=f"LoraLoaderFromWeight{NODE_SURFIX}",
|
||||
display_name=f"LoRA Loader From Weight {SYMBOL}",
|
||||
category=CATEGORY_NAME,
|
||||
inputs=[
|
||||
io.Custom("LoRA").Input("lora"),
|
||||
io.Model.Input("model"),
|
||||
io.Clip.Input("clip_optional", optional=True),
|
||||
],
|
||||
outputs=[
|
||||
io.Model.Output(),
|
||||
io.Clip.Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"lora": ("LoRA", ),
|
||||
"model": ("MODEL",),
|
||||
},
|
||||
"optional": {
|
||||
"clip_optional": ("CLIP", ),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("MODEL", "CLIP")
|
||||
FUNCTION = "load_lora_from_weight"
|
||||
|
||||
CATEGORY = CATEGORY_NAME
|
||||
|
||||
def load_lora_from_weight(self, lora, model, clip_optional=None):
|
||||
def execute(cls, lora, model, clip_optional=None) -> io.NodeOutput:
|
||||
lora_weight = lora["lora"]
|
||||
strength_model = lora["strength_model"]
|
||||
strength_clip = lora["strength_clip"]
|
||||
|
||||
if strength_model == 0 and strength_clip == 0:
|
||||
return (model, clip_optional)
|
||||
return io.NodeOutput(model, clip_optional)
|
||||
|
||||
model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip_optional, lora_weight, strength_model, strength_clip)
|
||||
return (model_lora, clip_lora)
|
||||
return io.NodeOutput(model_lora, clip_lora)
|
||||
|
||||
class LoraLoaderWeightOnly:
|
||||
def __init__(self):
|
||||
self.loaded_lora = None
|
||||
self.lbw = None
|
||||
# module-level cache replacing the old per-instance `self.loaded_lora` /
|
||||
# `self.lbw` state (execute() is a classmethod, no `self` to cache on).
|
||||
_weight_only_cache = {"loaded_lora": None, "lbw": None}
|
||||
|
||||
class LoraLoaderWeightOnly(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
return io.Schema(
|
||||
node_id=f"LoraLoaderWeightOnly{NODE_SURFIX}",
|
||||
display_name=f"LoRA Loader Weight Only {SYMBOL}",
|
||||
category=CATEGORY_NAME,
|
||||
inputs=[
|
||||
io.Combo.Input("lora_name", options=folder_paths.get_filename_list("loras")),
|
||||
io.Float.Input("strength_model", default=1.0, min=-20.0, max=20.0, step=0.01),
|
||||
io.Float.Input("strength_clip", default=1.0, min=-20.0, max=20.0, step=0.01),
|
||||
io.String.Input("lbw", multiline=False, default=""),
|
||||
],
|
||||
outputs=[
|
||||
io.Custom("LoRA").Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"lora_name": (folder_paths.get_filename_list("loras"), ),
|
||||
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||
"strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||
"lbw": ("STRING", {
|
||||
"multiline": False,
|
||||
"default": ""
|
||||
}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("LoRA", )
|
||||
FUNCTION = "load_lora_weight_only"
|
||||
|
||||
CATEGORY = CATEGORY_NAME
|
||||
|
||||
def load_lora_weight_only(self, lora_name, strength_model, strength_clip, lbw):
|
||||
def execute(cls, lora_name, strength_model, strength_clip, lbw) -> io.NodeOutput:
|
||||
lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||
lora = None
|
||||
|
||||
if self.loaded_lora is not None:
|
||||
if self.loaded_lora[0] == lora_path:
|
||||
lora = self.loaded_lora[1]
|
||||
if _weight_only_cache["loaded_lora"] is not None:
|
||||
if _weight_only_cache["loaded_lora"][0] == lora_path:
|
||||
lora = _weight_only_cache["loaded_lora"][1]
|
||||
else:
|
||||
temp = self.loaded_lora
|
||||
self.loaded_lora = None
|
||||
temp = _weight_only_cache["loaded_lora"]
|
||||
_weight_only_cache["loaded_lora"] = None
|
||||
del temp
|
||||
|
||||
if lora is None or self.lbw != lbw:
|
||||
if lora is None or _weight_only_cache["lbw"] != lbw:
|
||||
lora = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||
if lbw != "":
|
||||
weight_list = parse_weight_list(lbw)
|
||||
@@ -177,7 +177,7 @@ class LoraLoaderWeightOnly:
|
||||
if alpha_key in lora:
|
||||
del lora[alpha_key]
|
||||
|
||||
self.loaded_lora = (lora_path, lora)
|
||||
self.lbw = lbw
|
||||
_weight_only_cache["loaded_lora"] = (lora_path, lora)
|
||||
_weight_only_cache["lbw"] = lbw
|
||||
|
||||
return ({"lora": lora, "strength_model": strength_model, "strength_clip": strength_clip}, )
|
||||
return io.NodeOutput({"lora": lora, "strength_model": strength_model, "strength_clip": strength_clip})
|
||||
|
||||
@@ -1,56 +1,58 @@
|
||||
import comfy
|
||||
import math
|
||||
import torch
|
||||
from ... import ROOT_NAME
|
||||
from comfy_api.v0_0_2 import io
|
||||
from ... import ROOT_NAME, NODE_SURFIX, SYMBOL
|
||||
|
||||
CATEGORY_NAME = ROOT_NAME + "lora_merger"
|
||||
CLAMP_QUANTILE = 0.99
|
||||
REGULAR_LORA = "regular"
|
||||
DIFFUSERS_LORA = "diffusers"
|
||||
|
||||
class LoraMerge:
|
||||
def __init__(self):
|
||||
self.loaded_lora = None
|
||||
class LoraMerge(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
return io.Schema(
|
||||
node_id=f"LoraMerger{NODE_SURFIX}",
|
||||
display_name=f"LoRA Merge {SYMBOL}",
|
||||
category=CATEGORY_NAME,
|
||||
inputs=[
|
||||
io.Custom("LoRA").Input("lora_1"),
|
||||
io.Combo.Input("mode", options=["add", "concat", "svd", "svd_fast"]),
|
||||
io.Int.Input(
|
||||
"rank",
|
||||
default=16, # Minimum value
|
||||
min=1,
|
||||
max=320, # Maximum value
|
||||
step=1, # Slider's step
|
||||
display_mode=io.NumberDisplay.number, # Cosmetic only: display as "number" or "slider"
|
||||
),
|
||||
io.Float.Input(
|
||||
"threshold",
|
||||
default=1.0,
|
||||
min=0,
|
||||
max=1,
|
||||
step=0.01,
|
||||
),
|
||||
io.Combo.Input("device", options=["cuda", "cpu"]),
|
||||
io.Combo.Input("dtype", options=["float32", "float16", "bfloat16"]),
|
||||
io.Custom("LoRA").Input("lora_2", optional=True),
|
||||
],
|
||||
outputs=[
|
||||
io.Custom("LoRA").Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"lora_1": ("LoRA",),
|
||||
"mode": (["add", "concat", "svd", "svd_fast"], ),
|
||||
"rank": ("INT", {
|
||||
"default": 16,
|
||||
"min": 1, #Minimum value
|
||||
"max": 320, #Maximum value
|
||||
"step": 1, #Slider's step
|
||||
"display": "number" # Cosmetic only: display as "number" or "slider"
|
||||
}),
|
||||
"threshold": ("FLOAT", {
|
||||
"default": 1.0,
|
||||
"min": 0,
|
||||
"max": 1,
|
||||
"step": 0.01,
|
||||
}),
|
||||
"device": (["cuda", "cpu"], ),
|
||||
"dtype": (["float32", "float16", "bfloat16"], ),
|
||||
},
|
||||
"optional": {
|
||||
"lora_2": ("LoRA",),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("LoRA", )
|
||||
FUNCTION = "lora_merge"
|
||||
def execute(cls, lora_1, lora_2=None, mode=None, rank=None, threshold=None, device=None, dtype=None) -> io.NodeOutput:
|
||||
|
||||
CATEGORY = CATEGORY_NAME
|
||||
lora = cls.merge(lora_1, lora_2, mode, rank, threshold, device, dtype)
|
||||
|
||||
def lora_merge(self, lora_1, lora_2=None, mode=None, rank=None, threshold=None, device=None, dtype=None):
|
||||
|
||||
lora = self.merge(lora_1, lora_2, mode, rank, threshold, device, dtype)
|
||||
return io.NodeOutput(lora)
|
||||
|
||||
return (lora, )
|
||||
|
||||
@staticmethod
|
||||
@torch.no_grad()
|
||||
def merge(self, lora_1, lora_2, mode, rank, threshold, device, dtype):
|
||||
def merge(lora_1, lora_2, mode, rank, threshold, device, dtype):
|
||||
# lora = up @ down * alpha / rank
|
||||
|
||||
weight = {}
|
||||
@@ -123,31 +125,32 @@ class LoraMerge:
|
||||
|
||||
return {"lora":weight, "strength_model":1, "strength_clip":1}
|
||||
|
||||
class LoraSVDRank:
|
||||
def __init__(self):
|
||||
self.loaded_lora = None
|
||||
class LoraSVDRank(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
return io.Schema(
|
||||
node_id=f"LoraSVDRank{NODE_SURFIX}",
|
||||
display_name=f"LoRA SVD Rank {SYMBOL}",
|
||||
category=CATEGORY_NAME,
|
||||
inputs=[
|
||||
io.Custom("LoRA").Input("lora"),
|
||||
io.Float.Input(
|
||||
"threshold",
|
||||
default=1.0,
|
||||
min=0,
|
||||
max=1,
|
||||
step=0.001,
|
||||
),
|
||||
io.Combo.Input("device", options=["cuda", "cpu"]),
|
||||
],
|
||||
outputs=[
|
||||
io.String.Output(),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"lora": ("LoRA",),
|
||||
"threshold": ("FLOAT", {
|
||||
"default": 1.0,
|
||||
"min": 0,
|
||||
"max": 1,
|
||||
"step": 0.001,
|
||||
}),
|
||||
"device": (["cuda", "cpu"], ),
|
||||
},
|
||||
}
|
||||
RETURN_TYPES = ("STRING", )
|
||||
FUNCTION = "show"
|
||||
|
||||
CATEGORY = CATEGORY_NAME
|
||||
|
||||
@torch.no_grad()
|
||||
def show(self, lora, threshold, device):
|
||||
def execute(cls, lora, threshold, device) -> io.NodeOutput:
|
||||
|
||||
keys = lora_module_keys(lora)
|
||||
pber = comfy.utils.ProgressBar(len(keys))
|
||||
@@ -158,8 +161,8 @@ class LoraSVDRank:
|
||||
index = svd_show(up, down, threshold, device)
|
||||
content += f"{key}: {index}\n"
|
||||
pber.update(1)
|
||||
|
||||
return (content, )
|
||||
|
||||
return io.NodeOutput(content)
|
||||
|
||||
@torch.no_grad()
|
||||
def calc_up_down_alpha(key, lora, add=True):
|
||||
|
||||
+19
-18
@@ -2,28 +2,29 @@ import comfy
|
||||
import folder_paths
|
||||
import math
|
||||
import os
|
||||
from ... import ROOT_NAME
|
||||
from comfy_api.v0_0_2 import io
|
||||
from ... import ROOT_NAME, NODE_SURFIX, SYMBOL
|
||||
|
||||
CATEGORY_NAME = ROOT_NAME + "lora_merger"
|
||||
|
||||
class LoraSave:
|
||||
def __init__(self):
|
||||
self.loaded_lora = None
|
||||
class LoraSave(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls) -> io.Schema:
|
||||
return io.Schema(
|
||||
node_id=f"LoraSave{NODE_SURFIX}",
|
||||
display_name=f"LoRA Save {SYMBOL}",
|
||||
category=CATEGORY_NAME,
|
||||
inputs=[
|
||||
io.Custom("LoRA").Input("lora"),
|
||||
io.String.Input("file_name", multiline=False, default="merged"),
|
||||
io.Combo.Input("extension", options=["safetensors"]),
|
||||
],
|
||||
outputs=[],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "lora": ("LoRA",),
|
||||
"file_name": ("STRING", {"multiline": False, "default": "merged"}),
|
||||
"extension": (["safetensors"], ),
|
||||
}}
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "lora_save"
|
||||
|
||||
CATEGORY = CATEGORY_NAME
|
||||
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def lora_save(self, lora, file_name, extension):
|
||||
def execute(cls, lora, file_name, extension) -> io.NodeOutput:
|
||||
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:
|
||||
@@ -44,7 +45,7 @@ class LoraSave:
|
||||
print(f"Saving LoRA to {save_path}")
|
||||
comfy.utils.save_torch_file(new_state_dict, save_path)
|
||||
|
||||
return {}
|
||||
return io.NodeOutput()
|
||||
|
||||
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