Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d312c62982 | ||
|
|
7670b09c28 | ||
|
|
282c69c5f3 | ||
|
|
18121e97a2 | ||
|
|
4a7dd402b4 | ||
|
|
2545a3afb7 | ||
|
|
0668ee6d0a | ||
|
|
69fe0e2e9f |
+17
-5
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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])
|
||||
@@ -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:
|
||||
@@ -83,7 +81,7 @@ class DownloadAndLoadDynamiCrafterModel:
|
||||
'ckpt_name': model,
|
||||
'fp8_unet': fp8_unet
|
||||
}
|
||||
if not hasattr(self, 'model') or self.model == None or custom_config != self.current_config:
|
||||
if not hasattr(self, 'model') or self.model is None or custom_config != self.current_config:
|
||||
self.current_config = custom_config
|
||||
download_path = os.path.join(folder_paths.models_dir, "checkpoints", "dynamicrafter")
|
||||
model_path = os.path.join(download_path, model)
|
||||
@@ -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
|
||||
@@ -444,7 +564,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 +758,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 +767,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 +778,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:
|
||||
@@ -740,7 +869,14 @@ class ToonCrafterInterpolation:
|
||||
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 +958,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 +1126,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 +1143,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 +1165,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 +1186,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 +1287,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 +1297,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"
|
||||
}
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-dynamicrafterwrapper"
|
||||
description = "Wrapper nodes to use Dynami/ToonCrafter image2video and frame interpolation models in ComfyUI"
|
||||
version = "1.0.2"
|
||||
version = "1.0.3"
|
||||
license = "Apache-2.0"
|
||||
dependencies = ["einops>=0.3.0", "numpy>=1.24.2", "omegaconf>=2.1.1", "pytorch_lightning>=2.2.1", "tqdm>=4.65.0", "transformers>=4.25.1", "timm"]
|
||||
|
||||
|
||||
+5
-2
@@ -34,12 +34,15 @@ def instantiate_from_config(config):
|
||||
|
||||
|
||||
def get_obj_from_str(string, reload=False):
|
||||
package_directory_name = os.path.basename(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
module, cls = string.rsplit(".", 1)
|
||||
if reload:
|
||||
module_imp = importlib.import_module(module)
|
||||
importlib.reload(module_imp)
|
||||
return getattr(importlib.import_module(module, package=package_directory_name), cls)
|
||||
try:
|
||||
obj = getattr(importlib.import_module(module, package=os.path.basename(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))), cls)
|
||||
except:
|
||||
obj = getattr(importlib.import_module(module, package=os.path.dirname(os.path.dirname(os.path.abspath( __file__ )))), cls)
|
||||
return obj
|
||||
|
||||
|
||||
def load_npz_from_dir(data_dir):
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user