Fixes, add output to ADE

This commit is contained in:
kijai
2024-02-09 19:59:53 +02:00
parent 59603aab49
commit 555af1dba0
4 changed files with 85 additions and 20 deletions
+2 -2
View File
@@ -129,8 +129,8 @@ def load_weights(
dreambooth_state_dict = torch.load(dreambooth_model_path, map_location="cpu")
# 1. vae
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, animation_pipeline.vae.config, strict=False)
animation_pipeline.vae.load_state_dict(converted_vae_checkpoint)
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, animation_pipeline.vae.config)
animation_pipeline.vae.load_state_dict(converted_vae_checkpoint, strict=False)
# 2. unet
converted_unet_checkpoint = convert_ldm_unet_checkpoint(dreambooth_state_dict, animation_pipeline.unet.config)
animation_pipeline.unet.load_state_dict(converted_unet_checkpoint, strict=False)
+25
View File
@@ -0,0 +1,25 @@
class MotionLoraInfo:
def __init__(self, name: str, strength: float = 1.0, hash: str=""):
self.name = name
self.strength = strength
self.hash = ""
def set_hash(self, hash: str):
self.hash = hash
def clone(self):
return MotionLoraInfo(self.name, self.strength, self.hash)
class MotionLoraList:
def __init__(self):
self.loras: list[MotionLoraInfo] = []
def add_lora(self, lora: MotionLoraInfo):
self.loras.append(lora)
def clone(self):
new_list = MotionLoraList()
for lora in self.loras:
new_list.add_lora(lora.clone())
return new_list
+55 -16
View File
@@ -25,6 +25,8 @@ from .animatediff.utils.util import save_videos_grid, load_diffusers_lora, load_
from .animatediff.utils.lora_handler import LoraHandler
from .animatediff.utils.lora import extract_lora_child_module
from .motion_lora import MotionLoraInfo, MotionLoraList
from lion_pytorch import Lion
import comfy.model_management
import comfy.utils
@@ -295,8 +297,8 @@ class AD_MotionDirector_train:
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES =("image",)
RETURN_TYPES = ("IMAGE", "STRING",)
RETURN_NAMES =("image", "lora_path",)
FUNCTION = "process"
CATEGORY = "AD_MotionDirector"
@@ -362,7 +364,7 @@ class AD_MotionDirector_train:
name = lora_name
date_calendar = datetime.datetime.now().strftime("%Y-%m-%d")
date_time = datetime.datetime.now().strftime("-%H-%M-%S")
date_time = datetime.datetime.now().strftime("%H-%M-%S")
folder_name = "debug" if is_debug else name + date_time
output_dir = os.path.join(script_directory, "outputs", date_calendar, folder_name)
@@ -382,9 +384,9 @@ class AD_MotionDirector_train:
# Handle the output folder creation
#lora_path = create_save_paths(output_dir)
spatial_lora_path = os.path.join(folder_paths.models_dir,"loras", "trained_spatial", date_calendar, date_time, lora_name)
temporal_lora_path = os.path.join(folder_paths.models_dir,"animatediff_motion_lora", lora_name, date_calendar, date_time)
#OmegaConf.save(config, os.path.join(output_dir, 'config.yaml'))
temporal_lora_path = os.path.join(folder_paths.models_dir,"animatediff_motion_lora", date_calendar, date_time, lora_name)
# Load scheduler, tokenizer and models.
noise_scheduler_kwargs.update({"steps_offset": 1})
@@ -706,7 +708,7 @@ class AD_MotionDirector_train:
step=global_step,
use_safetensors=True,
lora_rank=lora_rank,
lora_name=lora_name + "_spatial"
lora_name=lora_name + "_r"+ str(lora_rank) + "_spatial",
)
if lora_manager_temporal is not None:
@@ -716,7 +718,7 @@ class AD_MotionDirector_train:
step=global_step,
use_safetensors=True,
lora_rank=lora_rank,
lora_name=lora_name + "_temporal",
lora_name=lora_name + "_r"+ str(lora_rank) + "_temporal",
use_motion_lora_format=True
)
@@ -784,16 +786,17 @@ class AD_MotionDirector_train:
if global_step >= max_train_steps:
break
final_temporal_lora_name = os.path.join(date_calendar, date_time, lora_name, (str(max_train_epoch) + "_" + lora_name + "_r"+ str(lora_rank) + "_temporal_unet.safetensors"))
print(final_temporal_lora_name)
samples = samples.view(*samples.shape[1:])
samples = samples.permute(1, 2, 3, 0).cpu()
return (samples,)
return (samples, final_temporal_lora_name,)
import folder_paths
class DiffusersLoaderForTraining:
@classmethod
def IS_CHANGED(s):
return ""
#@classmethod
#def IS_CHANGED(s):
# return ""
@classmethod
def INPUT_TYPES(cls):
paths = []
@@ -834,7 +837,7 @@ class DiffusersLoaderForTraining:
if os.path.exists(path):
model_path = path
break
print("Model path:", model_path)
config = OmegaConf.load(os.path.join(script_directory, f"configs/training/motion_director/training.yaml"))
vae = AutoencoderKL.from_pretrained(model_path, subfolder="vae")
tokenizer = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
@@ -922,16 +925,52 @@ class ValidationSettings:
}
print(validation_settings)
return validation_settings,
class AD_MotionLoraLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"lora_path": ("STRING", {"multiline": False, "default": "",}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
},
"optional": {
"prev_motion_lora": ("MOTION_LORA",),
}
}
RETURN_TYPES = ("MOTION_LORA",)
CATEGORY = "AD_MotionDirector"
FUNCTION = "load_motion_lora"
def load_motion_lora(self, lora_path: str, strength: float, prev_motion_lora: MotionLoraList=None):
if prev_motion_lora is None:
prev_motion_lora = MotionLoraList()
else:
prev_motion_lora = prev_motion_lora.clone()
full_lora_path = os.path.join(folder_paths.models_dir,"animatediff_motion_lora",lora_path)
# check if motion lora with name exists
if not Path(full_lora_path).is_file():
raise FileNotFoundError(f"Motion lora not found at {full_lora_path}")
# create motion lora info to be loaded in AnimateDiff Loader
lora_name = os.path.basename(lora_path)
lora_info = MotionLoraInfo(name=lora_path, strength=strength)
prev_motion_lora.add_lora(lora_info)
return (prev_motion_lora,)
NODE_CLASS_MAPPINGS = {
"AD_MotionDirector_train": AD_MotionDirector_train,
"DiffusersLoaderForTraining": DiffusersLoaderForTraining,
"ValidationModelSelect": ValidationModelSelect,
"ValidationSettings": ValidationSettings
"ValidationSettings": ValidationSettings,
"AD_MotionLoraLoader": AD_MotionLoraLoader
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AD_MotionDirector_train": "AD_MotionDirector_train",
"DiffusersLoaderForTraining": "DiffusersLoaderForTraining",
"ValidationModelSelect": "ValidationModelSelect",
"ValidationSettings": "ValidationSettings"
"ValidationSettings": "ValidationSettings",
"AD_MotionLoraLoader": "AD_MotionLoraLoader"
}
+3 -2
View File
@@ -1,8 +1,9 @@
diffusers>=0.26.0
diffusers>=0.26.2
huggingface_hub>=0.20.3
transformers>=4.27.4
loralib
einops
omegaconf
lion-pytorch
peft
peft
imageio>=2.33.1