wip
This commit is contained in:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user