diff --git a/README.md b/README.md index 37423a7..ee15e0f 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/__init__.py b/__init__.py index f5eb101..5514f0f 100644 --- a/__init__.py +++ b/__init__.py @@ -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, } diff --git a/sd.py b/sd.py index 27762be..ba4c4ef 100644 --- a/sd.py +++ b/sd.py @@ -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) \ No newline at end of file