continue
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user