From 0cef54066e36e71f4fe9f950ae2b1884e584381e Mon Sep 17 00:00:00 2001 From: thecooltechguy Date: Thu, 23 Nov 2023 10:34:01 +0000 Subject: [PATCH] wip --- nodes.py | 322 ++++++++++++++++++++++++++++++------------------------- svd.py | 47 ++++++++ 2 files changed, 224 insertions(+), 145 deletions(-) diff --git a/nodes.py b/nodes.py index c35c1f3..bddc436 100644 --- a/nodes.py +++ b/nodes.py @@ -1,146 +1,19 @@ -from .svd import load_model +from .svd import load_model, get_unique_embedder_keys_from_conditioner, get_batch import gc import folder_paths import torch import os - -class StableVideoDiffusion: - """ - Node for applying Stable Video Diffusion - - Class methods - ------------- - INPUT_TYPES (dict): - Tell the main program input parameters of nodes. - - Attributes - ---------- - RETURN_TYPES (`tuple`): - The type of each element in the output tulple. - RETURN_NAMES (`tuple`): - Optional: The name of each output in the output tulple. - FUNCTION (`str`): - The name of the entry-point method. For example, if `FUNCTION = "execute"` then it will run Example().execute() - OUTPUT_NODE ([`bool`]): - If this node is an output node that outputs a result/image from the graph. The SaveImage node is an example. - The backend iterates on these output nodes and tries to execute all their parents if their parent graph is properly connected. - Assumed to be False if not present. - CATEGORY (`str`): - The category the node should appear in the UI. - execute(s) -> tuple || None: - The entry point method. The name of this method must be the same as the value of property `FUNCTION`. - For example, if `FUNCTION = "execute"` then this method's name must be `execute`, if `FUNCTION = "foo"` then it must be `foo`. - """ - def __init__(self): - self.svd_model = None - - @classmethod - def INPUT_TYPES(s): - """ - Return a dictionary which contains config for all input fields. - Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT". - Input types "INT", "STRING" or "FLOAT" are special values for fields on the node. - The type can be a list for selection. - - Returns: `dict`: - - Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required` - - Value input_fields (`dict`): Contains input fields config: - * Key field_name (`string`): Name of a entry-point method's argument - * Value field_config (`tuple`): - + First value is a string indicate the type of field or a list for selection. - + Secound value is a config for type "INT", "STRING" or "FLOAT". - """ - return { - "required": { - "image": ("IMAGE",), - "model" : ("MODEL",), - # "int_field": ("INT", { - # "default": 0, - # "min": 0, #Minimum value - # "max": 4096, #Maximum value - # "step": 64, #Slider's step - # "display": "number" # Cosmetic only: display as "number" or "slider" - # }), - # "float_field": ("FLOAT", { - # "default": 1.0, - # "min": 0.0, - # "max": 10.0, - # "step": 0.01, - # "round": 0.001, #The value represeting the precision to round to, will be set to the step value by default. Can be set to False to disable rounding. - # "display": "number"}), - # "print_to_screen": (["enable", "disable"],), - # "string_field": ("STRING", { - # "multiline": False, #True if you want the field to look like the one on the ClipTextEncode node - # "default": "Hello World!" - # }), - }, - } - - RETURN_TYPES = ("IMAGE",) - #RETURN_NAMES = ("image_output_name",) - - FUNCTION = "sample_video" - - #OUTPUT_NODE = False - - CATEGORY = "Comfy Stable Video Diffusion" - - def sample_video(self, image, model): - print(type(image)) - print(image) - print(type(model)) - print(model) - return (image,) - +import math class SVDModelLoader: - """ - Node for applying Stable Video Diffusion - - Class methods - ------------- - INPUT_TYPES (dict): - Tell the main program input parameters of nodes. - - Attributes - ---------- - RETURN_TYPES (`tuple`): - The type of each element in the output tulple. - RETURN_NAMES (`tuple`): - Optional: The name of each output in the output tulple. - FUNCTION (`str`): - The name of the entry-point method. For example, if `FUNCTION = "execute"` then it will run Example().execute() - OUTPUT_NODE ([`bool`]): - If this node is an output node that outputs a result/image from the graph. The SaveImage node is an example. - The backend iterates on these output nodes and tries to execute all their parents if their parent graph is properly connected. - Assumed to be False if not present. - CATEGORY (`str`): - The category the node should appear in the UI. - execute(s) -> tuple || None: - The entry point method. The name of this method must be the same as the value of property `FUNCTION`. - For example, if `FUNCTION = "execute"` then this method's name must be `execute`, if `FUNCTION = "foo"` then it must be `foo`. - """ def __init__(self): self.svd_model = None @classmethod def INPUT_TYPES(s): - """ - Return a dictionary which contains config for all input fields. - Some types (string): "MODEL", "VAE", "CLIP", "CONDITIONING", "LATENT", "IMAGE", "INT", "STRING", "FLOAT". - Input types "INT", "STRING" or "FLOAT" are special values for fields on the node. - The type can be a list for selection. - - Returns: `dict`: - - Key input_fields_group (`string`): Can be either required, hidden or optional. A node class must have property `required` - - Value input_fields (`dict`): Contains input fields config: - * Key field_name (`string`): Name of a entry-point method's argument - * Value field_config (`tuple`): - + First value is a string indicate the type of field or a list for selection. - + Secound value is a config for type "INT", "STRING" or "FLOAT". - """ checkpoints = folder_paths.get_filename_list("svd") configs = folder_paths.get_filename_list("svd_configs") + devices = [] if torch.cuda.is_available(): devices.append("cuda") @@ -151,9 +24,6 @@ class SVDModelLoader: "checkpoint" : (checkpoints, { "default" : checkpoints[0], }), - "config" : (configs, { - "default" : configs[0], - }), "num_frames" : ("INT", { "default": 14, "min" : 0, @@ -166,20 +36,20 @@ class SVDModelLoader: }, } - RETURN_TYPES = ("MODEL", "CLIP") - #RETURN_NAMES = ("image_output_name",) + RETURN_TYPES = ("MODEL",) FUNCTION = "load_svd_model" CATEGORY = "Comfy Stable Video Diffusion" - def load_svd_model(self, checkpoint, config, num_frames, num_steps, device): + def load_svd_model(self, checkpoint, num_frames, num_steps, device): if self.svd_model is not None: del self.svd_model gc.collect() self.svd_model = None - print("Loading model...") - config = os.path.join(folder_paths.get_folder_paths("svd_configs")[0], config) + print("Loading SVD model...") + checkpoint_filename_without_extension = os.path.splitext(checkpoint)[0] + config = os.path.join(folder_paths.get_folder_paths("svd_configs")[0], f"{checkpoint_filename_without_extension}.yaml") checkpoint = os.path.join(folder_paths.get_folder_paths("svd")[0], checkpoint) self.svd_model = load_model( config=config, @@ -188,19 +58,181 @@ class SVDModelLoader: num_steps=num_steps, checkpoint=checkpoint, ) - conditioner = self.svd_model.conditioner - clip_model = conditioner.embedders[0].open_clip.model - print("Model loaded!") - return (self.svd_model, clip_model) + print("Loaded SVD model!") + return (self.svd_model,) + +class SVDSampler: + @classmethod + def INPUT_TYPES(s): + devices = [] + if torch.cuda.is_available(): + devices.append("cuda") + devices.append("cpu") + return { + "required": { + "image": ("IMAGE",), + "model" : ("MODEL",), + "motion_bucket_id" : ("INT", { + "default": 127, + }), + "fps_id" : ("INT", { + "default": 6, + }), + "cond_aug" : ("FLOAT", { + "default": 0.02, + }), + "seed" : ("INT", { + "default": 23, + }), + "decoding_t" : ("INT", { + "default": 14, + }), + "device" : (devices,), + }, + } + + RETURN_TYPES = ("LATENT",) + #RETURN_NAMES = ("image_output_name",) + + FUNCTION = "sample_video" + + #OUTPUT_NODE = False + + CATEGORY = "Comfy Stable Video Diffusion" + + def sample_video(self, image, model, motion_bucket_id, fps_id, cond_aug, seed, decoding_t, device): + # convert image tensor to PIL image + print(type(image)) + print(image.shape) + 1/0 + + if image.mode == "RGBA": + image = image.convert("RGB") + w, h = image.size + + if h % 64 != 0 or w % 64 != 0: + width, height = map(lambda x: x - x % 64, (w, h)) + image = image.resize((width, height)) + print( + f"WARNING: Your image is of size {h}x{w} which is not divisible by 64. We are resizing to {height}x{width}!" + ) + + image = ToTensor()(image) + image = image * 2.0 - 1.0 + + image = image.unsqueeze(0).to(device) + H, W = image.shape[2:] + assert image.shape[1] == 3 + F = 8 + C = 4 + shape = (num_frames, C, H // F, W // F) + if (H, W) != (576, 1024): + print( + "WARNING: The conditioning frame you provided is not 576x1024. This leads to suboptimal performance as model was only trained on 576x1024. Consider increasing `cond_aug`." + ) + if motion_bucket_id > 255: + print( + "WARNING: High motion bucket! This may lead to suboptimal performance." + ) + + if fps_id < 5: + print("WARNING: Small fps value! This may lead to suboptimal performance.") + + if fps_id > 30: + print("WARNING: Large fps value! This may lead to suboptimal performance.") + + value_dict = {} + value_dict["motion_bucket_id"] = motion_bucket_id + value_dict["fps_id"] = fps_id + value_dict["cond_aug"] = cond_aug + value_dict["cond_frames_without_noise"] = image + value_dict["cond_frames"] = image + cond_aug * torch.randn_like(image) + value_dict["cond_aug"] = cond_aug + + with torch.no_grad(): + with torch.autocast(device): + batch, batch_uc = get_batch( + get_unique_embedder_keys_from_conditioner(model.conditioner), + value_dict, + [1, num_frames], + T=num_frames, + device=device, + ) + c, uc = model.conditioner.get_unconditional_conditioning( + batch, + batch_uc=batch_uc, + force_uc_zero_embeddings=[ + "cond_frames", + "cond_frames_without_noise", + ], + ) + + for k in ["crossattn", "concat"]: + uc[k] = repeat(uc[k], "b ... -> b t ...", t=num_frames) + uc[k] = rearrange(uc[k], "b t ... -> (b t) ...", t=num_frames) + c[k] = repeat(c[k], "b ... -> b t ...", t=num_frames) + c[k] = rearrange(c[k], "b t ... -> (b t) ...", t=num_frames) + + randn = torch.randn(shape, device=device) + + additional_model_inputs = {} + additional_model_inputs["image_only_indicator"] = torch.zeros( + 2, num_frames + ).to(device) + additional_model_inputs["num_video_frames"] = batch["num_video_frames"] + + def denoiser(input, sigma, c): + return model.denoiser( + model.model, input, sigma, c, **additional_model_inputs + ) + + samples_z = model.sampler(denoiser, randn, cond=c, uc=uc) + return (samples_z,) + + +class SVDDecoder: + @classmethod + def INPUT_TYPES(s): + devices = [] + if torch.cuda.is_available(): + devices.append("cuda") + devices.append("cpu") + return { + "required": { + "samples_z": ("LATENT",), + "model" : ("MODEL",), + "decoding_t" : ("INT", { + "default": 14, + }), + "device" : (devices,), + }, + } + + RETURN_TYPES = ("IMAGE",) + + FUNCTION = "decode" + + CATEGORY = "Comfy Stable Video Diffusion" + + def decode(self, samples_z, model, decoding_t, device): + with torch.no_grad(): + with torch.autocast(device): + model.en_and_decode_n_samples_a_time = decoding_t + samples_x = model.decode_first_stage(samples_z) + samples = torch.clamp((samples_x + 1.0) / 2.0, min=0.0, max=1.0) + return (samples,) + # A dictionary that contains all nodes you want to export with their names NODE_CLASS_MAPPINGS = { "SVDModelLoader" : SVDModelLoader, - "StableVideoDiffusion": StableVideoDiffusion + "SVDSampler": SVDSampler, + "SVDDecoder": SVDDecoder, } # A dictionary that contains the friendly/humanly readable titles for the nodes NODE_DISPLAY_NAME_MAPPINGS = { - "SVDModelLoader" : "Load Stable Video Diffusion Model", - "StableVideoDiffusion": "Stable Video Diffusion" + "SVDModelLoader" : "Load SVD Model", + "SVDSampler": "SVD Sampler", + "SVDDecoder": "SVD Decoder", } \ No newline at end of file diff --git a/svd.py b/svd.py index a3a7870..7c981b3 100644 --- a/svd.py +++ b/svd.py @@ -3,8 +3,55 @@ from omegaconf import OmegaConf from .util import default, instantiate_from_config import torch import folder_paths +import math +import torch +from einops import rearrange, repeat from typing import Optional +def get_unique_embedder_keys_from_conditioner(conditioner): + return list(set([x.input_key for x in conditioner.embedders])) + +def get_batch(keys, value_dict, N, T, device): + batch = {} + batch_uc = {} + + for key in keys: + if key == "fps_id": + batch[key] = ( + torch.tensor([value_dict["fps_id"]]) + .to(device) + .repeat(int(math.prod(N))) + ) + elif key == "motion_bucket_id": + batch[key] = ( + torch.tensor([value_dict["motion_bucket_id"]]) + .to(device) + .repeat(int(math.prod(N))) + ) + elif key == "cond_aug": + batch[key] = repeat( + torch.tensor([value_dict["cond_aug"]]).to(device), + "1 -> b", + b=math.prod(N), + ) + elif key == "cond_frames": + batch[key] = repeat(value_dict["cond_frames"], "1 ... -> b ...", b=N[0]) + elif key == "cond_frames_without_noise": + batch[key] = repeat( + value_dict["cond_frames_without_noise"], "1 ... -> b ...", b=N[0] + ) + else: + batch[key] = value_dict[key] + + if T is not None: + batch["num_video_frames"] = T + + for key in batch.keys(): + if key not in batch_uc and isinstance(batch[key], torch.Tensor): + batch_uc[key] = torch.clone(batch[key]) + return batch, batch_uc + + def load_model( config: str, device: str,