This commit is contained in:
thecooltechguy
2023-11-23 10:34:01 +00:00
parent 72b74f5db8
commit 0cef54066e
2 changed files with 224 additions and 145 deletions
+177 -145
View File
@@ -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",
}