Author SHA1 Message Date
kijai 44143a4282 testing, doesn't work 2024-07-09 15:46:34 +03:00
6 changed files with 859 additions and 23 deletions
+1
View File
@@ -565,6 +565,7 @@ 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
+8 -3
View File
@@ -197,7 +197,7 @@ class DDIMSampler(object):
corrector_kwargs=corrector_kwargs,
unconditional_guidance_scale=unconditional_guidance_scale,
unconditional_conditioning=unconditional_conditioning,
mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale,
mask=mask,x0=x0,step=i,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,
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):
b, *_, device = *x.shape, x.device
if x.dim() == 5:
@@ -225,8 +225,13 @@ 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
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_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs)
else:
+237 -11
View File
@@ -18,6 +18,36 @@ 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 """
@@ -80,7 +110,8 @@ class CrossAttention(nn.Module):
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)
k_ip, v_ip, out_ip = None, None, None
@@ -99,9 +130,18 @@ 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]
@@ -110,6 +150,50 @@ 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
@@ -118,7 +202,6 @@ 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)
@@ -144,6 +227,147 @@ 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)
@@ -230,7 +454,7 @@ class BasicTransformerBlock(nn.Module):
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
input_tuple = (x,) ## should not be (x), otherwise *input_tuple will decouple x into multiple arguments
if context is not None:
@@ -238,12 +462,14 @@ 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):
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
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
x = self.ff(self.norm3(x)) + x
return x
@@ -365,7 +591,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):
def forward(self, x_in, context=None, frame_window_size=None, frame_window_stride=None, **kwargs):
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
@@ -394,7 +620,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)
x = block(x, mask=mask, **kwargs)
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()
@@ -406,7 +632,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)
x[j] = block(x[j], context=context_j, **kwargs)
if self.use_linear:
x = self.proj_out(x)
+8 -7
View File
@@ -35,15 +35,15 @@ class TimestepEmbedSequential(nn.Sequential, TimestepBlock):
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:
if isinstance(layer, TimestepBlock):
x = layer(x, emb, batch_size=batch_size)
elif isinstance(layer, SpatialTransformer):
x = layer(x, context)
x = layer(x, context, use_freetraj=use_freetraj, **kwargs)
elif isinstance(layer, TemporalTransformer):
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')
else:
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):
print("unet forward: ", 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)
@@ -582,9 +583,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)
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:
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
if ((id+1)%3 == 0) and features_adapter is not None:
h = h + features_adapter[adapter_idx]
@@ -593,10 +594,10 @@ 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)
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:
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)
y = self.out(h)
+295 -2
View File
@@ -4,6 +4,7 @@ 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
@@ -11,6 +12,8 @@ import comfy.utils
from contextlib import nullcontext
from .lvdm.models.samplers.ddim import DDIMSampler
from contextlib import nullcontext
try:
from accelerate import init_empty_weights
@@ -1108,6 +1111,294 @@ 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,
@@ -1117,7 +1408,8 @@ NODE_CLASS_MAPPINGS = {
"DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel,
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise,
"DynamiCrafterFreeTrajSampler": DynamiCrafterFreeTrajSampler
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -1129,5 +1421,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel",
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel",
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel",
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise"
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise",
"DynamiCrafterFreeTrajSampler": "DynamiCrafterFreeTrajSampler"
}
+310
View File
@@ -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