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, logdir=None,
rand_cond_frame=False, rand_cond_frame=False,
en_and_decode_n_samples_a_time=None, en_and_decode_n_samples_a_time=None,
control_scale=1.0,
*args, **kwargs): *args, **kwargs):
self.num_timesteps_cond = default(num_timesteps_cond, 1) self.num_timesteps_cond = default(num_timesteps_cond, 1)
self.scale_by_std = scale_by_std self.scale_by_std = scale_by_std
@@ -404,6 +405,7 @@ class LatentDiffusion(DDPM):
self.loop_video = loop_video self.loop_video = loop_video
self.fps_condition_type = fps_condition_type self.fps_condition_type = fps_condition_type
self.perframe_ae = perframe_ae self.perframe_ae = perframe_ae
self.control_scale = control_scale
self.logdir = logdir self.logdir = logdir
self.rand_cond_frame = rand_cond_frame self.rand_cond_frame = rand_cond_frame
@@ -565,7 +567,6 @@ class LatentDiffusion(DDPM):
return self.p_losses(x, c, t, **kwargs) return self.p_losses(x, c, t, **kwargs)
def apply_model(self, x_noisy, t, cond, **kwargs): def apply_model(self, x_noisy, t, cond, **kwargs):
print("apply_model:", kwargs)
if isinstance(cond, dict): if isinstance(cond, dict):
# hybrid case, cond is exptected to be a dict # hybrid case, cond is exptected to be a dict
pass pass
@@ -573,9 +574,20 @@ class LatentDiffusion(DDPM):
if not isinstance(cond, list): if not isinstance(cond, list):
cond = [cond] cond = [cond]
key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn' 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): if isinstance(x_recon, tuple):
return x_recon[0] return x_recon[0]
@@ -723,7 +735,7 @@ class DiffusionWrapper(pl.LightningModule):
self.diffusion_model = instantiate_from_config(diff_model_config) self.diffusion_model = instantiate_from_config(diff_model_config)
self.conditioning_key = conditioning_key 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): c_adm=None, s=None, mask=None, **kwargs):
# temporal_context = fps is foNone # temporal_context = fps is foNone
if self.conditioning_key is None: 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 ## it is just right [b,c,t,h,w]: concatenate in channel dim
xc = torch.cat([x] + c_concat, dim=1) xc = torch.cat([x] + c_concat, dim=1)
cc = torch.cat(c_crossattn, 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': elif self.conditioning_key == 'resblockcond':
cc = c_crossattn[0] cc = c_crossattn[0]
out = self.diffusion_model(x, t, context=cc) out = self.diffusion_model(x, t, context=cc)
+3 -8
View File
@@ -197,7 +197,7 @@ class DDIMSampler(object):
corrector_kwargs=corrector_kwargs, corrector_kwargs=corrector_kwargs,
unconditional_guidance_scale=unconditional_guidance_scale, unconditional_guidance_scale=unconditional_guidance_scale,
unconditional_conditioning=unconditional_conditioning, 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) **kwargs)
@@ -214,7 +214,7 @@ class DDIMSampler(object):
@torch.no_grad() @torch.no_grad()
def p_sample_ddim(self, x, c, t, sigmas, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False, 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, 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): uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs):
b, *_, device = *x.shape, x.device b, *_, device = *x.shape, x.device
if x.dim() == 5: if x.dim() == 5:
@@ -225,13 +225,8 @@ class DDIMSampler(object):
if unconditional_conditioning is None or unconditional_guidance_scale == 1.: if unconditional_conditioning is None or unconditional_guidance_scale == 1.:
model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser
else: 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 ### 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_cond = self.model.apply_model(x, t, c, **kwargs)
e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs) e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs)
else: else:
+11 -237
View File
@@ -18,36 +18,6 @@ from ...lvdm.basics import zero_module
import comfy.ops import comfy.ops
ops = comfy.ops.manual_cast 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): class RelativePosition(nn.Module):
""" https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """ """ 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.)) ) self.register_parameter('alpha', nn.Parameter(torch.tensor(0.)) )
def forward(self, x, context=None, mask=None, use_freetraj=False, idx_list=[], input_traj=[]): def forward(self, x, context=None, mask=None):
print("forward input_traj: ", input_traj)
spatial_self_attn = (context is None) spatial_self_attn = (context is None)
k_ip, v_ip, out_ip = None, None, None k_ip, v_ip, out_ip = None, None, None
@@ -130,18 +99,9 @@ class CrossAttention(nn.Module):
context = context[:,:self.text_context_len,:] context = context[:,:self.text_context_len,:]
k = self.to_k(context) k = self.to_k(context)
v = self.to_v(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)) 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 sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale
if self.relative_position: if self.relative_position:
len_q, len_k, len_v = q.shape[1], k.shape[1], v.shape[1] 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 sim += sim2
del k 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): if exists(mask):
## feasible for causal attention mask only ## feasible for causal attention mask only
max_neg_value = -torch.finfo(sim.dtype).max max_neg_value = -torch.finfo(sim.dtype).max
@@ -202,6 +118,7 @@ class CrossAttention(nn.Module):
# attention, what we cannot get enough of # attention, what we cannot get enough of
sim = sim.softmax(dim=-1) sim = sim.softmax(dim=-1)
out = torch.einsum('b i j, b j d -> b i d', sim, v) out = torch.einsum('b i j, b j d -> b i d', sim, v)
if self.relative_position: if self.relative_position:
v2 = self.relative_position_v(len_q, len_v) 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 out = out + self.image_cross_attention_scale * out_ip
return self.to_out(out) 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): def efficient_forward(self, x, context=None, mask=None):
spatial_self_attn = (context is None) spatial_self_attn = (context is None)
@@ -454,7 +230,7 @@ class BasicTransformerBlock(nn.Module):
self.checkpoint = checkpoint 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 ## 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 input_tuple = (x,) ## should not be (x), otherwise *input_tuple will decouple x into multiple arguments
if context is not None: if context is not None:
@@ -462,14 +238,12 @@ class BasicTransformerBlock(nn.Module):
if mask is not None: if mask is not None:
forward_mask = partial(self._forward, mask=mask) forward_mask = partial(self._forward, mask=mask)
return checkpoint(forward_mask, (x,), self.parameters(), self.checkpoint) 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) 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 def _forward(self, x, context=None, mask=None):
x = self.attn2(self.norm2(x), context=context, mask=mask, use_freetraj=use_freetraj, idx_list=idx_list, input_traj=input_traj) + x 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 x = self.ff(self.norm3(x)) + x
return x return x
@@ -591,7 +365,7 @@ class TemporalTransformer(nn.Module):
self.use_linear = use_linear 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 B, C, T, H, W = x_in.shape
def process_slice(x, t_start=None, t_end=None): def process_slice(x, t_start=None, t_end=None):
b, c, t, h, w = x.shape b, c, t, h, w = x.shape
@@ -620,7 +394,7 @@ class TemporalTransformer(nn.Module):
if self.only_self_att: if self.only_self_att:
## note: if no context is given, cross-attention defaults to self-attention ## note: if no context is given, cross-attention defaults to self-attention
for i, block in enumerate(self.transformer_blocks): 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() x = rearrange(x, "(b hw) t c -> b hw t c", b=b).contiguous()
else: else:
x = rearrange(x, "(b hw) t c -> b hw t c", b=b).contiguous() 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 context[j], "t l con -> (t r) l con", r=(h * w) // t, t=t
).contiguous() ).contiguous()
## note: causal mask will not applied in cross-attention case ## 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: if self.use_linear:
x = self.proj_out(x) x = self.proj_out(x)
+412 -10
View File
@@ -2,7 +2,9 @@ from functools import partial
from abc import abstractmethod from abc import abstractmethod
import torch import torch
import torch.nn as nn import torch.nn as nn
import numpy as np
from einops import rearrange from einops import rearrange
import math
import torch.nn.functional as F import torch.nn.functional as F
from ....lvdm.models.utils_diffusion import timestep_embedding from ....lvdm.models.utils_diffusion import timestep_embedding
from ....lvdm.common import checkpoint from ....lvdm.common import checkpoint
@@ -18,6 +20,9 @@ from ....lvdm.modules.attention import SpatialTransformer, TemporalTransformer
import comfy.ops import comfy.ops
ops = comfy.ops.manual_cast ops = comfy.ops.manual_cast
def exists(x):
return x is not None
class TimestepBlock(nn.Module): class TimestepBlock(nn.Module):
""" """
Any module where forward() takes timestep embeddings as a second argument. 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. 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: for layer in self:
if isinstance(layer, TimestepBlock): if isinstance(layer, TimestepBlock):
x = layer(x, emb, batch_size=batch_size) x = layer(x, emb, batch_size=batch_size)
elif isinstance(layer, SpatialTransformer): elif isinstance(layer, SpatialTransformer):
x = layer(x, context, use_freetraj=use_freetraj, **kwargs) x = layer(x, context)
elif isinstance(layer, TemporalTransformer): elif isinstance(layer, TemporalTransformer):
x = rearrange(x, '(b f) c h w -> b c f h w', b=batch_size) 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') x = rearrange(x, 'b c f h w -> (b f) c h w')
else: else:
x = layer(x) x = layer(x)
@@ -547,8 +552,7 @@ class UNetModel(nn.Module):
zero_module(conv_nd(dims, model_channels, out_channels, 3, padding=1)), 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): def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, control=None, **kwargs):
print("unet forward: ", kwargs)
b,_,t,_,_ = x.shape b,_,t,_,_ = x.shape
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype) t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype)
emb = self.time_embed(t_emb) emb = self.time_embed(t_emb)
@@ -583,9 +587,9 @@ class UNetModel(nn.Module):
adapter_idx = 0 adapter_idx = 0
hs = [] hs = []
for id, module in enumerate(self.input_blocks): 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: 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 ## plug-in adapter features
if ((id+1)%3 == 0) and features_adapter is not None: if ((id+1)%3 == 0) and features_adapter is not None:
h = h + features_adapter[adapter_idx] h = h + features_adapter[adapter_idx]
@@ -594,13 +598,411 @@ class UNetModel(nn.Module):
if features_adapter is not None: if features_adapter is not None:
assert len(features_adapter)==adapter_idx, 'Wrong features_adapter' 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: for module in self.output_blocks:
h = torch.cat([h, hs.pop()], dim=1) if control is None:
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs) 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) h = h.type(x.dtype)
y = self.out(h) y = self.out(h)
# reshape back to (b c t h w) # reshape back to (b c t h w)
y = rearrange(y, '(b t) c h w -> b c t h w', b=b) y = rearrange(y, '(b t) c h w -> b c t h w', b=b)
return y 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 import torch.nn.functional as F
from .scripts.evaluation.funcs import load_model_checkpoint, get_latent_z, get_latent_z_with_hidden_states 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 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 from einops import repeat
import folder_paths import folder_paths
import comfy.model_management as mm import comfy.model_management as mm
import comfy.utils import comfy.utils
from contextlib import nullcontext from contextlib import nullcontext
from .lvdm.models.samplers.ddim import DDIMSampler from .lvdm.models.samplers.ddim import DDIMSampler
from .lvdm.modules.networks.openaimodel3d import ControlNet
from contextlib import nullcontext from contextlib import nullcontext
try: try:
@@ -143,8 +141,130 @@ class DownloadAndLoadDynamiCrafterModel:
dcmodel = { dcmodel = {
'model': self.model, 'model': self.model,
'model_name': model, 'model_name': model,
'dtype': precision
} }
return (dcmodel,) 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: class DownloadAndLoadCLIPModel:
@classmethod @classmethod
@@ -392,6 +512,7 @@ class DynamiCrafterI2V:
self.model.to(device) self.model.to(device)
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(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(): 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) image = image.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0: if augmentation_level > 0:
image += torch.randn_like(image) * augmentation_level image += torch.randn_like(image) * augmentation_level
@@ -403,7 +524,7 @@ class DynamiCrafterI2V:
if H % 64 != 0: if H % 64 != 0:
H = H - (H % 64) H = H - (H % 64)
if orig_H % 64 != 0 or orig_W % 64 != 0: 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 B, C, H, W = image.shape
noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8] 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 z = get_latent_z(self.model, encode_pixels) #bc,1,hw
if image2 is not None: if image2 is not None:
image2 = image2 * 2 - 1
image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device) image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0: if augmentation_level > 0:
image2 += torch.randn_like(image2) * augmentation_level image2 += torch.randn_like(image2) * augmentation_level
if image2.shape != image.shape: 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 encode_pixels = image2.unsqueeze(2) * 2 - 1
z2 = get_latent_z(self.model, encode_pixels) #bc,1,hw z2 = get_latent_z(self.model, encode_pixels) #bc,1,hw
@@ -437,6 +557,8 @@ class DynamiCrafterI2V:
text_emb = positive[0][0].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) 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) img_emb = self.model.image_proj_model(cond_images)
@@ -444,7 +566,9 @@ class DynamiCrafterI2V:
del cond_images, img_emb, text_emb, encode_pixels del cond_images, img_emb, text_emb, encode_pixels
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device) 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: if noise_shape[-1] == 32:
timestep_spacing = "uniform" timestep_spacing = "uniform"
@@ -636,6 +760,7 @@ class ToonCrafterInterpolation:
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}), "augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"optional_latents": ("LATENT",), "optional_latents": ("LATENT",),
"ddpm_from": ("INT", {"default": 1000, "min": 1, "max": 1000, "step": 1}), "ddpm_from": ("INT", {"default": 1000, "min": 1, "max": 1000, "step": 1}),
"controlnet": ("DC_CONTROL",),
} }
} }
@@ -644,7 +769,8 @@ class ToonCrafterInterpolation:
FUNCTION = "process" FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper" 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() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
mm.unload_all_models() mm.unload_all_models()
@@ -654,6 +780,11 @@ class ToonCrafterInterpolation:
self.model = model['model'] 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 dtype = self.model.dtype
if vae_dtype == "auto": if vae_dtype == "auto":
try: try:
@@ -685,11 +816,12 @@ class ToonCrafterInterpolation:
pbar = comfy.utils.ProgressBar(len(images) - 1) pbar = comfy.utils.ProgressBar(len(images) - 1)
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(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(): 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 videos, videos2 = None, None
mm.soft_empty_cache() mm.soft_empty_cache()
image = images[i].unsqueeze(0) 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 B, C, H, W = image.shape
noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8] 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 image2 += torch.randn_like(image) * augmentation_level
encode_pixels = image.unsqueeze(2) * 2 - 1 encode_pixels = image.unsqueeze(2) * 2 - 1
videos = encode_pixels # bc1hw videos = encode_pixels # bc1hw
videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) 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 if len(images) == 1:
videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) videos = torch.cat([videos, videos], dim=2)
videos = torch.cat([videos, videos2], 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: try:
z, hs = get_latent_z_with_hidden_states(self.model, videos) 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 = torch.zeros_like(z)
img_tensor_repeat[:,:,:1,:,:] = z[:,:,:1,:,:] 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) self.model.first_stage_model.to(offload_device)
text_emb = positive[0][0].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)
cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
self.model.image_proj_model.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_emb = self.model.image_proj_model(cond_images)
img_emb2 = self.model.image_proj_model(cond_images2) if len(images) !=1:
img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio) 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) 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): if comfy.model_management.is_device_mps(device):
fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device) fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device)
else: else:
fs = torch.tensor([fs], dtype=torch.float64, device=self.model.device) 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: if noise_shape[-1] == 32:
timestep_spacing = "uniform" timestep_spacing = "uniform"
@@ -822,6 +967,41 @@ class ToonCrafterInterpolation:
} }
return (latent,) 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: class ToonCrafterDecode:
@classmethod @classmethod
@@ -955,8 +1135,9 @@ class DynamiCrafterBatchInterpolation:
mm.soft_empty_cache() mm.soft_empty_cache()
torch.manual_seed(seed) torch.manual_seed(seed)
dtype = model.dtype
self.model = model['model'] self.model = model['model']
dtype = self.model.dtype
if vae_dtype == "auto": if vae_dtype == "auto":
try: try:
@@ -971,7 +1152,6 @@ class DynamiCrafterBatchInterpolation:
print(f"VAE using dtype: {self.model.first_stage_model.dtype}") print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
self.model.to(device) self.model.to(device)
images = images * 2 - 1
images = images.permute(0, 3, 1, 2).to(dtype).to(device) images = images.permute(0, 3, 1, 2).to(dtype).to(device)
B, C, H, W = images.shape B, C, H, W = images.shape
orig_H, orig_W = H, W orig_H, orig_W = H, W
@@ -994,8 +1174,11 @@ class DynamiCrafterBatchInterpolation:
self.model.first_stage_model.to(device) self.model.first_stage_model.to(device)
z = get_latent_z(self.model, image.unsqueeze(2)) #bc,1,hw encode_pixels1 = image * 2 - 1
z2 = get_latent_z(self.model, image2.unsqueeze(2)) #bc,1,hw 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 = 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 = torch.zeros_like(img_tensor_repeat)
img_tensor_repeat[:,:,:1,:,:] = z img_tensor_repeat[:,:,:1,:,:] = z
@@ -1012,7 +1195,9 @@ class DynamiCrafterBatchInterpolation:
imtext_cond = torch.cat([text_emb, img_emb], dim=1) imtext_cond = torch.cat([text_emb, img_emb], dim=1)
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device) 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: if noise_shape[-1] == 32:
timestep_spacing = "uniform" timestep_spacing = "uniform"
@@ -1111,294 +1296,6 @@ class DynamiCrafterBatchInterpolation:
last_image = out_video[-1].unsqueeze(0) last_image = out_video[-1].unsqueeze(0)
return (out_video, last_image) 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 = { NODE_CLASS_MAPPINGS = {
"DynamiCrafterI2V": DynamiCrafterI2V, "DynamiCrafterI2V": DynamiCrafterI2V,
"DynamiCrafterModelLoader": DynamiCrafterModelLoader, "DynamiCrafterModelLoader": DynamiCrafterModelLoader,
@@ -1409,18 +1306,22 @@ NODE_CLASS_MAPPINGS = {
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel, "DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel, "DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise, "DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise,
"DynamiCrafterFreeTrajSampler": DynamiCrafterFreeTrajSampler "DownloadAndLoadDynamiCrafterCNModel": DownloadAndLoadDynamiCrafterCNModel,
"DynamiCrafterControlnetApply": DynamiCrafterControlnetApply,
"DynamiCrafterCNLoader": DynamiCrafterCNLoader
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"DynamiCrafterI2V": "DynamiCrafterI2V", "DynamiCrafterI2V": "DynamiCrafterI2V",
"DynamiCrafterModelLoader": "DynamiCrafterModelLoader", "DynamiCrafterModelLoader": "DynamiCrafter ModelLoader",
"DynamiCrafterBatchInterpolation": "DynamiCrafterBatchInterpolation", "DynamiCrafterBatchInterpolation": "DynamiCrafter BatchInterpolation",
"ToonCrafterInterpolation": "ToonCrafterInterpolation", "ToonCrafterInterpolation": "ToonCrafter Interpolation",
"ToonCrafterDecode": "ToonCrafterDecode", "ToonCrafterDecode": "ToonCrafter Decode",
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel", "DownloadAndLoadDynamiCrafterModel": "(Down)Load DynamiCrafterModel",
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel", "DownloadAndLoadCLIPModel": "(Down)Load CLIPModel",
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel", "DownloadAndLoadCLIPVisionModel": "(Down)Load CLIPVisionModel",
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise", "DynamiCrafterLoadInitNoise": "DynamiCrafter LoadInitNoise",
"DynamiCrafterFreeTrajSampler": "DynamiCrafterFreeTrajSampler" "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