diff --git a/nodes.py b/nodes.py index 359fbe2..dd70eb8 100644 --- a/nodes.py +++ b/nodes.py @@ -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, diff --git a/s2v/nodes.py b/s2v/nodes.py index 6c840b7..9bd8ab1 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -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 = { diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 88b35c9..03b1362 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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] diff --git a/wanvideo/modules/s2v/motioner.py b/wanvideo/modules/s2v/motioner.py index 699c570..043e476 100644 --- a/wanvideo/modules/s2v/motioner.py +++ b/wanvideo/modules/s2v/motioner.py @@ -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)