This commit is contained in:
kijai
2025-08-27 02:25:39 +03:00
parent 636b252c7e
commit 5d17484cc5
4 changed files with 265 additions and 118 deletions
+8 -3
View File
@@ -2226,6 +2226,7 @@ class WanVideoSampler:
s2v_ref_latent = s2v_audio_embeds["ref_latent"]
if s2v_ref_latent is not None:
s2v_ref_latent = s2v_ref_latent.to(device, dtype)
s2v_audio_input = s2v_audio_input[..., 0:image_embeds["num_frames"]]
#s2v_audio_input_all_layers = s2v_audio_embeds["audio_encoder_output"]["encoded_audio_all_layers"]
print(s2v_audio_input.shape)
##print(s2v_audio_input_all_layers[0].shape)
@@ -2509,7 +2510,7 @@ class WanVideoSampler:
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False,
mtv_motion_tokens=None):
mtv_motion_tokens=None, s2v_audio_input=None):
nonlocal transformer
z = z.to(dtype)
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
@@ -3202,6 +3203,10 @@ class WanVideoSampler:
log.info(f"context window: {c}")
log.info(f"motion_token_indices: {start_token_index}-{end_token_index}")
partial_s2v_audio_input = None
if s2v_audio_input is not None:
partial_s2v_audio_input = s2v_audio_input[..., c]
partial_add_cond = None
if add_cond is not None:
partial_add_cond = add_cond[:, :, c].to(device, dtype)
@@ -3219,7 +3224,7 @@ class WanVideoSampler:
text_embeds["negative_prompt_embeds"],
partial_timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input,
mtv_motion_tokens=partial_mtv_motion_tokens)
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input)
if cache_args is not None:
self.window_tracker.cache_states[window_id] = new_teacache
@@ -3653,7 +3658,7 @@ class WanVideoSampler:
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens)
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input)
if bidirectional_sampling:
noise_pred_flipped, self.cache_state = predict_with_cfg(
latent_model_input_flipped,
-2
View File
@@ -92,8 +92,6 @@ class WanVideoAddAudioEmbeds:
elif len(audio_embed_bucket.shape) == 4:
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1)
audio_embed_bucket = audio_embed_bucket[..., 0:frames]
print("audio_embed_bucket", audio_embed_bucket.shape)
new_entry = {
+226 -8
View File
@@ -30,6 +30,8 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
from ...MTV.mtv import apply_rotary_emb
from .s2v.motioner import MotionerTransformers, FramePackMotioner, rope_precompute
from diffusers.models.attention import AdaLayerNorm
__all__ = ['WanModel']
@@ -191,7 +193,6 @@ def sinusoidal_embedding_1d(dim, position):
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
assert dim % 2 == 0
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
@@ -1302,7 +1303,6 @@ class AudioInjector_WAN(nn.Module):
adain_dim=2048,
need_adain_ont=False):
super().__init__()
num_injector_layers = len(inject_layer)
self.injected_block_id = {}
audio_injector_id = 0
for mod_name, mod in zip(all_modules_names, all_modules):
@@ -1623,6 +1623,48 @@ class WanModel(torch.nn.Module):
self.adain_mode = adain_mode
self.zero_timestep = zero_timestep
# init motioner
enable_framepack = False
enable_motioner = False
add_last_motion = False
if enable_motioner and enable_framepack:
raise ValueError(
"enable_motioner and enable_framepack are mutually exclusive, please set one of them to False"
)
self.enable_motioner = enable_motioner
self.add_last_motion = add_last_motion
if enable_motioner:
motioner_dim = 2048
self.motioner = MotionerTransformers(
patch_size=(2, 4, 4),
dim=motioner_dim,
ffn_dim=motioner_dim,
freq_dim=256,
out_dim=16,
num_heads=16,
num_layers=13,
window_size=(-1, -1),
qk_norm=True,
cross_attn_norm=False,
eps=1e-6,
motion_token_num=4,
enable_tsm=False,
motion_stride=4,
expand_ratio=2,
trainable_token_pos_emb=False,
)
self.zip_motion_out = torch.nn.Sequential(
WanLayerNorm(motioner_dim),
zero_module(nn.Linear(motioner_dim, self.dim)))
self.enable_framepack = enable_framepack
if enable_framepack:
self.frame_packer = FramePackMotioner(
inner_dim=self.dim,
num_heads=self.num_heads,
zip_frame_buckets=[1, 2, 16],
drop_mode='padd')
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
@@ -1773,6 +1815,151 @@ class WanModel(torch.nn.Module):
block.to(self.offload_device, non_blocking=self.use_non_blocking)
return hints
def process_motion(self, motion_latents, drop_motion_frames=False):
if drop_motion_frames or motion_latents[0].shape[1] == 0:
return [], []
self.lat_motion_frames = motion_latents[0].shape[1]
mot = [self.patch_embedding(m.unsqueeze(0)) for m in motion_latents]
batch_size = len(mot)
mot_remb = []
flattern_mot = []
for bs in range(batch_size):
height, width = mot[bs].shape[3], mot[bs].shape[4]
flat_mot = mot[bs].flatten(2).transpose(1, 2).contiguous()
motion_grid_sizes = [[
torch.tensor([-self.lat_motion_frames, 0,
0]).unsqueeze(0).repeat(1, 1),
torch.tensor([0, height, width]).unsqueeze(0).repeat(1, 1),
torch.tensor([self.lat_motion_frames, height,
width]).unsqueeze(0).repeat(1, 1)
]]
motion_rope_emb = rope_precompute(
flat_mot.detach().view(1, flat_mot.shape[1], self.num_heads,
self.dim // self.num_heads),
motion_grid_sizes,
self.freqs,
start=None)
mot_remb.append(motion_rope_emb)
flattern_mot.append(flat_mot)
return flattern_mot, mot_remb
def process_motion_frame_pack(self,
motion_latents,
drop_motion_frames=False,
add_last_motion=2):
flattern_mot, mot_remb = self.frame_packer(motion_latents,
add_last_motion)
if drop_motion_frames:
return [m[:, :0] for m in flattern_mot
], [m[:, :0] for m in mot_remb]
else:
return flattern_mot, mot_remb
def process_motion_transformer_motioner(self,
motion_latents,
drop_motion_frames=False,
add_last_motion=True):
batch_size, height, width = len(
motion_latents), motion_latents[0].shape[2] // self.patch_size[
1], motion_latents[0].shape[3] // self.patch_size[2]
freqs = self.freqs
device = self.patch_embedding.weight.device
if freqs.device != device:
freqs = freqs.to(device)
if self.trainable_token_pos_emb:
token_freqs = self.token_freqs.to(torch.float64)
token_freqs = token_freqs / token_freqs.norm(
dim=-1, keepdim=True)
freqs = [freqs, torch.view_as_complex(token_freqs)]
if not drop_motion_frames and add_last_motion:
last_motion_latent = [u[:, -1:] for u in motion_latents]
last_mot = [
self.patch_embedding(m.unsqueeze(0)) for m in last_motion_latent
]
last_mot = [m.flatten(2).transpose(1, 2) for m in last_mot]
last_mot = torch.cat(last_mot)
gride_sizes = [[
torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([0, height,
width]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([1, height,
width]).unsqueeze(0).repeat(batch_size, 1)
]]
else:
last_mot = torch.zeros([batch_size, 0, self.dim],
device=motion_latents[0].device,
dtype=motion_latents[0].dtype)
gride_sizes = []
zip_motion = self.motioner(motion_latents)
zip_motion = self.zip_motion_out(zip_motion)
if drop_motion_frames:
zip_motion = zip_motion * 0.0
zip_motion_grid_sizes = [[
torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([
0, self.motioner.motion_side_len, self.motioner.motion_side_len
]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor(
[1 if not self.trainable_token_pos_emb else -1, height,
width]).unsqueeze(0).repeat(batch_size, 1),
]]
mot = torch.cat([last_mot, zip_motion], dim=1)
gride_sizes = gride_sizes + zip_motion_grid_sizes
motion_rope_emb = rope_precompute(
mot.detach().view(batch_size, mot.shape[1], self.num_heads,
self.dim // self.num_heads),
gride_sizes,
freqs,
start=None)
return [m.unsqueeze(0) for m in mot
], [r.unsqueeze(0) for r in motion_rope_emb]
def inject_motion(self,
x,
seq_lens,
rope_embs,
mask_input,
motion_latents,
drop_motion_frames=False,
add_last_motion=True):
# inject the motion frames token to the hidden states
if self.enable_motioner:
mot, mot_remb = self.process_motion_transformer_motioner(
motion_latents,
drop_motion_frames=drop_motion_frames,
add_last_motion=add_last_motion)
elif self.enable_framepack:
mot, mot_remb = self.process_motion_frame_pack(
motion_latents,
drop_motion_frames=drop_motion_frames,
add_last_motion=add_last_motion)
else:
mot, mot_remb = self.process_motion(
motion_latents, drop_motion_frames=drop_motion_frames)
if len(mot) > 0:
x = [torch.cat([u, m], dim=1) for u, m in zip(x, mot)]
seq_lens = seq_lens + torch.tensor([r.size(1) for r in mot],
dtype=torch.long)
rope_embs = [
torch.cat([u, m], dim=1) for u, m in zip(rope_embs, mot_remb)
]
mask_input = [
torch.cat([
m, 2 * torch.ones([1, u.shape[1] - m.shape[1]],
device=m.device,
dtype=m.dtype)
],
dim=1) for m, u in zip(mask_input, x)
]
return x, seq_lens, rope_embs, mask_input
def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0):
if block_idx in self.audio_injector.injected_block_id.keys():
@@ -1845,7 +2032,8 @@ class WanModel(torch.nn.Module):
mtv_strength=1.0,
s2v_audio_input=None,
s2v_ref_latent=None,
s2v_audio_scale=1.0
s2v_audio_scale=1.0,
motion_latents=None
):
r"""
@@ -1893,8 +2081,8 @@ class WanModel(torch.nn.Module):
#s2v
if self.model_type == 's2v' and s2v_audio_input is not None:
motion_frames=[17, 5]
#motion_frames=[17, 5]
motion_frames=[1, 0]
s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, motion_frames[0]), s2v_audio_input], dim=-1)
audio_emb_res = self.casual_audio_encoder(s2v_audio_input)
@@ -2108,10 +2296,40 @@ class WanModel(torch.nn.Module):
e0 = torch.cat([
e0.unsqueeze(2),
zero_e0.unsqueeze(2).repeat(e0.size(0), 1, 1, 1)
],
dim=2)
], dim=2)
e0 = [e0, self.original_seq_len]
mask_input = torch.zeros([1, x.shape[1]], dtype=torch.int32, device=x.device)
mask_input[:, self.original_seq_len:] = 1
if motion_latents is not None:
# compute the rope embeddings for the input
#x = torch.cat(x)
b, s, n, d = x.size(0), x.size(
1), self.num_heads, self.dim // self.num_heads
self.pre_compute_freqs = rope_precompute(
x.detach().view(b, s, n, d), grid_sizes, freqs, start=None)
x = [u.unsqueeze(0) for u in x]
self.pre_compute_freqs = [
u.unsqueeze(0) for u in self.pre_compute_freqs
]
x, seq_lens, self.pre_compute_freqs, mask_input = self.inject_motion(
x,
seq_lens,
self.pre_compute_freqs,
mask_input,
motion_latents,
drop_motion_frames=self.drop_motion_frames,
add_last_motion=True)
x = torch.cat(x, dim=0)
self.pre_compute_freqs = torch.cat(self.pre_compute_freqs, dim=0)
mask_input = torch.cat(mask_input, dim=0)
x = x + self.trainable_cond_mask(mask_input).to(x.dtype)
if x_ip is not None:
timestep_ip = torch.zeros_like(t) # [B] with 0s
t_ip = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep_ip.flatten()).to(x.dtype)) # b, dim )
@@ -2503,7 +2721,7 @@ class WanModel(torch.nn.Module):
x = x[:, :x_len]
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
x = x[:, :self.original_seq_len]
#x = x[:, :self.original_seq_len]
x = self.head(x, e.to(x.device))
x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type]
x = [u.float() for u in x]
+31 -105
View File
@@ -1,16 +1,16 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import math
from typing import Any, Dict, List, Literal, Optional, Union
from typing import Any, Dict
import numpy as np
import torch
import torch.cuda.amp as amp
import torch.nn as nn
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.utils import BaseOutput, is_torch_version
from diffusers.loaders import PeftAdapterMixin
from diffusers.utils import BaseOutput
from einops import rearrange, repeat
from ..model import flash_attention
from ..attention import attention
from .s2v_utils import rope_precompute
@@ -172,12 +172,11 @@ class SelfAttention(nn.Module):
q, k, v = qkv_fn(x)
x = flash_attention(
x = attention(
q=rope_apply(q, grid_sizes, freqs),
k=rope_apply(k, grid_sizes, freqs),
v=v,
k_lens=seq_lens,
window_size=self.window_size)
k_lens=seq_lens)
# output
x = x.flatten(2)
@@ -224,12 +223,7 @@ class SwinSelfAttention(SelfAttention):
# q: b (t h w) n d
# k: b (t h w) n d
out = flash_attention(
q=q,
k=k,
v=v,
# k_lens=torch.tensor([k.shape[1]] * k.shape[0], device=x.device, dtype=torch.long),
window_size=self.window_size)
out = attention(q, k, v)
out = torch.cat([out, ref_v[:1]], axis=0)
out = rearrange(out, '(b t) (h w) n d -> b (t h w) n d', t=T, h=H, w=W)
x = out
@@ -308,7 +302,7 @@ class CasualSelfAttention(SelfAttention):
# k: b (t h w) n d
outs = []
for i in range(q.shape[0]):
out = flash_attention(
out = attention(
q=q[i:i + 1],
k=k[i:i + 1],
v=v[i:i + 1],
@@ -479,10 +473,8 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin):
gride_sizes = [[
torch.tensor([0, 0, 0]).unsqueeze(0).repeat(1, 1),
torch.tensor([1, self.motion_side_len,
self.motion_side_len]).unsqueeze(0).repeat(1, 1),
torch.tensor([1, self.motion_side_len,
self.motion_side_len]).unsqueeze(0).repeat(1, 1),
torch.tensor([1, self.motion_side_len, self.motion_side_len]).unsqueeze(0).repeat(1, 1),
torch.tensor([1, self.motion_side_len, self.motion_side_len]).unsqueeze(0).repeat(1, 1),
]]
token_freqs = rope_apply(x, gride_sizes, self.freqs)
token_freqs = token_freqs[0, :, 0].reshape(motion_token_num, -1, 2)
@@ -502,7 +494,6 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin):
context: A list of text embeddings each with shape [L, C].
"""
# params
motion_frames = x[0].shape[1]
device = self.patch_embedding.weight.device
freqs = self.freqs
if freqs.device != device:
@@ -535,22 +526,16 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin):
seq_f, seq_h, seq_w = x[0].shape[-3:]
batch_size = len(x)
if not self.enable_tsm:
grid_sizes = torch.stack(
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
grid_sizes = [[
torch.zeros_like(grid_sizes), grid_sizes, grid_sizes
]]
grid_sizes = torch.stack([torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
grid_sizes = [[torch.zeros_like(grid_sizes), grid_sizes, grid_sizes]]
seq_f = 0
else:
grid_sizes = []
for idx in sample_idx[0][::-1][::self.sample_c]:
tsm_frame_grid_sizes = [[
torch.tensor([idx, 0,
0]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([idx + 1, seq_h,
seq_w]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([1, seq_h,
seq_w]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([idx, 0, 0]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([idx + 1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1),
]]
grid_sizes += tsm_frame_grid_sizes
seq_f = sample_idx[0][-1] + 1
@@ -563,14 +548,10 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin):
token_grid_sizes = [[
torch.tensor([seq_f, 0, 0]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor(
[seq_f + 1, self.motion_side_len,
self.motion_side_len]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor(
[1 if not self.trainable_token_pos_emb else -1, seq_h,
seq_w]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([seq_f + 1, self.motion_side_len,self.motion_side_len]).unsqueeze(0).repeat(batch_size, 1),
torch.tensor([1 if not self.trainable_token_pos_emb else -1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1),
] # 第三行代表rope emb的想要覆盖到的范围
]
]
grid_sizes = grid_sizes + token_grid_sizes
token_unpatch_grid_sizes = torch.stack([
@@ -589,31 +570,8 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin):
freqs=freqs,
)
# grad ckpt args
def create_custom_forward(module, return_dict=None):
def custom_forward(*inputs, **kwargs):
if return_dict is not None:
return module(*inputs, **kwargs, return_dict=return_dict)
else:
return module(*inputs, **kwargs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = ({
"use_reentrant": False
} if is_torch_version(">=", "1.11.0") else {})
for idx, block in enumerate(self.blocks):
if self.training and self.gradient_checkpointing:
x = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
x,
**kwargs,
**ckpt_kwargs,
)
else:
x = block(x, **kwargs)
x = block(x, **kwargs)
# head
out = x[:, -token_len:]
return out
@@ -628,18 +586,6 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin):
out.append(u)
return out
def init_weights(self):
# basic init
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.zeros_(m.bias)
# init embeddings
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
class FramePackMotioner(nn.Module):
def __init__(
@@ -677,7 +623,6 @@ class FramePackMotioner(nn.Module):
self.drop_mode = drop_mode
def forward(self, motion_latents, add_last_motion=2):
motion_frames = motion_latents[0].shape[1]
mot = []
mot_remb = []
for m in motion_latents:
@@ -702,21 +647,15 @@ class FramePackMotioner(nn.Module):
list(self.zip_frame_buckets)[::-1], dim=2) # 16, 2 ,1
# patchfy
clean_latents_post = self.proj(clean_latents_post).flatten(
2).transpose(1, 2)
clean_latents_2x = self.proj_2x(clean_latents_2x).flatten(
2).transpose(1, 2)
clean_latents_4x = self.proj_4x(clean_latents_4x).flatten(
2).transpose(1, 2)
clean_latents_post = self.proj(clean_latents_post).flatten(2).transpose(1, 2)
clean_latents_2x = self.proj_2x(clean_latents_2x).flatten(2).transpose(1, 2)
clean_latents_4x = self.proj_4x(clean_latents_4x).flatten(2).transpose(1, 2)
if add_last_motion < 2 and self.drop_mode == "drop":
clean_latents_post = clean_latents_post[:, :
0] if add_last_motion < 2 else clean_latents_post
clean_latents_2x = clean_latents_2x[:, :
0] if add_last_motion < 1 else clean_latents_2x
clean_latents_post = clean_latents_post[:, :0] if add_last_motion < 2 else clean_latents_post
clean_latents_2x = clean_latents_2x[:, :0] if add_last_motion < 1 else clean_latents_2x
motion_lat = torch.cat(
[clean_latents_post, clean_latents_2x, clean_latents_4x], dim=1)
motion_lat = torch.cat([clean_latents_post, clean_latents_2x, clean_latents_4x], dim=1)
# rope
start_time_id = -(self.zip_frame_buckets[:1].sum())
@@ -731,20 +670,18 @@ class FramePackMotioner(nn.Module):
start_time_id = -(self.zip_frame_buckets[:2].sum())
end_time_id = start_time_id + self.zip_frame_buckets[1] // 2
grid_sizes_2x = [] if add_last_motion < 1 and self.drop_mode == "drop" else \
[
[torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1),
[[
torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1),
torch.tensor([end_time_id, lat_height // 4, lat_width // 4]).unsqueeze(0).repeat(1, 1),
torch.tensor([self.zip_frame_buckets[1], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), ]
]
torch.tensor([self.zip_frame_buckets[1], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1),
]]
start_time_id = -(self.zip_frame_buckets[:3].sum())
end_time_id = start_time_id + self.zip_frame_buckets[2] // 4
grid_sizes_4x = [[
torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1),
torch.tensor([end_time_id, lat_height // 8,
lat_width // 8]).unsqueeze(0).repeat(1, 1),
torch.tensor([
self.zip_frame_buckets[2], lat_height // 2, lat_width // 2
torch.tensor([end_time_id, lat_height // 8, lat_width // 8]).unsqueeze(0).repeat(1, 1),
torch.tensor([self.zip_frame_buckets[2], lat_height // 2, lat_width // 2
]).unsqueeze(0).repeat(1, 1),
]]
@@ -781,14 +718,3 @@ def sample_indices(N, stride, expand_ratio, c):
return indices
if __name__ == '__main__':
device = "cuda"
model = FramePackMotioner(inner_dim=1024)
batch_size = 2
num_frame, height, width = (28, 32, 32)
single_input = torch.ones([16, num_frame, height, width], device=device)
for i in range(num_frame):
single_input[:, num_frame - 1 - i] *= i
x = [single_input] * batch_size
model.forward(x)