Add LoraLoaderSequence node

This commit is contained in:
sylym
2023-04-05 18:45:46 +08:00
parent 5f31a6f687
commit 19cde8e86f
3 changed files with 150 additions and 3 deletions
+6 -1
View File
@@ -166,6 +166,11 @@ Load the checkpoint model into UNet3DConditionModel. Usually used to generate a
- ckpt_name
- The name of the checkpoint model. The model should be in the `models/checkpoints` folder.
### LoraLoaderSequence
Same function as `LoraLoader` node, but acts on UNet3DConditionModel. The input and output of the model are both of `ORIGINAL_MODEL` type.
### TrainUnetSequence
Fine-tune the incoming model using latent vector and context, and convert the model to inference mode.
@@ -188,7 +193,7 @@ Fine-tune the incoming model using latent vector and context, and convert the mo
- The number of steps to fine-tune the model. If the steps is 0, the model will not be fine-tuned.
### KSamplerSequence
Same function as KSampler node, but added support for noise vector and image mask sequence.
Same function as `KSampler` node, but added support for noise vector and image mask sequence.
## Limits
+22 -1
View File
@@ -5,7 +5,7 @@ import numpy as np
from PIL import Image
from comfy import model_management
import comfy.samplers
from .sd import load_checkpoint_guess_config
from .sd import load_checkpoint_guess_config, load_lora_for_models
from .convert_from_ckpt import convert_scheduler_checkpoint
from .tuneavideo.util import ddim_inversion
import comfy.utils
@@ -334,6 +334,26 @@ class TrainUnetSequence:
return (model_train,)
class LoraLoaderSequence:
@classmethod
def INPUT_TYPES(s):
return {"required": { "model": ("ORIGINAL_MODEL",),
"clip": ("CLIP", ),
"lora_name": (folder_paths.get_filename_list("loras"), ),
"strength_model": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
"strength_clip": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.01}),
}}
RETURN_TYPES = ("ORIGINAL_MODEL", "CLIP")
FUNCTION = "load_lora"
CATEGORY = "vid2vid"
def load_lora(self, model, clip, lora_name, strength_model, strength_clip):
lora_path = folder_paths.get_full_path("loras", lora_name)
model_lora, clip_lora = load_lora_for_models(model, clip, lora_path, strength_model, strength_clip)
return (model_lora, clip_lora)
NODE_CLASS_MAPPINGS = {
"LoadImageSequence": LoadImageSequence,
"LoadImageMaskSequence": LoadImageMaskSequence,
@@ -341,6 +361,7 @@ NODE_CLASS_MAPPINGS = {
"DdimInversionSequence": DdimInversionSequence,
"SetLatentNoiseSequence": SetLatentNoiseSequence,
"CheckpointLoaderSimpleSequence": CheckpointLoaderSimpleSequence,
"LoraLoaderSequence": LoraLoaderSequence,
"TrainUnetSequence": TrainUnetSequence,
"KSamplerSequence": KSamplerSequence,
}
+122 -1
View File
@@ -1,6 +1,6 @@
import torch
from comfy import model_management
from comfy.sd import load_model_weights, ModelPatcher, VAE, CLIP
from comfy.sd import load_model_weights, ModelPatcher, VAE, CLIP, model_lora_keys
from comfy import utils
from comfy import clip_vision
from comfy.ldm.util import instantiate_from_config
@@ -149,3 +149,124 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, o
# model = model.half()
return (ModelPatcher(model), clip, vae, clipvision)
def load_lora(path, to_load):
lora = utils.load_torch_file(path)
patch_dict = {}
loaded_keys = set()
for x in to_load:
alpha_name = "{}.alpha".format(x)
alpha = None
if alpha_name in lora.keys():
alpha = lora[alpha_name].item()
loaded_keys.add(alpha_name)
A_name = "{}.lora_up.weight".format(x)
B_name = "{}.lora_down.weight".format(x)
mid_name = "{}.lora_mid.weight".format(x)
if A_name in lora.keys():
mid = None
if mid_name in lora.keys():
mid = lora[mid_name]
loaded_keys.add(mid_name)
patch_dict[to_load[x]] = (lora[A_name], lora[B_name], alpha, mid)
loaded_keys.add(A_name)
loaded_keys.add(B_name)
hada_w1_a_name = "{}.hada_w1_a".format(x)
hada_w1_b_name = "{}.hada_w1_b".format(x)
hada_w2_a_name = "{}.hada_w2_a".format(x)
hada_w2_b_name = "{}.hada_w2_b".format(x)
hada_t1_name = "{}.hada_t1".format(x)
hada_t2_name = "{}.hada_t2".format(x)
if hada_w1_a_name in lora.keys():
hada_t1 = None
hada_t2 = None
if hada_t1_name in lora.keys():
hada_t1 = lora[hada_t1_name]
hada_t2 = lora[hada_t2_name]
loaded_keys.add(hada_t1_name)
loaded_keys.add(hada_t2_name)
patch_dict[to_load[x]] = (lora[hada_w1_a_name], lora[hada_w1_b_name], alpha, lora[hada_w2_a_name], lora[hada_w2_b_name], hada_t1, hada_t2)
loaded_keys.add(hada_w1_a_name)
loaded_keys.add(hada_w1_b_name)
loaded_keys.add(hada_w2_a_name)
loaded_keys.add(hada_w2_b_name)
return patch_dict
def use_lora(pretrained_LoRA_path, model, alpha):
LORA_PREFIX_UNET = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
state_dict = utils.load_torch_file(pretrained_LoRA_path)
visited = []
# directly update weight in diffusers model
for key in state_dict:
# it is suggested to print out the key, it usually will be something like below
# "lora_te_text_model_encoder_layers_0_self_attn_k_proj.lora_down.weight"
# as we have set the alpha beforehand, so just skip
if ".alpha" in key or key in visited:
continue
if "text" in key:
continue
else:
layer_infos = key.split(".")[0].split(LORA_PREFIX_UNET + "_")[-1].split("_")
curr_layer = model.model.model.diffusion_model
# find the target layer
temp_name = layer_infos.pop(0)
while len(layer_infos) > -1:
try:
curr_layer = curr_layer.__getattr__(temp_name)
if len(layer_infos) > 0:
temp_name = layer_infos.pop(0)
elif len(layer_infos) == 0:
break
except Exception:
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
temp_name = layer_infos.pop(0)
pair_keys = []
if "lora_down" in key:
pair_keys.append(key.replace("lora_down", "lora_up"))
pair_keys.append(key)
else:
pair_keys.append(key)
pair_keys.append(key.replace("lora_up", "lora_down"))
# update weight
if len(state_dict[pair_keys[0]].shape) == 4:
weight_up = state_dict[pair_keys[0]].squeeze(3).squeeze(2).to(torch.float32)
weight_down = state_dict[pair_keys[1]].squeeze(3).squeeze(2).to(torch.float32)
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3)
else:
weight_up = state_dict[pair_keys[0]].to(torch.float32)
weight_down = state_dict[pair_keys[1]].to(torch.float32)
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down)
# update visited list
for item in pair_keys:
visited.append(item)
return model
def load_lora_for_models(model, clip, lora_path, strength_model, strength_clip):
key_map = model_lora_keys(model.model)
key_map = model_lora_keys(clip.cond_stage_model, key_map)
loaded = load_lora(lora_path, key_map)
new_modelpatcher = model.clone()
new_modelpatcher = use_lora(lora_path, new_modelpatcher, strength_model)
new_clip = clip.clone()
new_clip.add_patches(loaded, strength_clip)
return (new_modelpatcher, new_clip)