This commit is contained in:
thecooltechguy
2023-11-23 11:21:52 +00:00
parent 0cef54066e
commit a04ccd5dd3
+41 -18
View File
@@ -1,9 +1,12 @@
from .svd import load_model, get_unique_embedder_keys_from_conditioner, get_batch from .svd import load_model, get_unique_embedder_keys_from_conditioner, get_batch
from torchvision.transforms import ToTensor, ToPILImage
from einops import rearrange, repeat
import gc import gc
import folder_paths import folder_paths
import torch import torch
import os import os
import math import math
import numpy as np
class SVDModelLoader: class SVDModelLoader:
def __init__(self): def __init__(self):
@@ -84,9 +87,6 @@ class SVDSampler:
"seed" : ("INT", { "seed" : ("INT", {
"default": 23, "default": 23,
}), }),
"decoding_t" : ("INT", {
"default": 14,
}),
"device" : (devices,), "device" : (devices,),
}, },
} }
@@ -100,31 +100,37 @@ class SVDSampler:
CATEGORY = "Comfy Stable Video Diffusion" CATEGORY = "Comfy Stable Video Diffusion"
def sample_video(self, image, model, motion_bucket_id, fps_id, cond_aug, seed, decoding_t, device): def sample_video(self, image, model, motion_bucket_id, fps_id, cond_aug, seed, device):
# convert image tensor to PIL image # convert image torch tensor to PIL image
print(type(image)) # image shape: (1, H, W, C)
print(image.shape) image = image.squeeze(0)
1/0 image = image.permute(2, 0, 1)
image = (image + 1.0) / 2.0
image = image.clamp(min=0.0, max=1.0)
image = ToPILImage()(image)
if image.mode == "RGBA": if image.mode == "RGBA":
image = image.convert("RGB") image = image.convert("RGB")
w, h = image.size
w, h = image.size
if h % 64 != 0 or w % 64 != 0: if h % 64 != 0 or w % 64 != 0:
width, height = map(lambda x: x - x % 64, (w, h)) width, height = map(lambda x: x - x % 64, (w, h))
image = image.resize((width, height)) image = image.resize((width, height))
print( print(
f"WARNING: Your image is of size {h}x{w} which is not divisible by 64. We are resizing to {height}x{width}!" 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 = ToTensor()(image)
image = image * 2.0 - 1.0 image = image * 2.0 - 1.0
image = image.unsqueeze(0).to(device) image = image.unsqueeze(0).to(device)
H, W = image.shape[2:] H, W = image.shape[2:]
assert image.shape[1] == 3 assert image.shape[1] == 3
F = 8 F = 8
C = 4 C = 4
num_frames = model.sampler.guider.num_frames
shape = (num_frames, C, H // F, W // F) shape = (num_frames, C, H // F, W // F)
if (H, W) != (576, 1024): if (H, W) != (576, 1024):
print( print(
@@ -149,8 +155,11 @@ class SVDSampler:
value_dict["cond_frames"] = image + cond_aug * torch.randn_like(image) value_dict["cond_frames"] = image + cond_aug * torch.randn_like(image)
value_dict["cond_aug"] = cond_aug value_dict["cond_aug"] = cond_aug
all_samples_z = []
with torch.no_grad(): with torch.no_grad():
with torch.autocast(device): with torch.autocast(device):
print("getting batch")
batch, batch_uc = get_batch( batch, batch_uc = get_batch(
get_unique_embedder_keys_from_conditioner(model.conditioner), get_unique_embedder_keys_from_conditioner(model.conditioner),
value_dict, value_dict,
@@ -158,6 +167,8 @@ class SVDSampler:
T=num_frames, T=num_frames,
device=device, device=device,
) )
print("getting conditioning")
c, uc = model.conditioner.get_unconditional_conditioning( c, uc = model.conditioner.get_unconditional_conditioning(
batch, batch,
batch_uc=batch_uc, batch_uc=batch_uc,
@@ -167,6 +178,7 @@ class SVDSampler:
], ],
) )
print("repeating conditioning")
for k in ["crossattn", "concat"]: for k in ["crossattn", "concat"]:
uc[k] = repeat(uc[k], "b ... -> b t ...", t=num_frames) uc[k] = repeat(uc[k], "b ... -> b t ...", t=num_frames)
uc[k] = rearrange(uc[k], "b t ... -> (b t) ...", t=num_frames) uc[k] = rearrange(uc[k], "b t ... -> (b t) ...", t=num_frames)
@@ -186,7 +198,10 @@ class SVDSampler:
model.model, input, sigma, c, **additional_model_inputs model.model, input, sigma, c, **additional_model_inputs
) )
print("sampling: ", randn.shape)
samples_z = model.sampler(denoiser, randn, cond=c, uc=uc) samples_z = model.sampler(denoiser, randn, cond=c, uc=uc)
print("sampled: ", samples_z.shape)
return (samples_z,) return (samples_z,)
@@ -215,12 +230,20 @@ class SVDDecoder:
CATEGORY = "Comfy Stable Video Diffusion" CATEGORY = "Comfy Stable Video Diffusion"
def decode(self, samples_z, model, decoding_t, device): def decode(self, samples_z, model, decoding_t, device):
print("decoding: ", samples_z.shape)
with torch.no_grad(): with torch.no_grad():
with torch.autocast(device): with torch.autocast(device):
model.en_and_decode_n_samples_a_time = decoding_t model.en_and_decode_n_samples_a_time = decoding_t
samples_x = model.decode_first_stage(samples_z) samples_x = model.decode_first_stage(samples_z)
samples = torch.clamp((samples_x + 1.0) / 2.0, min=0.0, max=1.0) samples = torch.clamp((samples_x + 1.0) / 2.0, min=0.0, max=1.0)
return (samples,)
vid = (
(rearrange(samples, "t c h w -> t h w c") * 255)
.cpu()
.numpy()
.astype(np.uint8)
)
return (vid,)
# A dictionary that contains all nodes you want to export with their names # A dictionary that contains all nodes you want to export with their names