Author SHA1 Message Date
kijai 70b3f8b0e5 Update nodes.py 2024-08-16 00:17:03 +03:00
kijai 2e107312c0 Merge branch 'main' into develop 2024-08-16 00:15:38 +03:00
kijai df0f87f020 Update nodes.py 2024-08-16 00:14:08 +03:00
kijai c55b6b8a20 Update nodes.py 2024-08-16 00:13:00 +03:00
kijai 7670b09c28 Fix other nodes that don't use controlnet (yet) 2024-07-20 19:22:46 +03:00
kijai 282c69c5f3 Update nodes.py 2024-07-20 00:27:22 +03:00
kijai 18121e97a2 Update nodes.py 2024-07-19 22:08:31 +03:00
kijai 4a7dd402b4 Support other SD 2.1 controlnets as well 2024-07-19 20:27:52 +03:00
kijai 2545a3afb7 Initial sketch controlnet support 2024-07-19 19:50:26 +03:00
kijai 0668ee6d0a fix DynamiCrafterBatchInterpolation too 2024-07-09 16:15:03 +03:00
kijai 69fe0e2e9f don't normalize image2 twice 2024-07-09 16:06:04 +03:00
6 changed files with 672 additions and 898 deletions
+17 -5
View File
@@ -387,6 +387,7 @@ class LatentDiffusion(DDPM):
logdir=None,
rand_cond_frame=False,
en_and_decode_n_samples_a_time=None,
control_scale=1.0,
*args, **kwargs):
self.num_timesteps_cond = default(num_timesteps_cond, 1)
self.scale_by_std = scale_by_std
@@ -404,6 +405,7 @@ class LatentDiffusion(DDPM):
self.loop_video = loop_video
self.fps_condition_type = fps_condition_type
self.perframe_ae = perframe_ae
self.control_scale = control_scale
self.logdir = logdir
self.rand_cond_frame = rand_cond_frame
@@ -565,7 +567,6 @@ class LatentDiffusion(DDPM):
return self.p_losses(x, c, t, **kwargs)
def apply_model(self, x_noisy, t, cond, **kwargs):
print("apply_model:", kwargs)
if isinstance(cond, dict):
# hybrid case, cond is exptected to be a dict
pass
@@ -573,9 +574,20 @@ class LatentDiffusion(DDPM):
if not isinstance(cond, list):
cond = [cond]
key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn'
cond = {key: cond}
cond = {key: [cond[0]]}
x_recon = self.model(x_noisy, t, **cond, **kwargs)
control_cond = cond["control_cond"]
if control_cond is not None:
control_cond = rearrange(control_cond, 'b c t h w-> (b t) c h w')
control_x = rearrange(x_noisy, 'b c t h w-> (b t) c h w')
control_context = repeat(cond["c_crossattn"][0], "b c l-> (repeat b) c l", repeat=16)
control = self.control_model(x=control_x, hint=control_cond, timesteps=t, context=control_context)
control = [c * self.control_model.control_scale for c in control]
else:
control = None
x_recon = self.model(x_noisy, t, c_crossattn=cond["c_crossattn"], c_concat=cond["c_concat"], control=control, **kwargs)
if isinstance(x_recon, tuple):
return x_recon[0]
@@ -723,7 +735,7 @@ class DiffusionWrapper(pl.LightningModule):
self.diffusion_model = instantiate_from_config(diff_model_config)
self.conditioning_key = conditioning_key
def forward(self, x, t, c_concat: list = None, c_crossattn: list = None,
def forward(self, x, t, c_concat: list = None, c_crossattn: list = None, control = None,
c_adm=None, s=None, mask=None, **kwargs):
# temporal_context = fps is foNone
if self.conditioning_key is None:
@@ -738,7 +750,7 @@ class DiffusionWrapper(pl.LightningModule):
## it is just right [b,c,t,h,w]: concatenate in channel dim
xc = torch.cat([x] + c_concat, dim=1)
cc = torch.cat(c_crossattn, 1)
out = self.diffusion_model(xc, t, context=cc, **kwargs)
out = self.diffusion_model(xc, t, context=cc, control=control, **kwargs)
elif self.conditioning_key == 'resblockcond':
cc = c_crossattn[0]
out = self.diffusion_model(x, t, context=cc)
+3 -8
View File
@@ -197,7 +197,7 @@ class DDIMSampler(object):
corrector_kwargs=corrector_kwargs,
unconditional_guidance_scale=unconditional_guidance_scale,
unconditional_conditioning=unconditional_conditioning,
mask=mask,x0=x0,step=i,fs=fs,guidance_rescale=guidance_rescale,
mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale,
**kwargs)
@@ -214,7 +214,7 @@ class DDIMSampler(object):
@torch.no_grad()
def p_sample_ddim(self, x, c, t, sigmas, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
unconditional_guidance_scale=1., unconditional_conditioning=None, step=0, ddim_edit=0,
unconditional_guidance_scale=1., unconditional_conditioning=None,
uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs):
b, *_, device = *x.shape, x.device
if x.dim() == 5:
@@ -225,13 +225,8 @@ class DDIMSampler(object):
if unconditional_conditioning is None or unconditional_guidance_scale == 1.:
model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser
else:
if step < ddim_edit:
print("applying ddim_edit step", step)
print(kwargs)
e_t_cond = self.model.apply_model(x, t, c, use_freetraj=True, **kwargs)
e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, use_freetraj=True, **kwargs)
### do_classifier_free_guidance
elif isinstance(c, torch.Tensor) or isinstance(c, dict):
if isinstance(c, torch.Tensor) or isinstance(c, dict):
e_t_cond = self.model.apply_model(x, t, c, **kwargs)
e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs)
else:
+11 -237
View File
@@ -18,36 +18,6 @@ from ...lvdm.basics import zero_module
import comfy.ops
ops = comfy.ops.manual_cast
from ...utils.utils_freetraj import get_path, plan_path
import math
def gaussian_2d(x=0, y=0, mx=0, my=0, sx=1, sy=1):
""" 2d Gaussian weight function
"""
gaussian_map = (
1
/ (2 * math.pi * sx * sy)
* torch.exp(-((x - mx) ** 2 / (2 * sx**2) + (y - my) ** 2 / (2 * sy**2)))
)
gaussian_map.div_(gaussian_map.max())
return gaussian_map
def gaussian_weight(height=32, width=32, KERNEL_DIVISION=3.0):
x = torch.linspace(0, height, height)
y = torch.linspace(0, width, width)
x, y = torch.meshgrid(x, y, indexing="ij")
noise_patch = (
gaussian_2d(
x,
y,
mx=int(height / 2),
my=int(width / 2),
sx=float(height / KERNEL_DIVISION),
sy=float(width / KERNEL_DIVISION),
)
).half()
return noise_patch
class RelativePosition(nn.Module):
""" https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """
@@ -110,8 +80,7 @@ class CrossAttention(nn.Module):
self.register_parameter('alpha', nn.Parameter(torch.tensor(0.)) )
def forward(self, x, context=None, mask=None, use_freetraj=False, idx_list=[], input_traj=[]):
print("forward input_traj: ", input_traj)
def forward(self, x, context=None, mask=None):
spatial_self_attn = (context is None)
k_ip, v_ip, out_ip = None, None, None
@@ -130,18 +99,9 @@ class CrossAttention(nn.Module):
context = context[:,:self.text_context_len,:]
k = self.to_k(context)
v = self.to_v(context)
hw = q.shape[0]
w_base = 64
h_base = 40
w_len = int((hw / w_base / h_base) ** 0.5 * h_base)
h_len = int(hw / w_len)
BOX_SIZE_H = input_traj[0][2] - input_traj[0][1]
BOX_SIZE_W = input_traj[0][4] - input_traj[0][3]
PATHS = plan_path(input_traj)
sub_h = int(BOX_SIZE_H * h_len)
sub_w = int(BOX_SIZE_W * w_len)
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale
if self.relative_position:
len_q, len_k, len_v = q.shape[1], k.shape[1], v.shape[1]
@@ -150,50 +110,6 @@ class CrossAttention(nn.Module):
sim += sim2
del k
if use_freetraj:
sim = rearrange(sim, '(y x h) i j -> y x h i j', h=h, y=h_len)
sim_mask = torch.zeros_like(sim)
for i in range(sim.shape[3]):
h_start1 = int(PATHS[i][0] * h_len)
h_end1 = h_start1 + sub_h
w_start1 = int(PATHS[i][2] * w_len)
w_end1 = w_start1 + sub_w
h_fg1 = list(range(h_start1, h_end1))
h_fg_tensor1 = torch.zeros(h_len, device=sim.device)
h_fg_tensor1[h_fg1] = 1
w_fg1 = list(range(w_start1, w_end1))
w_fg_tensor1 = torch.zeros(w_len, device=sim.device)
w_fg_tensor1[w_fg1] = 1
fg_tensor1 = h_fg_tensor1.view(-1, 1) * w_fg_tensor1.view(1, -1)
bg_tensor1 = 1 - fg_tensor1
for j in range(sim.shape[4]):
h_start2 = int(PATHS[j][0] * h_len)
h_end2 = h_start2 + sub_h
w_start2 = int(PATHS[j][2] * w_len)
w_end2 = w_start2 + sub_w
h_fg2 = list(range(h_start2, h_end2))
h_fg_tensor2 = torch.zeros(h_len, device=sim.device)
h_fg_tensor2[h_fg2] = 1
w_fg2 = list(range(w_start2, w_end2))
w_fg_tensor2 = torch.zeros(w_len, device=sim.device)
w_fg_tensor2[w_fg2] = 1
fg_tensor2 = h_fg_tensor2.view(-1, 1) * w_fg_tensor2.view(1, -1)
bg_tensor2 = 1 - fg_tensor2
fg_tensor = fg_tensor1 * fg_tensor2
bg_tensor = bg_tensor1 * bg_tensor2
coef = 0.01
sim_mask[:, :, :, i, j] = coef * torch.ones_like(sim_mask[:, :, :, i, j])
sim_mask[:, :, :, i, j] += (1 - coef) * torch.ones_like(sim_mask[:, :, :, i, j]) * (fg_tensor.view(h_len, w_len, 1) + bg_tensor.view(h_len, w_len, 1))
sim *= sim_mask
sim = rearrange(sim, 'y x h i j -> (y x h) i j')
del sim_mask
if exists(mask):
## feasible for causal attention mask only
max_neg_value = -torch.finfo(sim.dtype).max
@@ -202,6 +118,7 @@ class CrossAttention(nn.Module):
# attention, what we cannot get enough of
sim = sim.softmax(dim=-1)
out = torch.einsum('b i j, b j d -> b i d', sim, v)
if self.relative_position:
v2 = self.relative_position_v(len_q, len_v)
@@ -227,147 +144,6 @@ class CrossAttention(nn.Module):
out = out + self.image_cross_attention_scale * out_ip
return self.to_out(out)
def space_forward(self, x, context=None, mask=None, use_freetraj=False, idx_list=[], input_traj=[]):
if context is None:
SA_flag = True
else:
SA_flag = False
h = self.heads
q = self.to_q(x)
context = default(context, x)
## considering image token additionally
if context is not None and self.img_cross_attention:
context, context_img = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:]
k = self.to_k(context)
v = self.to_v(context)
k_ip = self.to_k_ip(context_img)
v_ip = self.to_v_ip(context_img)
else:
k = self.to_k(context)
v = self.to_v(context)
hw = q.shape[1]
w_base = 64
h_base = 40
w_len = int((hw / h_base / w_base) ** 0.5 * w_base)
h_len = int(hw / w_len)
BOX_SIZE_H = input_traj[0][2] - input_traj[0][1]
BOX_SIZE_W = input_traj[0][4] - input_traj[0][3]
PATHS = plan_path(input_traj)
sub_h = int(BOX_SIZE_H * h_len)
sub_w = int(BOX_SIZE_W * w_len)
q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v))
sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale
if self.relative_position:
len_q, len_k, len_v = q.shape[1], k.shape[1], v.shape[1]
k2 = self.relative_position_k(len_q, len_k)
sim2 = einsum('b t d, t s d -> b t s', q, k2) * self.scale # TODO check
sim += sim2
del k
if use_freetraj:
coef_a = 0.25 / (BOX_SIZE_H * BOX_SIZE_W) / len(idx_list)
weight = gaussian_weight(sub_h, sub_w).to(x.device)
if SA_flag:
weight_add = 0
sim = rearrange(sim, '(t h) (y x) (y0 x0) -> t h y x y0 x0', h=h, y=h_len, y0=h_len)
sim_mask = torch.zeros_like(sim)
for i in range(sim.shape[0]):
h_start = int(PATHS[i][0] * h_len)
h_end = h_start + sub_h
w_start = int(PATHS[i][2] * w_len)
w_end = w_start + sub_w
h_fg = list(range(h_start, h_end))
h_fg_tensor = torch.zeros(h_len, device=sim.device)
h_fg_tensor[h_fg] = 1
w_fg = list(range(w_start, w_end))
w_fg_tensor = torch.zeros(w_len, device=sim.device)
w_fg_tensor[w_fg] = 1
fg_tensor = h_fg_tensor.view(-1, 1) * w_fg_tensor.view(1, -1)
bg_tensor = 1 - fg_tensor
coef = 0.01
sim_mask[i] = coef * torch.ones_like(sim_mask[i])
sim_mask[i] += (1-coef) * (torch.ones_like(sim_mask[i]) * fg_tensor.view(1, h_len, w_len, 1, 1) * fg_tensor.view(1, 1, 1, h_len, w_len) + torch.ones_like(sim_mask[i]) * bg_tensor.view(1, h_len, w_len, 1, 1) * bg_tensor.view(1, 1, 1, h_len, w_len))
sim *= sim_mask
sim = rearrange(sim, 't h y x y0 x0 -> (t h) (y x) (y0 x0)')
else:
sim = rearrange(sim, '(t h) (y x) d -> t h y x d', h=h, y=h_len)
sim_mask = torch.zeros_like(sim)
weight_add = torch.zeros_like(sim)
weight_map = torch.zeros([sim.shape[0], h_len, w_len], device=sim.device)
for i in range(sim.shape[0]):
h_start = int(PATHS[i][0] * h_len)
h_end = h_start + sub_h
w_start = int(PATHS[i][2] * w_len)
w_end = w_start + sub_w
h_fg = list(range(h_start, h_end))
h_fg_tensor = torch.zeros(h_len, device=sim.device)
h_fg_tensor[h_fg] = 1
w_fg = list(range(w_start, w_end))
w_fg_tensor = torch.zeros(w_len, device=sim.device)
w_fg_tensor[w_fg] = 1
fg_tensor = h_fg_tensor.view(-1, 1) * w_fg_tensor.view(1, -1)
bg_tensor = 1 - fg_tensor
if idx_list == []:
p_fg = [2]
else:
p_fg = idx_list
p_bg = list(range(77))
for j in p_fg:
p_bg.remove(j)
weight_map[i, h_start:h_end, w_start:w_end] = weight * coef_a
sim_mask[i, :, :, :, p_bg] = torch.ones_like(sim_mask[i, :, :, :, p_bg]) * bg_tensor.view(1, h_len, w_len, 1)
weight_add[i, :, :, :, p_fg] = torch.ones_like(sim_mask[i, :, :, :, p_fg]) * weight_map[i].view(1, h_len, w_len, 1)
max_neg_value = -torch.finfo(sim.dtype).max
sim.masked_fill_(~(sim_mask>0.5), max_neg_value)
sim = rearrange(sim, 't h y x d -> (t h) (y x) d')
weight_add = rearrange(weight_add, 't h y x d -> (t h) (y x) d')
del sim_mask
if exists(mask):
## feasible for causal attention mask only
max_neg_value = -torch.finfo(sim.dtype).max
mask = repeat(mask, 'b i j -> (b h) i j', h=h)
sim.masked_fill_(~(mask>0.5), max_neg_value)
# attention, what we cannot get enough of
sim = sim.softmax(dim=-1)
if use_freetraj:
sim += weight_add
out = torch.einsum('b i j, b j d -> b i d', sim, v)
if self.relative_position:
v2 = self.relative_position_v(len_q, len_v)
out2 = einsum('b t s, t s d -> b t d', sim, v2) # TODO check
out += out2
out = rearrange(out, '(b h) n d -> b n (h d)', h=h)
## considering image token additionally
if context is not None and self.img_cross_attention:
k_ip, v_ip = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (k_ip, v_ip))
sim_ip = torch.einsum('b i d, b j d -> b i j', q, k_ip) * self.scale
del k_ip
sim_ip = sim_ip.softmax(dim=-1)
out_ip = torch.einsum('b i j, b j d -> b i d', sim_ip, v_ip)
out_ip = rearrange(out_ip, '(b h) n d -> b n (h d)', h=h)
out = out + self.image_cross_attention_scale * out_ip
del q
return self.to_out(out)
def efficient_forward(self, x, context=None, mask=None):
spatial_self_attn = (context is None)
@@ -454,7 +230,7 @@ class BasicTransformerBlock(nn.Module):
self.checkpoint = checkpoint
def forward(self, x, context=None, mask=None, use_freetraj=False, idx_list=[], input_traj=[], **kwargs):
def forward(self, x, context=None, mask=None, **kwargs):
## implementation tricks: because checkpointing doesn't support non-tensor (e.g. None or scalar) arguments
input_tuple = (x,) ## should not be (x), otherwise *input_tuple will decouple x into multiple arguments
if context is not None:
@@ -462,14 +238,12 @@ class BasicTransformerBlock(nn.Module):
if mask is not None:
forward_mask = partial(self._forward, mask=mask)
return checkpoint(forward_mask, (x,), self.parameters(), self.checkpoint)
if context is not None and mask is not None:
input_tuple = (x, context, mask)
input_tuple = (x, context, mask, use_freetraj, idx_list, input_traj)
return checkpoint(self._forward, input_tuple, self.parameters(), self.checkpoint)
def _forward(self, x, context=None, mask=None, use_freetraj=False, idx_list=[], input_traj=[]):
x = self.attn1(self.norm1(x), context=context if self.disable_self_attn else None, mask=mask, use_freetraj=use_freetraj, idx_list=idx_list, input_traj=input_traj) + x
x = self.attn2(self.norm2(x), context=context, mask=mask, use_freetraj=use_freetraj, idx_list=idx_list, input_traj=input_traj) + x
def _forward(self, x, context=None, mask=None):
x = self.attn1(self.norm1(x), context=context if self.disable_self_attn else None, mask=mask) + x
x = self.attn2(self.norm2(x), context=context, mask=mask) + x
x = self.ff(self.norm3(x)) + x
return x
@@ -591,7 +365,7 @@ class TemporalTransformer(nn.Module):
self.use_linear = use_linear
def forward(self, x_in, context=None, frame_window_size=None, frame_window_stride=None, **kwargs):
def forward(self, x_in, context=None, frame_window_size=None, frame_window_stride=None):
B, C, T, H, W = x_in.shape
def process_slice(x, t_start=None, t_end=None):
b, c, t, h, w = x.shape
@@ -620,7 +394,7 @@ class TemporalTransformer(nn.Module):
if self.only_self_att:
## note: if no context is given, cross-attention defaults to self-attention
for i, block in enumerate(self.transformer_blocks):
x = block(x, mask=mask, **kwargs)
x = block(x, mask=mask)
x = rearrange(x, "(b hw) t c -> b hw t c", b=b).contiguous()
else:
x = rearrange(x, "(b hw) t c -> b hw t c", b=b).contiguous()
@@ -632,7 +406,7 @@ class TemporalTransformer(nn.Module):
context[j], "t l con -> (t r) l con", r=(h * w) // t, t=t
).contiguous()
## note: causal mask will not applied in cross-attention case
x[j] = block(x[j], context=context_j, **kwargs)
x[j] = block(x[j], context=context_j)
if self.use_linear:
x = self.proj_out(x)
+412 -10
View File
@@ -2,7 +2,9 @@ from functools import partial
from abc import abstractmethod
import torch
import torch.nn as nn
import numpy as np
from einops import rearrange
import math
import torch.nn.functional as F
from ....lvdm.models.utils_diffusion import timestep_embedding
from ....lvdm.common import checkpoint
@@ -18,6 +20,9 @@ from ....lvdm.modules.attention import SpatialTransformer, TemporalTransformer
import comfy.ops
ops = comfy.ops.manual_cast
def exists(x):
return x is not None
class TimestepBlock(nn.Module):
"""
Any module where forward() takes timestep embeddings as a second argument.
@@ -35,15 +40,15 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
support it as an extra input.
"""
def forward(self, x, emb, context=None, batch_size=None, use_freetraj=False, frame_window_size=None, frame_window_stride=None, **kwargs):
def forward(self, x, emb, context=None, batch_size=None, frame_window_size=None, frame_window_stride=None):
for layer in self:
if isinstance(layer, TimestepBlock):
x = layer(x, emb, batch_size=batch_size)
elif isinstance(layer, SpatialTransformer):
x = layer(x, context, use_freetraj=use_freetraj, **kwargs)
x = layer(x, context)
elif isinstance(layer, TemporalTransformer):
x = rearrange(x, '(b f) c h w -> b c f h w', b=batch_size)
x = layer(x, context, use_freetraj=use_freetraj, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
x = layer(x, context, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
x = rearrange(x, 'b c f h w -> (b f) c h w')
else:
x = layer(x)
@@ -547,8 +552,7 @@ class UNetModel(nn.Module):
zero_module(conv_nd(dims, model_channels, out_channels, 3, padding=1)),
)
def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, **kwargs):
print("unet forward: ", kwargs)
def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, control=None, **kwargs):
b,_,t,_,_ = x.shape
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype)
emb = self.time_embed(t_emb)
@@ -583,9 +587,9 @@ class UNetModel(nn.Module):
adapter_idx = 0
hs = []
for id, module in enumerate(self.input_blocks):
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
if id ==0 and self.addition_attention:
h = self.init_attn(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
h = self.init_attn(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
## plug-in adapter features
if ((id+1)%3 == 0) and features_adapter is not None:
h = h + features_adapter[adapter_idx]
@@ -594,13 +598,411 @@ class UNetModel(nn.Module):
if features_adapter is not None:
assert len(features_adapter)==adapter_idx, 'Wrong features_adapter'
h = self.middle_block(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
h = self.middle_block(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
if control is not None:
h += control.pop()
for module in self.output_blocks:
h = torch.cat([h, hs.pop()], dim=1)
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
if control is None:
h = torch.cat([h, hs.pop()], dim=1)
else:
h = torch.cat([h, hs.pop() + control.pop()], dim=1)
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
h = h.type(x.dtype)
y = self.out(h)
# reshape back to (b c t h w)
y = rearrange(y, '(b t) c h w -> b c t h w', b=b)
return y
class ControlNet(nn.Module):
def __init__(
self,
image_size,
in_channels,
model_channels,
hint_channels,
num_res_blocks,
attention_resolutions,
dropout=0,
channel_mult=(1, 2, 4, 8),
conv_resample=True,
dims=2,
use_checkpoint=False,
use_fp16=False,
num_heads=-1,
num_head_channels=-1,
num_heads_upsample=-1,
use_scale_shift_norm=False,
resblock_updown=False,
use_new_attention_order=False,
use_spatial_transformer=False, # custom transformer support
transformer_depth=1, # custom transformer support
context_dim=None, # custom transformer support
n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model
legacy=True,
disable_self_attentions=None,
num_attention_blocks=None,
disable_middle_self_attn=False,
use_linear_in_transformer=False,
):
super().__init__()
if use_spatial_transformer:
assert context_dim is not None, 'Fool!! You forgot to include the dimension of your cross-attention conditioning...'
if context_dim is not None:
assert use_spatial_transformer, 'Fool!! You forgot to use the spatial transformer for your cross-attention conditioning...'
from omegaconf.listconfig import ListConfig
if type(context_dim) == ListConfig:
context_dim = list(context_dim)
if num_heads_upsample == -1:
num_heads_upsample = num_heads
if num_heads == -1:
assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set'
if num_head_channels == -1:
assert num_heads != -1, 'Either num_heads or num_head_channels has to be set'
self.dims = dims
self.image_size = image_size
self.in_channels = in_channels
self.model_channels = model_channels
if isinstance(num_res_blocks, int):
self.num_res_blocks = len(channel_mult) * [num_res_blocks]
else:
if len(num_res_blocks) != len(channel_mult):
raise ValueError("provide num_res_blocks either as an int (globally constant) or "
"as a list/tuple (per-level) with the same length as channel_mult")
self.num_res_blocks = num_res_blocks
if disable_self_attentions is not None:
# should be a list of booleans, indicating whether to disable self-attention in TransformerBlocks or not
assert len(disable_self_attentions) == len(channel_mult)
if num_attention_blocks is not None:
assert len(num_attention_blocks) == len(self.num_res_blocks)
assert all(map(lambda i: self.num_res_blocks[i] >= num_attention_blocks[i], range(len(num_attention_blocks))))
print(f"Constructor of UNetModel received num_attention_blocks={num_attention_blocks}. "
f"This option has LESS priority than attention_resolutions {attention_resolutions}, "
f"i.e., in cases where num_attention_blocks[i] > 0 but 2**i not in attention_resolutions, "
f"attention will still not be set.")
self.attention_resolutions = attention_resolutions
self.dropout = dropout
self.channel_mult = channel_mult
self.conv_resample = conv_resample
self.use_checkpoint = use_checkpoint
self.dtype = torch.float16 if use_fp16 else torch.float32
self.num_heads = num_heads
self.num_head_channels = num_head_channels
self.num_heads_upsample = num_heads_upsample
self.predict_codebook_ids = n_embed is not None
time_embed_dim = model_channels * 4
self.time_embed = nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
)
self.input_blocks = nn.ModuleList(
[
TimestepEmbedSequential(
conv_nd(dims, in_channels, model_channels, 3, padding=1)
)
]
)
self.zero_convs = nn.ModuleList([self.make_zero_conv(model_channels)])
self.input_hint_block = TimestepEmbedSequential(
conv_nd(dims, hint_channels, 16, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 16, 16, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 16, 32, 3, padding=1, stride=2),
nn.SiLU(),
conv_nd(dims, 32, 32, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 32, 96, 3, padding=1, stride=2),
nn.SiLU(),
conv_nd(dims, 96, 96, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 96, 256, 3, padding=1, stride=2),
nn.SiLU(),
zero_module(conv_nd(dims, 256, model_channels, 3, padding=1))
)
self._feature_size = model_channels
input_block_chans = [model_channels]
ch = model_channels
ds = 1
for level, mult in enumerate(channel_mult):
for nr in range(self.num_res_blocks[level]):
layers = [
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=mult * model_channels,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = mult * model_channels
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
if legacy:
# num_heads = 1
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
if exists(disable_self_attentions):
disabled_sa = disable_self_attentions[level]
else:
disabled_sa = False
if not exists(num_attention_blocks) or nr < num_attention_blocks[level]:
layers.append(
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=dim_head,
use_new_attention_order=use_new_attention_order,
) if not use_spatial_transformer else SpatialTransformer(
ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim,
disable_self_attn=disabled_sa, use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint
)
)
self.input_blocks.append(TimestepEmbedSequential(*layers))
self.zero_convs.append(self.make_zero_conv(ch))
self._feature_size += ch
input_block_chans.append(ch)
if level != len(channel_mult) - 1:
out_ch = ch
self.input_blocks.append(
TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
down=True,
)
if resblock_updown
else Downsample(
ch, conv_resample, dims=dims, out_channels=out_ch
)
)
)
ch = out_ch
input_block_chans.append(ch)
self.zero_convs.append(self.make_zero_conv(ch))
ds *= 2
self._feature_size += ch
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
if legacy:
# num_heads = 1
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
self.middle_block = TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=dim_head,
use_new_attention_order=use_new_attention_order,
) if not use_spatial_transformer else SpatialTransformer( # always uses a self-attn
ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim,
disable_self_attn=disable_middle_self_attn, use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint
),
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
)
self.middle_block_out = self.make_zero_conv(ch)
self._feature_size += ch
def make_zero_conv(self, channels):
return TimestepEmbedSequential(zero_module(conv_nd(self.dims, channels, channels, 1, padding=0)))
def forward(self, x, hint, timesteps, context, **kwargs):
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False)
emb = self.time_embed(t_emb)
guided_hint = self.input_hint_block(hint, emb, context)
outs = []
h = x.type(self.dtype)
for module, zero_conv in zip(self.input_blocks, self.zero_convs):
if guided_hint is not None:
h = module(h, emb, context)
h += guided_hint
guided_hint = None
else:
h = module(h, emb, context)
outs.append(zero_conv(h, emb, context, True))
h = self.middle_block(h, emb, context)
outs.append(self.middle_block_out(h, emb, context))
return outs
class AttentionBlock(nn.Module):
"""
An attention block that allows spatial positions to attend to each other.
Originally ported from here, but adapted to the N-d case.
https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/models/unet.py#L66.
"""
def __init__(
self,
channels,
num_heads=1,
num_head_channels=-1,
use_checkpoint=False,
use_new_attention_order=False,
):
super().__init__()
self.channels = channels
if num_head_channels == -1:
self.num_heads = num_heads
else:
assert (
channels % num_head_channels == 0
), f"q,k,v channels {channels} is not divisible by num_head_channels {num_head_channels}"
self.num_heads = channels // num_head_channels
self.use_checkpoint = use_checkpoint
self.norm = normalization(channels)
self.qkv = conv_nd(1, channels, channels * 3, 1)
if use_new_attention_order:
# split qkv before split heads
self.attention = QKVAttention(self.num_heads)
else:
# split heads before split qkv
self.attention = QKVAttentionLegacy(self.num_heads)
self.proj_out = zero_module(conv_nd(1, channels, channels, 1))
def forward(self, x):
return checkpoint(self._forward, (x,), self.parameters(), True) # TODO: check checkpoint usage, is True # TODO: fix the .half call!!!
#return pt_checkpoint(self._forward, x) # pytorch
def _forward(self, x):
b, c, *spatial = x.shape
x = x.reshape(b, c, -1)
qkv = self.qkv(self.norm(x))
h = self.attention(qkv)
h = self.proj_out(h)
return (x + h).reshape(b, c, *spatial)
class QKVAttention(nn.Module):
"""
A module which performs QKV attention and splits in a different order.
"""
def __init__(self, n_heads):
super().__init__()
self.n_heads = n_heads
def forward(self, qkv):
"""
Apply QKV attention.
:param qkv: an [N x (3 * H * C) x T] tensor of Qs, Ks, and Vs.
:return: an [N x (H * C) x T] tensor after attention.
"""
bs, width, length = qkv.shape
assert width % (3 * self.n_heads) == 0
ch = width // (3 * self.n_heads)
q, k, v = qkv.chunk(3, dim=1)
scale = 1 / math.sqrt(math.sqrt(ch))
weight = torch.einsum(
"bct,bcs->bts",
(q * scale).view(bs * self.n_heads, ch, length),
(k * scale).view(bs * self.n_heads, ch, length),
) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
a = torch.einsum("bts,bcs->bct", weight, v.reshape(bs * self.n_heads, ch, length))
return a.reshape(bs, -1, length)
@staticmethod
def count_flops(model, _x, y):
return count_flops_attn(model, _x, y)
class QKVAttentionLegacy(nn.Module):
"""
A module which performs QKV attention. Matches legacy QKVAttention + input/ouput heads shaping
"""
def __init__(self, n_heads):
super().__init__()
self.n_heads = n_heads
def forward(self, qkv):
"""
Apply QKV attention.
:param qkv: an [N x (H * 3 * C) x T] tensor of Qs, Ks, and Vs.
:return: an [N x (H * C) x T] tensor after attention.
"""
bs, width, length = qkv.shape
assert width % (3 * self.n_heads) == 0
ch = width // (3 * self.n_heads)
q, k, v = qkv.reshape(bs * self.n_heads, ch * 3, length).split(ch, dim=1)
scale = 1 / math.sqrt(math.sqrt(ch))
weight = torch.einsum(
"bct,bcs->bts", q * scale, k * scale
) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
a = torch.einsum("bts,bcs->bct", weight, v)
return a.reshape(bs, -1, length)
@staticmethod
def count_flops(model, _x, y):
return count_flops_attn(model, _x, y)
def count_flops_attn(model, _x, y):
"""
A counter for the `thop` package to count the operations in an
attention operation.
Meant to be used like:
macs, params = thop.profile(
model,
inputs=(inputs, timestamps),
custom_ops={QKVAttention: QKVAttention.count_flops},
)
"""
b, c, *spatial = y[0].shape
num_spatial = int(np.prod(spatial))
# We perform two matmuls with the same number of ops.
# The first computes the weight matrix, the second computes
# the combination of the value vectors.
matmul_ops = 2 * b * (num_spatial ** 2) * c
model.total_ops += torch.DoubleTensor([matmul_ops])
+229 -328
View File
@@ -4,15 +4,13 @@ import torch
import torch.nn.functional as F
from .scripts.evaluation.funcs import load_model_checkpoint, get_latent_z, get_latent_z_with_hidden_states
from .utils.utils import instantiate_from_config
from .utils.utils_freetraj import get_freq_filter, freq_mix_3d, get_path, plan_path, load_idx, load_traj
from einops import repeat
import folder_paths
import comfy.model_management as mm
import comfy.utils
from contextlib import nullcontext
from .lvdm.models.samplers.ddim import DDIMSampler
from .lvdm.modules.networks.openaimodel3d import ControlNet
from contextlib import nullcontext
try:
@@ -143,8 +141,130 @@ class DownloadAndLoadDynamiCrafterModel:
dcmodel = {
'model': self.model,
'model_name': model,
'dtype': precision
}
return (dcmodel,)
class DownloadAndLoadDynamiCrafterCNModel:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": (
[
'sketch_encoder-fp16.safetensors',
],
),
},
}
RETURN_TYPES = ("DC_CN_MODEL",)
RETURN_NAMES = ("DynCraft_CN_model",)
FUNCTION = "loadmodel"
CATEGORY = "DynamiCrafterWrapper"
def loadmodel(self, model):
custom_config = {
'ckpt_name': model,
}
if not hasattr(self, 'cn_model') or self.cn_model == None or custom_config != self.current_config:
download_path = os.path.join(folder_paths.models_dir, "checkpoints", "dynamicrafter", "controlnet")
cn_model_path = os.path.join(download_path, model)
if not os.path.exists(cn_model_path):
print(f"Downloading model to: {cn_model_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="Kijai/DynamiCrafter_pruned",
allow_patterns=[f"*{model}*"],
local_dir=download_path,
local_dir_use_symlinks=False)
cn_config = {
"use_checkpoint": True,
"image_size": 32, # unused
"in_channels": 4,
"hint_channels": 3,
"model_channels": 320,
"attention_resolutions": [4, 2, 1],
"num_res_blocks": 2,
"channel_mult": [1, 2, 4, 4],
"num_head_channels": 64, # need to fix for flash-attn
"use_spatial_transformer": True,
"use_linear_in_transformer": True,
"transformer_depth": 1,
"context_dim": 1024,
"legacy": False
}
if "sketch_encoder" in model:
cn_config["hint_channels"] = 1
self.cn_model = ControlNet(**cn_config)
print("Loading ControlNet")
cn_sd = comfy.utils.load_torch_file(cn_model_path)
self.cn_model.load_state_dict(cn_sd, strict=True)
print("ControlNet loaded")
controlnet = {
'model': self.cn_model,
'config': cn_config,
}
return (controlnet,)
class DynamiCrafterCNLoader:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"ckpt_name": (folder_paths.get_filename_list("controlnet"), ),
},
}
RETURN_TYPES = ("DC_CN_MODEL",)
RETURN_NAMES = ("DynCraft_CN_model",)
FUNCTION = "loadmodel"
CATEGORY = "DynamiCrafterWrapper"
def loadmodel(self, ckpt_name):
custom_config = {
'ckpt_name': ckpt_name,
}
if not hasattr(self, 'cn_model') or self.cn_model == None or custom_config != self.current_config:
self.current_config = custom_config
model_path = folder_paths.get_full_path("controlnet", ckpt_name)
print(f"Loading ControlNet from: {model_path}")
cn_config = {
"use_checkpoint": True,
"image_size": 32, # unused
"in_channels": 4,
"hint_channels": 3,
"model_channels": 320,
"attention_resolutions": [4, 2, 1],
"num_res_blocks": 2,
"channel_mult": [1, 2, 4, 4],
"num_head_channels": 64, # need to fix for flash-attn
"use_spatial_transformer": True,
"use_linear_in_transformer": True,
"transformer_depth": 1,
"context_dim": 1024,
"legacy": False
}
if "sketch_encoder" in ckpt_name:
cn_config["hint_channels"] = 1
self.cn_model = ControlNet(**cn_config)
print("Loading ControlNet")
cn_sd = comfy.utils.load_torch_file(model_path)
self.cn_model.load_state_dict(cn_sd, strict=True)
del cn_sd
print("ControlNet loaded")
controlnet = {
'model': self.cn_model,
'config': cn_config,
}
return (controlnet,)
class DownloadAndLoadCLIPModel:
@classmethod
@@ -392,6 +512,7 @@ class DynamiCrafterI2V:
self.model.to(device)
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
image = image.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image += torch.randn_like(image) * augmentation_level
@@ -403,7 +524,7 @@ class DynamiCrafterI2V:
if H % 64 != 0:
H = H - (H % 64)
if orig_H % 64 != 0 or orig_W % 64 != 0:
image = F.interpolate(image, size=(H, W), mode="bicubic")
image = F.interpolate(image, size=(H, W), mode="bilinear")
B, C, H, W = image.shape
noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8]
@@ -413,14 +534,13 @@ class DynamiCrafterI2V:
z = get_latent_z(self.model, encode_pixels) #bc,1,hw
if image2 is not None:
image2 = image2 * 2 - 1
image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image2 += torch.randn_like(image2) * augmentation_level
if image2.shape != image.shape:
image2 = F.interpolate(image, size=(H, W), mode="bicubic")
image2 = F.interpolate(image, size=(H, W), mode="bilinear")
encode_pixels = image2.unsqueeze(2) * 2 - 1
z2 = get_latent_z(self.model, encode_pixels) #bc,1,hw
@@ -437,6 +557,8 @@ class DynamiCrafterI2V:
text_emb = positive[0][0].to(device)
cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))['last_hidden_state'].to(device)
cond_images = torch.sum(cond_images, dim=0).unsqueeze(0)
cond_images = torch.mean(cond_images, dim=0).unsqueeze(0)
img_emb = self.model.image_proj_model(cond_images)
@@ -444,7 +566,9 @@ class DynamiCrafterI2V:
del cond_images, img_emb, text_emb, encode_pixels
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]}
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat], "control_cond": None}
self.model.control_model = None
if noise_shape[-1] == 32:
timestep_spacing = "uniform"
@@ -636,6 +760,7 @@ class ToonCrafterInterpolation:
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"optional_latents": ("LATENT",),
"ddpm_from": ("INT", {"default": 1000, "min": 1, "max": 1000, "step": 1}),
"controlnet": ("DC_CONTROL",),
}
}
@@ -644,7 +769,8 @@ class ToonCrafterInterpolation:
FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper"
def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0, augmentation_level=0, optional_latents=None, ddpm_from=1000):
def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames,
vae_dtype, image_embed_ratio=1.0, augmentation_level=0, optional_latents=None, ddpm_from=1000, controlnet=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
@@ -654,6 +780,11 @@ class ToonCrafterInterpolation:
self.model = model['model']
if controlnet is not None:
self.model.control_model = controlnet["model"]
else:
self.model.control_model = None
dtype = self.model.dtype
if vae_dtype == "auto":
try:
@@ -685,11 +816,12 @@ class ToonCrafterInterpolation:
pbar = comfy.utils.ProgressBar(len(images) - 1)
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
for i in range(len(images) - 1):
for i in range(len(images) - 1) if len(images) > 1 else range(len(images)):
videos, videos2 = None, None
mm.soft_empty_cache()
image = images[i].unsqueeze(0)
image2 = images[i+1].unsqueeze(0)
if len(images) !=1:
image2 = images[i+1].unsqueeze(0)
B, C, H, W = image.shape
noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8]
@@ -701,12 +833,16 @@ class ToonCrafterInterpolation:
image2 += torch.randn_like(image) * augmentation_level
encode_pixels = image.unsqueeze(2) * 2 - 1
videos = encode_pixels # bc1hw
videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames//2)
encode_pixels = image2.unsqueeze(2) * 2 - 1
videos2 = encode_pixels # bc1hw
videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames//2)
videos = torch.cat([videos, videos2], dim=2)
videos = encode_pixels # bc1hw
videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames // 2)
if len(images) == 1:
videos = torch.cat([videos, videos], dim=2)
else:
encode_pixels = image2.unsqueeze(2) * 2 - 1
videos2 = encode_pixels # bc1hw
videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames // 2)
videos = torch.cat([videos, videos2], dim=2)
try:
z, hs = get_latent_z_with_hidden_states(self.model, videos)
@@ -718,29 +854,38 @@ class ToonCrafterInterpolation:
img_tensor_repeat = torch.zeros_like(z)
img_tensor_repeat[:,:,:1,:,:] = z[:,:,:1,:,:]
img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:]
if len(images) !=1:
img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:]
self.model.first_stage_model.to(offload_device)
text_emb = positive[0][0].to(device)
cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
self.model.image_proj_model.to(device)
cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
img_emb = self.model.image_proj_model(cond_images)
img_emb2 = self.model.image_proj_model(cond_images2)
img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio)
if len(images) !=1:
cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
img_emb2 = self.model.image_proj_model(cond_images2)
img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio)
else:
img_embeds = img_emb
imtext_cond = torch.cat([text_emb, img_embeds], dim=1)
del cond_images, img_emb, img_emb2, text_emb
del cond_images, img_emb, text_emb
if comfy.model_management.is_device_mps(device):
fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device)
else:
fs = torch.tensor([fs], dtype=torch.float64, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]}
if controlnet is not None:
cn_videos = controlnet["cn_videos"]
cn_videos = cn_videos.to(dtype).to(device)
else:
cn_videos = None
cond = {"c_crossattn": [imtext_cond], "fs": fs, "c_concat": [img_tensor_repeat], "control_cond": cn_videos}
if noise_shape[-1] == 32:
timestep_spacing = "uniform"
@@ -822,6 +967,41 @@ class ToonCrafterInterpolation:
}
return (latent,)
class DynamiCrafterControlnetApply:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"controlnet": ("DC_CN_MODEL",),
"images": ("IMAGE",),
"control_scale": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("DC_CONTROL",)
RETURN_NAMES = ("controlnet",)
FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper"
def process(self, controlnet, images, control_scale):
controlnet['model'].control_scale = control_scale
#images = images * 2.0 - 1.0
cn_tensor = images.permute(3, 0, 1, 2).unsqueeze(0)
print("control frame: ", cn_tensor.shape) # b c t h w
print(controlnet["config"])
if controlnet["config"]["hint_channels"] == 1:
cn_tensor = cn_tensor[:, :1, :, :, :]
print("control frame: ", cn_tensor.shape) # b c t h w
controlnet = {
"model": controlnet['model'],
"cn_videos": cn_tensor,
}
return (controlnet,)
class ToonCrafterDecode:
@classmethod
@@ -955,8 +1135,9 @@ class DynamiCrafterBatchInterpolation:
mm.soft_empty_cache()
torch.manual_seed(seed)
dtype = model.dtype
self.model = model['model']
dtype = self.model.dtype
if vae_dtype == "auto":
try:
@@ -971,7 +1152,6 @@ class DynamiCrafterBatchInterpolation:
print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
self.model.to(device)
images = images * 2 - 1
images = images.permute(0, 3, 1, 2).to(dtype).to(device)
B, C, H, W = images.shape
orig_H, orig_W = H, W
@@ -994,8 +1174,11 @@ class DynamiCrafterBatchInterpolation:
self.model.first_stage_model.to(device)
z = get_latent_z(self.model, image.unsqueeze(2)) #bc,1,hw
z2 = get_latent_z(self.model, image2.unsqueeze(2)) #bc,1,hw
encode_pixels1 = image * 2 - 1
encode_pixels2 = image2 * 2 - 1
z = get_latent_z(self.model, encode_pixels1.unsqueeze(2)) #bc,1,hw
z2 = get_latent_z(self.model, encode_pixels2.unsqueeze(2)) #bc,1,hw
img_tensor_repeat = repeat(z, 'b c t h w -> b c (repeat t) h w', repeat=frames)
img_tensor_repeat = torch.zeros_like(img_tensor_repeat)
img_tensor_repeat[:,:,:1,:,:] = z
@@ -1012,7 +1195,9 @@ class DynamiCrafterBatchInterpolation:
imtext_cond = torch.cat([text_emb, img_emb], dim=1)
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]}
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat], "control_cond": None}
self.model.control_model = None
if noise_shape[-1] == 32:
timestep_spacing = "uniform"
@@ -1111,294 +1296,6 @@ class DynamiCrafterBatchInterpolation:
last_image = out_video[-1].unsqueeze(0)
return (out_video, last_image)
class DynamiCrafterFreeTrajSampler:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("DCMODEL",),
"clip_vision": ("CLIP_VISION", ),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"image": ("IMAGE",),
"steps": ("INT", {"default": 50, "min": 1, "max": 200, "step": 1}),
"cfg": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 20.0, "step": 0.01}),
"eta": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"frames": ("INT", {"default": 16, "min": 1, "max": 100, "step": 1}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"fs": ("INT", {"default": 10, "min": 2, "max": 100, "step": 1}),
"n_samples": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"ddim_edit": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"keep_model_loaded": ("BOOLEAN", {"default": True}),
"vae_dtype": (
[
'fp32',
'fp16',
'bf16',
'auto'
], {
"default": 'auto'
}),
},
"optional": {
"image2": ("IMAGE",),
"mask": ("MASK",),
"frame_window_size": ("INT", {"default": 16, "min": 1, "max": 200, "step": 1}),
"frame_window_stride": ("INT", {"default": 4, "min": 1, "max": 200, "step": 1}),
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"init_noise": ("DCNOISE",),
}
}
RETURN_TYPES = ("IMAGE", "IMAGE",)
RETURN_NAMES = ("images", "last_image",)
FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper"
def process(self, model, image, clip_vision, positive, negative, cfg, steps, eta, seed, fs, n_samples, keep_model_loaded, ddim_edit,
frames, vae_dtype, frame_window_size=16, frame_window_stride=4, mask=None, image2=None, augmentation_level=0, init_noise=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
mm.soft_empty_cache()
self.model = model['model']
channels = self.model.model.diffusion_model.out_channels
torch.manual_seed(seed)
dtype = self.model.dtype
if vae_dtype == "auto":
try:
if mm.should_use_bf16():
self.model.first_stage_model.to(convert_dtype('bf16'))
else:
self.model.first_stage_model.to(convert_dtype('fp32'))
except:
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtype manually.")
else:
self.model.first_stage_model.to(convert_dtype(vae_dtype))
print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
self.model.to(device)
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
image = image.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image += torch.randn_like(image) * augmentation_level
B, C, H, W = image.shape
orig_H, orig_W = H, W
if W % 64 != 0:
W = W - (W % 64)
if H % 64 != 0:
H = H - (H % 64)
if orig_H % 64 != 0 or orig_W % 64 != 0:
image = F.interpolate(image, size=(H, W), mode="bicubic")
B, C, H, W = image.shape
noise_h, noise_w = H // 8, W // 8
noise_shape = [B, channels, frames, noise_h, noise_w]
self.model.first_stage_model.to(device)
encode_pixels = image.unsqueeze(2) * 2 - 1
z = get_latent_z(self.model, encode_pixels) #bc,1,hw
if image2 is not None:
image2 = image2 * 2 - 1
image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image2 += torch.randn_like(image2) * augmentation_level
if image2.shape != image.shape:
image2 = F.interpolate(image, size=(H, W), mode="bicubic")
encode_pixels = image2.unsqueeze(2) * 2 - 1
z2 = get_latent_z(self.model, encode_pixels) #bc,1,hw
img_tensor_repeat = repeat(z, 'b c t h w -> b c (repeat t) h w', repeat=frames)
img_tensor_repeat = torch.zeros_like(img_tensor_repeat)
img_tensor_repeat[:,:,:1,:,:] = z
img_tensor_repeat[:,:,-1:,:,:] = z2
else:
img_tensor_repeat = repeat(z, 'b c t h w -> b c (repeat t) h w', repeat=frames)
self.model.first_stage_model.to(offload_device)
self.model.image_proj_model.to(device)
text_emb = positive[0][0].to(device)
cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))['last_hidden_state'].to(device)
img_emb = self.model.image_proj_model(cond_images)
imtext_cond = torch.cat([text_emb, img_emb], dim=1)
del cond_images, img_emb, text_emb, encode_pixels
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]}
if noise_shape[-1] == 32:
timestep_spacing = "uniform"
guidance_rescale = 0.0
else:
timestep_spacing = "uniform_trailing"
guidance_rescale = 0.7
## construct unconditional guidance
if cfg != 1.0:
uc_emb = negative[0][0].to(device)
## process image embedding token
if hasattr(self.model, 'embedder'):
uc_img = torch.rand(noise_shape[0],3,224,224).to(self.model.device)
## img: b c h w >> b l c
uc_img = clip_vision.encode_image(uc_img.permute(0, 2, 3, 1))['last_hidden_state'].to(self.model.device)
uc_img = self.model.image_proj_model(uc_img)
uc_emb = torch.cat([uc_emb, uc_img], dim=1)
if isinstance(cond, dict):
uc = {key:cond[key] for key in cond.keys()}
uc.update({'c_crossattn': [uc_emb]})
else:
uc = uc_emb
else:
uc = None
self.model.image_proj_model.to(offload_device)
if mask is not None:
mask = mask.to(dtype).to(device)
mask = F.interpolate(mask.unsqueeze(0), size=(H // 8, W // 8), mode="nearest").squeeze(0)
mask = (1 - mask)
mask = mask.unsqueeze(1)
B, C, H, W = mask.shape
if B < frames:
mask = mask.unsqueeze(2)
mask = mask.expand(-1, -1, frames, -1, -1)
else:
mask = mask.unsqueeze(0)
mask = mask.permute(0, 2, 1, 3, 4)
mask = torch.where(mask < 1.0, torch.tensor(0.0, device=device, dtype=dtype), torch.tensor(1.0, device=device, dtype=dtype))
if init_noise is not None:
if init_noise['analytic_init']:
eps=torch.randn_like(init_noise['mu_p'])
sigma_p = init_noise['sigma_p']
init = (init_noise['mu_p'] + sigma_p*eps).to(dtype).to(device)
if noise_shape[2] % init.shape[2] == 0:
init = init.repeat(1, 1, noise_shape[2] // init.shape[2], 1, 1)
else:
raise ValueError("The target dimension size is not an integral multiple of the original dimension size.")
else:
init = None
timestep_spacing = "uniform_trailing"
guidance_rescale = 0.7
ddpm_from = init_noise['M']
else:
init = None
ddpm_from = 1000
total_shape = [n_samples, 1, channels, frames, noise_h, noise_w]
print('total_shape', total_shape)
x_T_total = None
input_traj = [[0, 0, 0.3, 0, 0.3], [15, 0.7, 1, 0.7, 14]]
idx_list = [1, 2]
if x_T_total is None:
x_T_total = torch.randn(total_shape, device=device).repeat(1, B, 1, 1, 1, 1)
noise_flow = True
if noise_flow:
print('noise_flow')
BOX_SIZE_H = input_traj[0][2] - input_traj[0][1]
BOX_SIZE_W = input_traj[0][4] - input_traj[0][3]
print("box_size_h: ", BOX_SIZE_H, "box_size_h: ", BOX_SIZE_W)
PATHS = plan_path(input_traj)
print("PATHS: ",PATHS)
sub_h = int(BOX_SIZE_H * noise_h)
sub_w = int(BOX_SIZE_W * noise_w)
x_T_sub = torch.randn([n_samples, 1, channels, sub_h, sub_w], device=device)
for i in range(frames):
h_start = int(PATHS[i][0] * noise_h)
h_end = h_start + sub_h
w_start = int(PATHS[i][2] * noise_w)
w_end = w_start + sub_w
# no mix
x_T_total[:, :, :, i, h_start:h_end, w_start:w_end] = x_T_sub
filter_shape = [1, channels, frames, noise_h, noise_w]
freq_filter = get_freq_filter(
filter_shape,
device = device,
filter_type='butterworth',
n=4,
d_s=0.25,
d_t=0.1
)
x_T_rand = torch.randn([1, 1, channels, frames, noise_h, noise_w], device=device)
x_T_total = freq_mix_3d(x_T_total.to(dtype=torch.float32), x_T_rand, LPF=freq_filter).to(dtype=dtype).to(device)
#inference
ddim_sampler = DDIMSampler(self.model)
samples, _ = ddim_sampler.sample(
S=steps,
conditioning=cond,
batch_size=noise_shape[0],
shape=noise_shape[1:],
verbose=True,
unconditional_guidance_scale=cfg,
unconditional_conditioning=uc,
eta=eta,
temporal_length=noise_shape[2],
conditional_guidance_scale_temporal=None,
x_T=x_T_total[0],
fs=fs,
timestep_spacing=timestep_spacing,
guidance_rescale=guidance_rescale,
clean_cond=True,
mask=mask,
x0=img_tensor_repeat.clone() if mask is not None else None,
frame_window_size = frame_window_size,
frame_window_stride = frame_window_stride,
ddpm_from=ddpm_from,
idx_list=idx_list,
input_traj=input_traj,
ddim_edit = ddim_edit,
)
assert not torch.isnan(samples).any().item(), "Resulting tensor containts NaNs. I'm unsure why this happens, changing step count and/or image dimensions might help."
## reconstruct from latent to pixel space
self.model.first_stage_model.to(device)
self.model.en_and_decode_n_samples_a_time = 1
decoded_images = self.model.decode_first_stage(samples) #b c t h w
self.model.first_stage_model.to(offload_device)
video = decoded_images.detach().cpu()
video = torch.clamp(video.float(), -1., 1.)
video = (video + 1.0) / 2.0
video = video.squeeze(0).permute(1, 2, 3, 0)
del decoded_images, samples
if not keep_model_loaded:
self.model.to(offload_device)
mm.soft_empty_cache()
# Ensure the final dimensions are divisible by 2
final_H = (orig_H // 2) * 2
final_W = (orig_W // 2) * 2
if video.shape[1] != final_H or video.shape[2] != final_W:
video = F.interpolate(video.permute(0, 3, 1, 2), size=(final_H, final_W), mode="bicubic").permute(0, 2, 3, 1)
last_image = video[-1].unsqueeze(0)
return (video, last_image)
NODE_CLASS_MAPPINGS = {
"DynamiCrafterI2V": DynamiCrafterI2V,
"DynamiCrafterModelLoader": DynamiCrafterModelLoader,
@@ -1409,18 +1306,22 @@ NODE_CLASS_MAPPINGS = {
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise,
"DynamiCrafterFreeTrajSampler": DynamiCrafterFreeTrajSampler
"DownloadAndLoadDynamiCrafterCNModel": DownloadAndLoadDynamiCrafterCNModel,
"DynamiCrafterControlnetApply": DynamiCrafterControlnetApply,
"DynamiCrafterCNLoader": DynamiCrafterCNLoader
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DynamiCrafterI2V": "DynamiCrafterI2V",
"DynamiCrafterModelLoader": "DynamiCrafterModelLoader",
"DynamiCrafterBatchInterpolation": "DynamiCrafterBatchInterpolation",
"ToonCrafterInterpolation": "ToonCrafterInterpolation",
"ToonCrafterDecode": "ToonCrafterDecode",
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel",
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel",
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel",
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise",
"DynamiCrafterFreeTrajSampler": "DynamiCrafterFreeTrajSampler"
"DynamiCrafterModelLoader": "DynamiCrafter ModelLoader",
"DynamiCrafterBatchInterpolation": "DynamiCrafter BatchInterpolation",
"ToonCrafterInterpolation": "ToonCrafter Interpolation",
"ToonCrafterDecode": "ToonCrafter Decode",
"DownloadAndLoadDynamiCrafterModel": "(Down)Load DynamiCrafterModel",
"DownloadAndLoadCLIPModel": "(Down)Load CLIPModel",
"DownloadAndLoadCLIPVisionModel": "(Down)Load CLIPVisionModel",
"DynamiCrafterLoadInitNoise": "DynamiCrafter LoadInitNoise",
"DownloadAndLoadDynamiCrafterCNModel": "(Down)Load DynamiCrafter CNModel",
"DynamiCrafterControlnetApply": "DynamiCrafter ControlnetApply",
"DynamiCrafterCNLoader": "DynamiCrafter CNLoader"
}
-310
View File
@@ -1,310 +0,0 @@
import torch
import torch.fft as fft
import math
def get_longpath(BOX_SIZE_H=0.3, BOX_SIZE_W=0.3, input_mode=4):
if input_mode == 1:
# mode 1
inputs = [[0, 0, 0 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W],
[7, 1-BOX_SIZE_H, 1, (1-BOX_SIZE_W) / 15 * 7, (1-BOX_SIZE_W) / 15 * 7 + BOX_SIZE_W],
[8, 1-BOX_SIZE_H, 1, (1-BOX_SIZE_W) / 15 * 8, (1-BOX_SIZE_W) / 15 * 8 + BOX_SIZE_W],
[15, 0, 0 + BOX_SIZE_H, 1-BOX_SIZE_W, 1],
[16, 0.1, 0.1 + BOX_SIZE_H, 0.9-BOX_SIZE_W, 0.9],
[25, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W],
[31, 0.9-BOX_SIZE_H, 0.9, 0.1, 0.1 + BOX_SIZE_W],
[32, 1-BOX_SIZE_H, 1, 0, 0 + BOX_SIZE_W],
[39, 0, 0 + BOX_SIZE_H, (1-BOX_SIZE_W) / 15 * 7, (1-BOX_SIZE_W) / 15 * 7 + BOX_SIZE_W],
[40, 0, 0 + BOX_SIZE_H, (1-BOX_SIZE_W) / 15 * 8, (1-BOX_SIZE_W) / 15 * 8 + BOX_SIZE_W],
[47, 1-BOX_SIZE_H, 1, 1-BOX_SIZE_W, 1],
[48, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9],
[57, 0.9-BOX_SIZE_H, 0.9, 0.1, 0.1 + BOX_SIZE_W],
[63, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W]]
elif input_mode == 2:
# mode 2
inputs = [[0, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W],
[6, 0.9-BOX_SIZE_H, 0.9, 0.1, 0.1 + BOX_SIZE_W],
[15, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9],
[16, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9],
[22, 0.1, 0.1 + BOX_SIZE_H, 0.9-BOX_SIZE_W, 0.9],
[31, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W],
[32, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W],
[41, 0.1, 0.1 + BOX_SIZE_H, 0.9-BOX_SIZE_W, 0.9],
[47, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9],
[48, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9],
[57, 0.9-BOX_SIZE_H, 0.9, 0.1, 0.1 + BOX_SIZE_W],
[63, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W]]
elif input_mode == 3:
# mode 3 ||||
inputs = [[0, 0, 0 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W],
[9, 1-BOX_SIZE_H, 1, (1-BOX_SIZE_W) / 7 * 1, (1-BOX_SIZE_W) / 7 * 1 + BOX_SIZE_W],
[18, 0, 0 + BOX_SIZE_H, (1-BOX_SIZE_W) / 7 * 2, (1-BOX_SIZE_W) / 7 * 2 + BOX_SIZE_W],
[27, 1-BOX_SIZE_H, 1, (1-BOX_SIZE_W) / 7 * 3, (1-BOX_SIZE_W) / 7 * 3 + BOX_SIZE_W],
[36, 0, 0 + BOX_SIZE_H, (1-BOX_SIZE_W) / 7 * 4, (1-BOX_SIZE_W) / 7 * 4 + BOX_SIZE_W],
[45, 1-BOX_SIZE_H, 1, (1-BOX_SIZE_W) / 7 * 5, (1-BOX_SIZE_W) / 7 * 5 + BOX_SIZE_W],
[54, 0, 0 + BOX_SIZE_H, (1-BOX_SIZE_W) / 7 * 6, (1-BOX_SIZE_W) / 7 * 6 + BOX_SIZE_W],
[63, 1-BOX_SIZE_H, 1, 1-BOX_SIZE_W, 1]]
elif input_mode == 4:
# mode 4 ----
inputs = [[0, 0, 0 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W],
[9, (1-BOX_SIZE_H) / 7 * 1, (1-BOX_SIZE_H) / 7 * 1 + BOX_SIZE_H, 1-BOX_SIZE_W, 1],
[18, (1-BOX_SIZE_H) / 7 * 2, (1-BOX_SIZE_H) / 7 * 2 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W],
[27, (1-BOX_SIZE_H) / 7 * 3, (1-BOX_SIZE_H) / 7 * 3 + BOX_SIZE_H, 1-BOX_SIZE_W, 1],
[36, (1-BOX_SIZE_H) / 7 * 4, (1-BOX_SIZE_H) / 7 * 4 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W],
[45, (1-BOX_SIZE_H) / 7 * 5, (1-BOX_SIZE_H) / 7 * 5 + BOX_SIZE_H, 1-BOX_SIZE_W, 1],
[54, (1-BOX_SIZE_H) / 7 * 6, (1-BOX_SIZE_H) / 7 * 6 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W],
[63, 1-BOX_SIZE_H, 1, 1-BOX_SIZE_W, 1]]
else:
print('error')
exit()
outputs = plan_path(inputs)
# print(outputs)
return outputs
def get_path(BOX_SIZE_H=0.3, BOX_SIZE_W=0.3, input_mode=0):
if input_mode == 0:
# \ d
inputs = [[0, 0, 0 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W], [15, 1-BOX_SIZE_H, 1, 1-BOX_SIZE_W, 1]]
elif input_mode == 1:
# / re d
inputs = [[0, 0, 0 + BOX_SIZE_H, 1-BOX_SIZE_W, 1], [15, 1-BOX_SIZE_H, 1, 0, 0 + BOX_SIZE_W]]
elif input_mode == 2:
# L
inputs = [[0, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W], [6, 0.9-BOX_SIZE_H, 0.9, 0.1, 0.1 + BOX_SIZE_W], [15, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9]]
elif input_mode == 3:
# re L
inputs = [[0, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9], [6, 0.1, 0.1 + BOX_SIZE_H, 0.9-BOX_SIZE_W, 0.9], [15, 0.1, 0.1 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W]]
elif input_mode == 4:
# V
inputs = [[0, 0, 0 + BOX_SIZE_H, 0, 0 + BOX_SIZE_W], [7, 1-BOX_SIZE_H, 1, (1-BOX_SIZE_W) / 15 * 7, (1-BOX_SIZE_W) / 15 * 7 + BOX_SIZE_W], [8, 1-BOX_SIZE_H, 1, (1-BOX_SIZE_W) / 15 * 8, (1-BOX_SIZE_W) / 15 * 8 + BOX_SIZE_W], [15, 0, 0 + BOX_SIZE_H, 1-BOX_SIZE_W, 1]]
elif input_mode == 5:
# re V
inputs = [[0, 1-BOX_SIZE_H, 1, 1-BOX_SIZE_W, 1], [7, 0, 0 + BOX_SIZE_H, (1-BOX_SIZE_W) / 15 * 8, (1-BOX_SIZE_W) / 15 * 8 + BOX_SIZE_W], [8, 0, 0 + BOX_SIZE_H, (1-BOX_SIZE_W) / 15 * 7, (1-BOX_SIZE_W) / 15 * 7 + BOX_SIZE_W], [15, 1-BOX_SIZE_H, 1, 0, 0 + BOX_SIZE_W]]
elif input_mode == 6:
# -- goback
inputs = [[0, 0.35, 0.35 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W], [7, 0.35, 0.35 + BOX_SIZE_H, 0.9-BOX_SIZE_W, 0.9], [8, 0.35, 0.35 + BOX_SIZE_H, 0.9-BOX_SIZE_W, 0.9], [15, 0.35, 0.35 + BOX_SIZE_H, 0.1, 0.1 + BOX_SIZE_W]]
elif input_mode == 7:
# tri
inputs = [[0, 0.1, 0.1 + BOX_SIZE_H, 0.35, 0.35 + BOX_SIZE_W], [5, 0.9-BOX_SIZE_H, 0.9, 0.9-BOX_SIZE_W, 0.9], [10, 0.9-BOX_SIZE_H, 0.9, 0.1, 0.1 + BOX_SIZE_W], [15, 0.1, 0.1 + BOX_SIZE_H, 0.35, 0.35 + BOX_SIZE_W]]
outputs = plan_path(inputs)
return outputs
# input: List([frame, h_start, h_end, w_start, w_end], ...)
# return: List([h_start, h_end, w_start, w_end], ...)
def plan_path(input):
len_input = len(input)
path = [input[0][1:]]
for i in range(1, len_input):
start = input[i-1]
end = input[i]
start_frame = start[0]
end_frame = end[0]
h_start_change = (end[1] - start[1]) / (end_frame - start_frame)
h_end_change = (end[2] - start[2]) / (end_frame - start_frame)
w_start_change = (end[3] - start[3]) / (end_frame - start_frame)
w_end_change = (end[4] - start[4]) / (end_frame - start_frame)
for j in range(start_frame+1, end_frame + 1):
increase_frame = j - start_frame
path += [[increase_frame * h_start_change + start[1], increase_frame * h_end_change + start[2], increase_frame * w_start_change + start[3], increase_frame * w_end_change + start[4]]]
return path
def gaussian_2d(x=0, y=0, mx=0, my=0, sx=1, sy=1):
""" 2d Gaussian weight function
"""
gaussian_map = (
1
/ (2 * math.pi * sx * sy)
* torch.exp(-((x - mx) ** 2 / (2 * sx**2) + (y - my) ** 2 / (2 * sy**2)))
)
gaussian_map.div_(gaussian_map.max())
return gaussian_map
def gaussian_weight(height=32, width=32, KERNEL_DIVISION=3.0):
x = torch.linspace(0, height, height)
y = torch.linspace(0, width, width)
x, y = torch.meshgrid(x, y, indexing="ij")
noise_patch = (
gaussian_2d(
x,
y,
mx=int(height / 2),
my=int(width / 2),
sx=float(height / KERNEL_DIVISION),
sy=float(width / KERNEL_DIVISION),
)
).half()
return noise_patch
def freq_mix_3d(x, noise, LPF):
"""
Noise reinitialization.
Args:
x: diffused latent
noise: randomly sampled noise
LPF: low pass filter
"""
# FFT
x_freq = fft.fftn(x, dim=(-3, -2, -1))
x_freq = fft.fftshift(x_freq, dim=(-3, -2, -1))
noise_freq = fft.fftn(noise, dim=(-3, -2, -1))
noise_freq = fft.fftshift(noise_freq, dim=(-3, -2, -1))
# frequency mix
HPF = 1 - LPF
x_freq_low = x_freq * LPF
noise_freq_high = noise_freq * HPF
x_freq_mixed = x_freq_low + noise_freq_high # mix in freq domain
# IFFT
x_freq_mixed = fft.ifftshift(x_freq_mixed, dim=(-3, -2, -1))
x_mixed = fft.ifftn(x_freq_mixed, dim=(-3, -2, -1)).real
return x_mixed
def get_freq_filter(shape, device, filter_type, n, d_s, d_t):
"""
Form the frequency filter for noise reinitialization.
Args:
shape: shape of latent (B, C, T, H, W)
filter_type: type of the freq filter
n: (only for butterworth) order of the filter, larger n ~ ideal, smaller n ~ gaussian
d_s: normalized stop frequency for spatial dimensions (0.0-1.0)
d_t: normalized stop frequency for temporal dimension (0.0-1.0)
"""
if filter_type == "gaussian":
return gaussian_low_pass_filter(shape=shape, d_s=d_s, d_t=d_t).to(device)
elif filter_type == "ideal":
return ideal_low_pass_filter(shape=shape, d_s=d_s, d_t=d_t).to(device)
elif filter_type == "box":
return box_low_pass_filter(shape=shape, d_s=d_s, d_t=d_t).to(device)
elif filter_type == "butterworth":
return butterworth_low_pass_filter(shape=shape, n=n, d_s=d_s, d_t=d_t).to(device)
else:
raise NotImplementedError
def gaussian_low_pass_filter(shape, d_s=0.25, d_t=0.25):
"""
Compute the gaussian low pass filter mask.
Args:
shape: shape of the filter (volume)
d_s: normalized stop frequency for spatial dimensions (0.0-1.0)
d_t: normalized stop frequency for temporal dimension (0.0-1.0)
"""
T, H, W = shape[-3], shape[-2], shape[-1]
mask = torch.zeros(shape)
if d_s==0 or d_t==0:
return mask
for t in range(T):
for h in range(H):
for w in range(W):
d_square = (((d_s/d_t)*(2*t/T-1))**2 + (2*h/H-1)**2 + (2*w/W-1)**2)
mask[..., t,h,w] = math.exp(-1/(2*d_s**2) * d_square)
return mask
def butterworth_low_pass_filter(shape, n=4, d_s=0.25, d_t=0.25):
"""
Compute the butterworth low pass filter mask.
Args:
shape: shape of the filter (volume)
n: order of the filter, larger n ~ ideal, smaller n ~ gaussian
d_s: normalized stop frequency for spatial dimensions (0.0-1.0)
d_t: normalized stop frequency for temporal dimension (0.0-1.0)
"""
T, H, W = shape[-3], shape[-2], shape[-1]
mask = torch.zeros(shape)
if d_s==0 or d_t==0:
return mask
for t in range(T):
for h in range(H):
for w in range(W):
d_square = (((d_s/d_t)*(2*t/T-1))**2 + (2*h/H-1)**2 + (2*w/W-1)**2)
mask[..., t,h,w] = 1 / (1 + (d_square / d_s**2)**n)
return mask
def ideal_low_pass_filter(shape, d_s=0.25, d_t=0.25):
"""
Compute the ideal low pass filter mask.
Args:
shape: shape of the filter (volume)
d_s: normalized stop frequency for spatial dimensions (0.0-1.0)
d_t: normalized stop frequency for temporal dimension (0.0-1.0)
"""
T, H, W = shape[-3], shape[-2], shape[-1]
mask = torch.zeros(shape)
if d_s==0 or d_t==0:
return mask
for t in range(T):
for h in range(H):
for w in range(W):
d_square = (((d_s/d_t)*(2*t/T-1))**2 + (2*h/H-1)**2 + (2*w/W-1)**2)
mask[..., t,h,w] = 1 if d_square <= d_s*2 else 0
return mask
def box_low_pass_filter(shape, d_s=0.25, d_t=0.25):
"""
Compute the ideal low pass filter mask (approximated version).
Args:
shape: shape of the filter (volume)
d_s: normalized stop frequency for spatial dimensions (0.0-1.0)
d_t: normalized stop frequency for temporal dimension (0.0-1.0)
"""
T, H, W = shape[-3], shape[-2], shape[-1]
mask = torch.zeros(shape)
if d_s==0 or d_t==0:
return mask
threshold_s = round(int(H // 2) * d_s)
threshold_t = round(T // 2 * d_t)
cframe, crow, ccol = T // 2, H // 2, W //2
mask[..., cframe - threshold_t:cframe + threshold_t, crow - threshold_s:crow + threshold_s, ccol - threshold_s:ccol + threshold_s] = 1.0
return mask
def load_idx(prompt_file):
f = open(prompt_file, 'r')
idx_list = []
for idx, line in enumerate(f.readlines()):
l = line.strip()
if len(l) != 0:
indices = l.split(',')
indices_list = []
for index in indices:
indices_list.append(int(index))
idx_list.append(indices_list)
f.close()
return idx_list
def load_traj(prompt_file):
f = open(prompt_file, 'r')
traj_list = []
for idx, line in enumerate(f.readlines()):
l = line.strip()
if len(l) != 0:
numbers = l.split(',')
numbers_list = []
for number_index in range(len(numbers)):
if number_index == 0:
numbers_list.append(int(numbers[number_index]))
else:
numbers_list.append(float(numbers[number_index]))
traj_list.append(numbers_list)
f.close()
return traj_list