Fixes, add output to ADE
This commit is contained in:
@@ -129,8 +129,8 @@ def load_weights(
|
|||||||
dreambooth_state_dict = torch.load(dreambooth_model_path, map_location="cpu")
|
dreambooth_state_dict = torch.load(dreambooth_model_path, map_location="cpu")
|
||||||
|
|
||||||
# 1. vae
|
# 1. vae
|
||||||
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, animation_pipeline.vae.config, strict=False)
|
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, animation_pipeline.vae.config)
|
||||||
animation_pipeline.vae.load_state_dict(converted_vae_checkpoint)
|
animation_pipeline.vae.load_state_dict(converted_vae_checkpoint, strict=False)
|
||||||
# 2. unet
|
# 2. unet
|
||||||
converted_unet_checkpoint = convert_ldm_unet_checkpoint(dreambooth_state_dict, animation_pipeline.unet.config)
|
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)
|
animation_pipeline.unet.load_state_dict(converted_unet_checkpoint, strict=False)
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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_handler import LoraHandler
|
||||||
from .animatediff.utils.lora import extract_lora_child_module
|
from .animatediff.utils.lora import extract_lora_child_module
|
||||||
|
|
||||||
|
from .motion_lora import MotionLoraInfo, MotionLoraList
|
||||||
|
|
||||||
from lion_pytorch import Lion
|
from lion_pytorch import Lion
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
@@ -295,8 +297,8 @@ class AD_MotionDirector_train:
|
|||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE",)
|
RETURN_TYPES = ("IMAGE", "STRING",)
|
||||||
RETURN_NAMES =("image",)
|
RETURN_NAMES =("image", "lora_path",)
|
||||||
FUNCTION = "process"
|
FUNCTION = "process"
|
||||||
|
|
||||||
CATEGORY = "AD_MotionDirector"
|
CATEGORY = "AD_MotionDirector"
|
||||||
@@ -362,7 +364,7 @@ class AD_MotionDirector_train:
|
|||||||
|
|
||||||
name = lora_name
|
name = lora_name
|
||||||
date_calendar = datetime.datetime.now().strftime("%Y-%m-%d")
|
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
|
folder_name = "debug" if is_debug else name + date_time
|
||||||
|
|
||||||
output_dir = os.path.join(script_directory, "outputs", date_calendar, folder_name)
|
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
|
# Handle the output folder creation
|
||||||
#lora_path = create_save_paths(output_dir)
|
#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)
|
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)
|
|
||||||
|
temporal_lora_path = os.path.join(folder_paths.models_dir,"animatediff_motion_lora", date_calendar, date_time, lora_name)
|
||||||
#OmegaConf.save(config, os.path.join(output_dir, 'config.yaml'))
|
|
||||||
|
|
||||||
# Load scheduler, tokenizer and models.
|
# Load scheduler, tokenizer and models.
|
||||||
noise_scheduler_kwargs.update({"steps_offset": 1})
|
noise_scheduler_kwargs.update({"steps_offset": 1})
|
||||||
@@ -706,7 +708,7 @@ class AD_MotionDirector_train:
|
|||||||
step=global_step,
|
step=global_step,
|
||||||
use_safetensors=True,
|
use_safetensors=True,
|
||||||
lora_rank=lora_rank,
|
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:
|
if lora_manager_temporal is not None:
|
||||||
@@ -716,7 +718,7 @@ class AD_MotionDirector_train:
|
|||||||
step=global_step,
|
step=global_step,
|
||||||
use_safetensors=True,
|
use_safetensors=True,
|
||||||
lora_rank=lora_rank,
|
lora_rank=lora_rank,
|
||||||
lora_name=lora_name + "_temporal",
|
lora_name=lora_name + "_r"+ str(lora_rank) + "_temporal",
|
||||||
use_motion_lora_format=True
|
use_motion_lora_format=True
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -784,16 +786,17 @@ class AD_MotionDirector_train:
|
|||||||
|
|
||||||
if global_step >= max_train_steps:
|
if global_step >= max_train_steps:
|
||||||
break
|
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.view(*samples.shape[1:])
|
||||||
samples = samples.permute(1, 2, 3, 0).cpu()
|
samples = samples.permute(1, 2, 3, 0).cpu()
|
||||||
return (samples,)
|
return (samples, final_temporal_lora_name,)
|
||||||
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
class DiffusersLoaderForTraining:
|
class DiffusersLoaderForTraining:
|
||||||
@classmethod
|
#@classmethod
|
||||||
def IS_CHANGED(s):
|
#def IS_CHANGED(s):
|
||||||
return ""
|
# return ""
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
paths = []
|
paths = []
|
||||||
@@ -834,7 +837,7 @@ class DiffusersLoaderForTraining:
|
|||||||
if os.path.exists(path):
|
if os.path.exists(path):
|
||||||
model_path = path
|
model_path = path
|
||||||
break
|
break
|
||||||
|
print("Model path:", model_path)
|
||||||
config = OmegaConf.load(os.path.join(script_directory, f"configs/training/motion_director/training.yaml"))
|
config = OmegaConf.load(os.path.join(script_directory, f"configs/training/motion_director/training.yaml"))
|
||||||
vae = AutoencoderKL.from_pretrained(model_path, subfolder="vae")
|
vae = AutoencoderKL.from_pretrained(model_path, subfolder="vae")
|
||||||
tokenizer = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
|
tokenizer = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
|
||||||
@@ -922,16 +925,52 @@ class ValidationSettings:
|
|||||||
}
|
}
|
||||||
print(validation_settings)
|
print(validation_settings)
|
||||||
return 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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"AD_MotionDirector_train": AD_MotionDirector_train,
|
"AD_MotionDirector_train": AD_MotionDirector_train,
|
||||||
"DiffusersLoaderForTraining": DiffusersLoaderForTraining,
|
"DiffusersLoaderForTraining": DiffusersLoaderForTraining,
|
||||||
"ValidationModelSelect": ValidationModelSelect,
|
"ValidationModelSelect": ValidationModelSelect,
|
||||||
"ValidationSettings": ValidationSettings
|
"ValidationSettings": ValidationSettings,
|
||||||
|
"AD_MotionLoraLoader": AD_MotionLoraLoader
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"AD_MotionDirector_train": "AD_MotionDirector_train",
|
"AD_MotionDirector_train": "AD_MotionDirector_train",
|
||||||
"DiffusersLoaderForTraining": "DiffusersLoaderForTraining",
|
"DiffusersLoaderForTraining": "DiffusersLoaderForTraining",
|
||||||
"ValidationModelSelect": "ValidationModelSelect",
|
"ValidationModelSelect": "ValidationModelSelect",
|
||||||
"ValidationSettings": "ValidationSettings"
|
"ValidationSettings": "ValidationSettings",
|
||||||
|
"AD_MotionLoraLoader": "AD_MotionLoraLoader"
|
||||||
}
|
}
|
||||||
+3
-2
@@ -1,8 +1,9 @@
|
|||||||
diffusers>=0.26.0
|
diffusers>=0.26.2
|
||||||
huggingface_hub>=0.20.3
|
huggingface_hub>=0.20.3
|
||||||
transformers>=4.27.4
|
transformers>=4.27.4
|
||||||
loralib
|
loralib
|
||||||
einops
|
einops
|
||||||
omegaconf
|
omegaconf
|
||||||
lion-pytorch
|
lion-pytorch
|
||||||
peft
|
peft
|
||||||
|
imageio>=2.33.1
|
||||||
Reference in New Issue
Block a user