Compare commits
1
Commits
main
...
freetraj_testing
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
44143a4282 |
@@ -565,6 +565,7 @@ 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
|
||||||
|
|||||||
@@ -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,fs=fs,guidance_rescale=guidance_rescale,
|
mask=mask,x0=x0,step=i,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,
|
unconditional_guidance_scale=1., unconditional_conditioning=None, step=0, ddim_edit=0,
|
||||||
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,8 +225,13 @@ 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
|
||||||
if isinstance(c, torch.Tensor) or isinstance(c, dict):
|
elif 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:
|
||||||
|
|||||||
+237
-11
@@ -18,6 +18,36 @@ 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 """
|
||||||
@@ -80,7 +110,8 @@ 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):
|
def forward(self, x, context=None, mask=None, use_freetraj=False, idx_list=[], input_traj=[]):
|
||||||
|
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
|
||||||
|
|
||||||
@@ -99,9 +130,18 @@ 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]
|
||||||
@@ -110,6 +150,50 @@ 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
|
||||||
@@ -118,7 +202,6 @@ 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)
|
||||||
@@ -144,6 +227,147 @@ 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)
|
||||||
@@ -230,7 +454,7 @@ class BasicTransformerBlock(nn.Module):
|
|||||||
self.checkpoint = checkpoint
|
self.checkpoint = checkpoint
|
||||||
|
|
||||||
|
|
||||||
def forward(self, x, context=None, mask=None, **kwargs):
|
def forward(self, x, context=None, mask=None, use_freetraj=False, idx_list=[], input_traj=[], **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:
|
||||||
@@ -238,12 +462,14 @@ 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=[]):
|
||||||
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, 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, use_freetraj=use_freetraj, idx_list=idx_list, input_traj=input_traj) + 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
|
||||||
|
|
||||||
@@ -365,7 +591,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):
|
def forward(self, x_in, context=None, frame_window_size=None, frame_window_stride=None, **kwargs):
|
||||||
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
|
||||||
@@ -394,7 +620,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)
|
x = block(x, mask=mask, **kwargs)
|
||||||
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()
|
||||||
@@ -406,7 +632,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)
|
x[j] = block(x[j], context=context_j, **kwargs)
|
||||||
|
|
||||||
if self.use_linear:
|
if self.use_linear:
|
||||||
x = self.proj_out(x)
|
x = self.proj_out(x)
|
||||||
|
|||||||
@@ -35,15 +35,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, frame_window_size=None, frame_window_stride=None):
|
def forward(self, x, emb, context=None, batch_size=None, use_freetraj=False, frame_window_size=None, frame_window_stride=None, **kwargs):
|
||||||
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)
|
x = layer(x, context, use_freetraj=use_freetraj, **kwargs)
|
||||||
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, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
|
x = layer(x, context, use_freetraj=use_freetraj, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
|
||||||
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)
|
||||||
@@ -548,6 +548,7 @@ class UNetModel(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
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, **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)
|
||||||
@@ -582,9 +583,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)
|
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
|
||||||
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)
|
h = self.init_attn(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
|
||||||
## 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]
|
||||||
@@ -593,10 +594,10 @@ 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)
|
h = self.middle_block(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
|
||||||
for module in self.output_blocks:
|
for module in self.output_blocks:
|
||||||
h = torch.cat([h, hs.pop()], dim=1)
|
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)
|
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride, **kwargs)
|
||||||
h = h.type(x.dtype)
|
h = h.type(x.dtype)
|
||||||
y = self.out(h)
|
y = self.out(h)
|
||||||
|
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ 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
|
||||||
@@ -11,6 +12,8 @@ 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 contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
try:
|
try:
|
||||||
from accelerate import init_empty_weights
|
from accelerate import init_empty_weights
|
||||||
@@ -1108,6 +1111,294 @@ 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,
|
||||||
@@ -1117,7 +1408,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel,
|
"DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel,
|
||||||
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
|
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
|
||||||
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
|
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
|
||||||
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise
|
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise,
|
||||||
|
"DynamiCrafterFreeTrajSampler": DynamiCrafterFreeTrajSampler
|
||||||
|
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
@@ -1129,5 +1421,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel",
|
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel",
|
||||||
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel",
|
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel",
|
||||||
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel",
|
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel",
|
||||||
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise"
|
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise",
|
||||||
|
"DynamiCrafterFreeTrajSampler": "DynamiCrafterFreeTrajSampler"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,310 @@
|
|||||||
|
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