init
This commit is contained in:
@@ -12,6 +12,7 @@ from .nodes_model_loading import NODE_CLASS_MAPPINGS as MODEL_LOADING_NODE_CLASS
|
||||
from .nodes_utility import NODE_CLASS_MAPPINGS as UTILITY_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UTILITY_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODE_CACHE_DISPLAY_NAME_MAPPINGS
|
||||
from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as S2V_NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
try:
|
||||
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
||||
@@ -58,6 +59,7 @@ NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
@@ -75,5 +77,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -2212,6 +2212,16 @@ class WanVideoSampler:
|
||||
log.info(f"mtv_motion_rotary_emb: {motion_rotary_emb[0].shape}")
|
||||
mtv_freqs = mtv_freqs.to(device, dtype)
|
||||
|
||||
#region S2V
|
||||
s2v_audio_input = None
|
||||
s2v_audio_embeds = image_embeds.get("audio_embeds", None)
|
||||
if s2v_audio_embeds is not None:
|
||||
log.info(f"Using S2V audio embeddings")
|
||||
s2v_audio_input = s2v_audio_embeds["audio_embed_bucket"].to(device, dtype)
|
||||
#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)
|
||||
|
||||
# vid2vid
|
||||
noise_mask=original_image=None
|
||||
if samples is not None and not multitalk_sampling:
|
||||
@@ -2669,6 +2679,7 @@ class WanVideoSampler:
|
||||
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None, # MTV-Crafter RoPE
|
||||
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, # MTV-Crafter scaling
|
||||
"mtv_freqs": mtv_freqs if mtv_input is not None else None, # MTV-Crafter extra RoPE freqs
|
||||
"s2v_audio_input": s2v_audio_input #official speech-to-video
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
|
||||
@@ -731,7 +731,7 @@ class WanVideoSetLoRAs:
|
||||
|
||||
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
|
||||
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
|
||||
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio"}
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
@@ -1079,7 +1079,9 @@ class WanVideoModelLoader:
|
||||
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
|
||||
|
||||
model_type = "t2v"
|
||||
if not "text_embedding.0.weight" in sd:
|
||||
if "audio_injector.injector.0.k.weight" in sd:
|
||||
model_type = "s2v"
|
||||
elif not "text_embedding.0.weight" in sd:
|
||||
model_type = "no_cross_attn" #minimaxremover
|
||||
elif "model_type.Wan2_1-FLF2V-14B-720P" in sd or "img_emb.emb_pos" in sd or "flf2v" in model.lower():
|
||||
model_type = "fl2v"
|
||||
@@ -1186,7 +1188,10 @@ class WanVideoModelLoader:
|
||||
"add_ref_conv": True if "ref_conv.weight" in sd else False,
|
||||
"in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None,
|
||||
"add_control_adapter": True if "control_adapter.conv.weight" in sd else False,
|
||||
"use_motion_attn": True if "blocks.0.motion_attn.k.weight" in sd else False
|
||||
"use_motion_attn": True if "blocks.0.motion_attn.k.weight" in sd else False,
|
||||
"enable_adain": True if "audio_injector.injector_adain_layers.0.linear.weight" in sd else False,
|
||||
"cond_dim": sd["cond_encoder.weight"].shape[1] if "cond_encoder.weight" in sd else 0
|
||||
|
||||
}
|
||||
|
||||
with init_empty_weights():
|
||||
|
||||
+164
@@ -0,0 +1,164 @@
|
||||
import folder_paths
|
||||
import math
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
def get_sample_indices(original_fps,
|
||||
total_frames,
|
||||
target_fps,
|
||||
num_sample,
|
||||
fixed_start=None):
|
||||
required_duration = num_sample / target_fps
|
||||
required_origin_frames = int(np.ceil(required_duration * original_fps))
|
||||
if required_duration > total_frames / original_fps:
|
||||
raise ValueError("required_duration must be less than video length")
|
||||
|
||||
if not fixed_start is None and fixed_start >= 0:
|
||||
start_frame = fixed_start
|
||||
else:
|
||||
max_start = total_frames - required_origin_frames
|
||||
if max_start < 0:
|
||||
raise ValueError("video length is too short")
|
||||
start_frame = np.random.randint(0, max_start + 1)
|
||||
start_time = start_frame / original_fps
|
||||
|
||||
end_time = start_time + required_duration
|
||||
time_points = np.linspace(start_time, end_time, num_sample, endpoint=False)
|
||||
|
||||
frame_indices = np.round(np.array(time_points) * original_fps).astype(int)
|
||||
frame_indices = np.clip(frame_indices, 0, total_frames - 1)
|
||||
return frame_indices
|
||||
|
||||
def linear_interpolation(features, input_fps, output_fps, output_len=None):
|
||||
"""
|
||||
features: shape=[1, T, 512]
|
||||
input_fps: fps for audio, f_a
|
||||
output_fps: fps for video, f_m
|
||||
output_len: video length
|
||||
"""
|
||||
features = features.transpose(1, 2) # [1, 512, T]
|
||||
seq_len = features.shape[2] / float(input_fps) # T/f_a
|
||||
if output_len is None:
|
||||
output_len = int(seq_len * output_fps) # f_m*T/f_a
|
||||
output_features = F.interpolate(
|
||||
features, size=output_len, align_corners=True,
|
||||
mode='linear') # [1, 512, output_len]
|
||||
return output_features.transpose(1, 2) # [1, output_len, 512]
|
||||
|
||||
class WanVideoAddAudioEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",),
|
||||
"input_fps": ("FLOAT", {"default": 30.0, "min": 1.0, "max": 120.0, "step": 1.0, "tooltip": "Frames per second for the audio"}),
|
||||
"output_fps": ("FLOAT", {"default": 30.0, "min": 1.0, "max": 120.0, "step": 1.0, "tooltip": "Frames per second for the video"}),
|
||||
"frames": ("INT", {"default": 81, "min": 1, "max": 120, "step": 1, "tooltip": "Number of frames to process"})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, input_fps, output_fps, frames, audio_encoder_output):
|
||||
# Prepare the new audio entry
|
||||
|
||||
#audio_feat = audio_encoder_output["encoded_audio"]
|
||||
#print("audio_feat", audio_feat.shape)
|
||||
all_layers = audio_encoder_output["encoded_audio_all_layers"]
|
||||
audio_feat = torch.stack(all_layers, dim=0).squeeze(1) # shape: [num_layers, T, 512]
|
||||
|
||||
print("audio_feat", audio_feat.shape)
|
||||
|
||||
if input_fps != output_fps:
|
||||
audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps)
|
||||
|
||||
self.video_rate = output_fps
|
||||
|
||||
audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps(
|
||||
audio_feat,
|
||||
fps=output_fps,
|
||||
batch_frames=frames
|
||||
)
|
||||
|
||||
audio_embed_bucket = audio_embed_bucket.unsqueeze(0)
|
||||
if len(audio_embed_bucket.shape) == 3:
|
||||
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 1)
|
||||
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 = {
|
||||
"audio_embed_bucket": audio_embed_bucket,
|
||||
"num_repeat": num_repeat
|
||||
}
|
||||
updated = dict(embeds)
|
||||
updated["audio_embeds"] = new_entry
|
||||
return (updated,)
|
||||
|
||||
def get_audio_embed_bucket_fps(self, audio_embed, fps=16, batch_frames=81, m=0):
|
||||
num_layers, audio_frame_num, audio_dim = audio_embed.shape
|
||||
|
||||
if num_layers > 1:
|
||||
return_all_layers = True
|
||||
else:
|
||||
return_all_layers = False
|
||||
|
||||
scale = self.video_rate / fps
|
||||
|
||||
min_batch_num = int(audio_frame_num / (batch_frames * scale)) + 1
|
||||
|
||||
bucket_num = min_batch_num * batch_frames
|
||||
padd_audio_num = math.ceil(min_batch_num * batch_frames / fps *
|
||||
self.video_rate) - audio_frame_num
|
||||
batch_idx = get_sample_indices(
|
||||
original_fps=self.video_rate,
|
||||
total_frames=audio_frame_num + padd_audio_num,
|
||||
target_fps=fps,
|
||||
num_sample=bucket_num,
|
||||
fixed_start=0)
|
||||
batch_audio_eb = []
|
||||
audio_sample_stride = int(self.video_rate / fps)
|
||||
for bi in batch_idx:
|
||||
if bi < audio_frame_num:
|
||||
|
||||
chosen_idx = list(
|
||||
range(bi - m * audio_sample_stride,
|
||||
bi + (m + 1) * audio_sample_stride,
|
||||
audio_sample_stride))
|
||||
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
|
||||
chosen_idx = [
|
||||
audio_frame_num - 1 if c >= audio_frame_num else c
|
||||
for c in chosen_idx
|
||||
]
|
||||
|
||||
if return_all_layers:
|
||||
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
|
||||
start_dim=-2, end_dim=-1)
|
||||
else:
|
||||
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
|
||||
else:
|
||||
frame_audio_embed = \
|
||||
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
|
||||
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
|
||||
batch_audio_eb.append(frame_audio_embed)
|
||||
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
|
||||
dim=0)
|
||||
|
||||
return batch_audio_eb, min_batch_num
|
||||
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddAudioEmbeds": WanVideoAddAudioEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddAudioEmbeds": "WanVideo Add Audio Embeds",
|
||||
}
|
||||
+255
-34
@@ -30,11 +30,38 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
|
||||
|
||||
from ...MTV.mtv import apply_rotary_emb
|
||||
|
||||
from diffusers.models.attention import AdaLayerNorm
|
||||
|
||||
__all__ = ['WanModel']
|
||||
|
||||
from comfy import model_management as mm
|
||||
|
||||
|
||||
def zero_module(module):
|
||||
"""
|
||||
Zero out the parameters of a module and return it.
|
||||
"""
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
def torch_dfs(model: nn.Module, parent_name='root'):
|
||||
module_names, modules = [], []
|
||||
current_name = parent_name if parent_name else 'root'
|
||||
module_names.append(current_name)
|
||||
modules.append(model)
|
||||
|
||||
for name, child in model.named_children():
|
||||
if parent_name:
|
||||
child_name = f'{parent_name}.{name}'
|
||||
else:
|
||||
child_name = name
|
||||
child_modules, child_names = torch_dfs(child, child_name)
|
||||
module_names += child_names
|
||||
modules += child_modules
|
||||
return modules, module_names
|
||||
|
||||
#from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
|
||||
def apply_rope_comfy(xq, xk, freqs_cis):
|
||||
xq_ = xq.to(dtype=freqs_cis.dtype).reshape(*xq.shape[:-1], -1, 1, 2)
|
||||
@@ -1175,41 +1202,143 @@ class MLPProj(torch.nn.Module):
|
||||
clip_extra_context_tokens = self.proj(image_embeds)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
from .s2v.auxi_blocks import MotionEncoder_tc
|
||||
|
||||
|
||||
class CausalAudioEncoder(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim=5120,
|
||||
num_layers=25,
|
||||
out_dim=2048,
|
||||
video_rate=8,
|
||||
num_token=4,
|
||||
need_global=False):
|
||||
super().__init__()
|
||||
self.encoder = MotionEncoder_tc(
|
||||
in_dim=dim,
|
||||
hidden_dim=out_dim,
|
||||
num_heads=num_token,
|
||||
need_global=need_global)
|
||||
weight = torch.ones((1, num_layers, 1, 1)) * 0.01
|
||||
|
||||
self.weights = torch.nn.Parameter(weight)
|
||||
self.act = torch.nn.SiLU()
|
||||
|
||||
def forward(self, features):
|
||||
# features B * num_layers * dim * video_length
|
||||
weights = self.act(self.weights)
|
||||
weights_sum = weights.sum(dim=1, keepdims=True)
|
||||
weighted_feat = ((features * weights) / weights_sum).sum(
|
||||
dim=1) # b dim f
|
||||
weighted_feat = weighted_feat.permute(0, 2, 1) # b f dim
|
||||
res = self.encoder(weighted_feat) # b f n dim
|
||||
|
||||
return res # b f n dim
|
||||
|
||||
|
||||
class AudioCrossAttention(WanT2VCrossAttention):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
|
||||
class AudioInjector_WAN(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
all_modules,
|
||||
all_modules_names,
|
||||
dim=2048,
|
||||
num_heads=32,
|
||||
inject_layer=[0, 27],
|
||||
root_net=None,
|
||||
enable_adain=False,
|
||||
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):
|
||||
if isinstance(mod, WanAttentionBlock):
|
||||
for inject_id in inject_layer:
|
||||
if f'transformer_blocks.{inject_id}' in mod_name:
|
||||
self.injected_block_id[inject_id] = audio_injector_id
|
||||
audio_injector_id += 1
|
||||
|
||||
self.injector = nn.ModuleList([
|
||||
AudioCrossAttention(
|
||||
in_features=dim,
|
||||
out_features=dim,
|
||||
num_heads=num_heads,
|
||||
qk_norm=True,
|
||||
) for _ in range(audio_injector_id)
|
||||
])
|
||||
self.injector_pre_norm_feat = nn.ModuleList([
|
||||
nn.LayerNorm(
|
||||
dim,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
) for _ in range(audio_injector_id)
|
||||
])
|
||||
self.injector_pre_norm_vec = nn.ModuleList([
|
||||
nn.LayerNorm(
|
||||
dim,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
) for _ in range(audio_injector_id)
|
||||
])
|
||||
if enable_adain:
|
||||
self.injector_adain_layers = nn.ModuleList([
|
||||
AdaLayerNorm(
|
||||
output_dim=dim * 2, embedding_dim=adain_dim, chunk_dim=1)
|
||||
for _ in range(audio_injector_id)
|
||||
])
|
||||
if need_adain_ont:
|
||||
self.injector_adain_output_layers = nn.ModuleList(
|
||||
[nn.Linear(dim, dim) for _ in range(audio_injector_id)])
|
||||
|
||||
class WanModel(torch.nn.Module):
|
||||
def __init__(self,
|
||||
model_type='t2v',
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=512,
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
in_features=5120,
|
||||
out_features=5120,
|
||||
ffn_dim=8192,
|
||||
ffn2_dim=8192,
|
||||
freq_dim=256,
|
||||
text_dim=4096,
|
||||
out_dim=16,
|
||||
num_heads=16,
|
||||
num_layers=32,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func='comfy',
|
||||
main_device=torch.device('cuda'),
|
||||
offload_device=torch.device('cpu'),
|
||||
teacache_coefficients=[],
|
||||
magcache_ratios=[],
|
||||
vace_layers=None,
|
||||
vace_in_dim=None,
|
||||
inject_sample_info=False,
|
||||
add_ref_conv=False,
|
||||
in_dim_ref_conv=16,
|
||||
add_control_adapter=False,
|
||||
in_dim_control_adapter=24,
|
||||
use_motion_attn=False
|
||||
):
|
||||
model_type='t2v',
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=512,
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
in_features=5120,
|
||||
out_features=5120,
|
||||
ffn_dim=8192,
|
||||
ffn2_dim=8192,
|
||||
freq_dim=256,
|
||||
text_dim=4096,
|
||||
out_dim=16,
|
||||
num_heads=16,
|
||||
num_layers=32,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func='comfy',
|
||||
main_device=torch.device('cuda'),
|
||||
offload_device=torch.device('cpu'),
|
||||
teacache_coefficients=[],
|
||||
magcache_ratios=[],
|
||||
vace_layers=None,
|
||||
vace_in_dim=None,
|
||||
inject_sample_info=False,
|
||||
add_ref_conv=False,
|
||||
in_dim_ref_conv=16,
|
||||
add_control_adapter=False,
|
||||
in_dim_control_adapter=24,
|
||||
use_motion_attn=False,
|
||||
#s2v
|
||||
cond_dim=0,
|
||||
audio_dim=1024,
|
||||
num_audio_token=4,
|
||||
enable_adain=False,
|
||||
adain_mode="attn_norm",
|
||||
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
|
||||
@@ -1360,7 +1489,7 @@ class WanModel(torch.nn.Module):
|
||||
])
|
||||
else:
|
||||
# blocks
|
||||
if model_type == 't2v':
|
||||
if model_type == 't2v' or model_type == 's2v':
|
||||
cross_attn_type = 't2v_cross_attn'
|
||||
elif model_type == 'i2v' or model_type == 'fl2v':
|
||||
cross_attn_type = 'i2v_cross_attn'
|
||||
@@ -1415,6 +1544,36 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
self.block_mask=None
|
||||
|
||||
#S2V
|
||||
if cond_dim > 0:
|
||||
self.cond_encoder = nn.Conv3d(
|
||||
cond_dim,
|
||||
self.dim,
|
||||
kernel_size=self.patch_size,
|
||||
stride=self.patch_size)
|
||||
self.enable_adain = enable_adain
|
||||
self.casual_audio_encoder = CausalAudioEncoder(
|
||||
dim=audio_dim,
|
||||
out_dim=self.dim,
|
||||
num_token=num_audio_token,
|
||||
need_global=enable_adain)
|
||||
all_modules, all_modules_names = torch_dfs(
|
||||
self.blocks, parent_name="root.transformer_blocks")
|
||||
self.audio_injector = AudioInjector_WAN(
|
||||
all_modules,
|
||||
all_modules_names,
|
||||
dim=self.dim,
|
||||
num_heads=self.num_heads,
|
||||
inject_layer=audio_inject_layers,
|
||||
root_net=self,
|
||||
enable_adain=enable_adain,
|
||||
adain_dim=self.dim,
|
||||
need_adain_ont=adain_mode != "attn_norm",
|
||||
)
|
||||
self.adain_mode = adain_mode
|
||||
|
||||
self.trainable_cond_mask = nn.Embedding(3, self.dim)
|
||||
|
||||
@staticmethod
|
||||
def _prepare_blockwise_causal_attn_mask(
|
||||
device: torch.device | str, num_frames: int = 21,
|
||||
@@ -1565,6 +1724,46 @@ class WanModel(torch.nn.Module):
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
return hints
|
||||
|
||||
def audio_injector_forward(self, block_idx, hidden_states, merged_audio_emb):
|
||||
if block_idx in self.audio_injector.injected_block_id.keys():
|
||||
audio_attn_id = self.audio_injector.injected_block_id[block_idx]
|
||||
audio_emb = merged_audio_emb # b f n c
|
||||
num_frames = audio_emb.shape[1]
|
||||
|
||||
input_hidden_states = hidden_states[:, :self.original_seq_len].clone() # b (f h w) c
|
||||
input_hidden_states = rearrange(
|
||||
input_hidden_states, "b (t n) c -> (b t) n c", t=num_frames)
|
||||
|
||||
if self.enable_adain and self.adain_mode == "attn_norm":
|
||||
audio_emb_global = self.audio_emb_global
|
||||
audio_emb_global = rearrange(audio_emb_global,
|
||||
"b t n c -> (b t) n c")
|
||||
adain_hidden_states = self.audio_injector.injector_adain_layers[
|
||||
audio_attn_id](
|
||||
input_hidden_states, temb=audio_emb_global[:, 0])
|
||||
attn_hidden_states = adain_hidden_states
|
||||
else:
|
||||
attn_hidden_states = self.audio_injector.injector_pre_norm_feat[
|
||||
audio_attn_id](
|
||||
input_hidden_states)
|
||||
audio_emb = rearrange(
|
||||
audio_emb, "b t n c -> (b t) n c", t=num_frames)
|
||||
attn_audio_emb = audio_emb
|
||||
residual_out = self.audio_injector.injector[audio_attn_id](
|
||||
x=attn_hidden_states,
|
||||
context=attn_audio_emb,
|
||||
context_lens=torch.ones(
|
||||
attn_hidden_states.shape[0],
|
||||
dtype=torch.long,
|
||||
device=attn_hidden_states.device) * attn_audio_emb.shape[1])
|
||||
residual_out = rearrange(
|
||||
residual_out, "(b t) n c -> b (t n) c", t=num_frames)
|
||||
hidden_states[:, :self.
|
||||
original_seq_len] = hidden_states[:, :self.
|
||||
original_seq_len] + residual_out
|
||||
|
||||
return hidden_states
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -1610,6 +1809,7 @@ class WanModel(torch.nn.Module):
|
||||
mtv_motion_rotary_emb=None,
|
||||
mtv_freqs=None,
|
||||
mtv_strength=1.0,
|
||||
s2v_audio_input=None
|
||||
|
||||
):
|
||||
r"""
|
||||
@@ -1654,8 +1854,25 @@ class WanModel(torch.nn.Module):
|
||||
if isinstance(submodule, nn.Linear):
|
||||
if hasattr(submodule, 'step'):
|
||||
submodule.step = current_step
|
||||
|
||||
#s2v
|
||||
if self.model_type == 's2v' and s2v_audio_input is not None:
|
||||
motion_frames=[17, 5]
|
||||
|
||||
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)
|
||||
if self.enable_adain:
|
||||
audio_emb_global, audio_emb = audio_emb_res
|
||||
self.audio_emb_global = audio_emb_global[:, motion_frames[1]:].clone()
|
||||
else:
|
||||
audio_emb = audio_emb_res
|
||||
merged_audio_emb = audio_emb[:, motion_frames[1]:, :]
|
||||
|
||||
|
||||
# params
|
||||
device = self.patch_embedding.weight.device
|
||||
|
||||
if freqs is not None and freqs.device != device:
|
||||
freqs = freqs.to(device)
|
||||
|
||||
@@ -1737,7 +1954,9 @@ class WanModel(torch.nn.Module):
|
||||
F += phantom_ref_frames
|
||||
x = [torch.concat([u, phantom_ref.unsqueeze(0)], dim=1) for phantom_ref, u in zip(phantom_ref, x)]
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.float32)
|
||||
self.original_seq_len = x[0].size(1)
|
||||
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
||||
@@ -2163,6 +2382,8 @@ class WanModel(torch.nn.Module):
|
||||
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
||||
continue
|
||||
x, x_ip = block(x, x_ip=x_ip, **kwargs) #run block
|
||||
if self.audio_injector is not None and s2v_audio_input is not None:
|
||||
x = self.audio_injector_forward(b, x, merged_audio_emb) #s2v
|
||||
if self.block_swap_debug:
|
||||
compute_end = time.perf_counter()
|
||||
compute_time = compute_end - compute_start
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import math
|
||||
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor
|
||||
|
||||
|
||||
def get_sample_indices(original_fps,
|
||||
total_frames,
|
||||
target_fps,
|
||||
num_sample,
|
||||
fixed_start=None):
|
||||
required_duration = num_sample / target_fps
|
||||
required_origin_frames = int(np.ceil(required_duration * original_fps))
|
||||
if required_duration > total_frames / original_fps:
|
||||
raise ValueError("required_duration must be less than video length")
|
||||
|
||||
if not fixed_start is None and fixed_start >= 0:
|
||||
start_frame = fixed_start
|
||||
else:
|
||||
max_start = total_frames - required_origin_frames
|
||||
if max_start < 0:
|
||||
raise ValueError("video length is too short")
|
||||
start_frame = np.random.randint(0, max_start + 1)
|
||||
start_time = start_frame / original_fps
|
||||
|
||||
end_time = start_time + required_duration
|
||||
time_points = np.linspace(start_time, end_time, num_sample, endpoint=False)
|
||||
|
||||
frame_indices = np.round(np.array(time_points) * original_fps).astype(int)
|
||||
frame_indices = np.clip(frame_indices, 0, total_frames - 1)
|
||||
return frame_indices
|
||||
|
||||
|
||||
def linear_interpolation(features, input_fps, output_fps, output_len=None):
|
||||
"""
|
||||
features: shape=[1, T, 512]
|
||||
input_fps: fps for audio, f_a
|
||||
output_fps: fps for video, f_m
|
||||
output_len: video length
|
||||
"""
|
||||
features = features.transpose(1, 2) # [1, 512, T]
|
||||
seq_len = features.shape[2] / float(input_fps) # T/f_a
|
||||
if output_len is None:
|
||||
output_len = int(seq_len * output_fps) # f_m*T/f_a
|
||||
output_features = F.interpolate(
|
||||
features, size=output_len, align_corners=True,
|
||||
mode='linear') # [1, 512, output_len]
|
||||
return output_features.transpose(1, 2) # [1, output_len, 512]
|
||||
|
||||
|
||||
class AudioEncoder():
|
||||
|
||||
def __init__(self, device='cpu', model_id="facebook/wav2vec2-base-960h"):
|
||||
# load pretrained model
|
||||
self.processor = Wav2Vec2Processor.from_pretrained(model_id)
|
||||
self.model = Wav2Vec2ForCTC.from_pretrained(model_id)
|
||||
|
||||
self.model = self.model.to(device)
|
||||
|
||||
self.video_rate = 30
|
||||
|
||||
def extract_audio_feat(self,
|
||||
audio_path,
|
||||
return_all_layers=False,
|
||||
dtype=torch.float32):
|
||||
audio_input, sample_rate = librosa.load(audio_path, sr=16000)
|
||||
|
||||
input_values = self.processor(
|
||||
audio_input, sampling_rate=sample_rate,
|
||||
return_tensors="pt").input_values
|
||||
|
||||
# INFERENCE
|
||||
|
||||
# retrieve logits & take argmax
|
||||
res = self.model(
|
||||
input_values.to(self.model.device), output_hidden_states=True)
|
||||
if return_all_layers:
|
||||
feat = torch.cat(res.hidden_states)
|
||||
else:
|
||||
feat = res.hidden_states[-1]
|
||||
feat = linear_interpolation(
|
||||
feat, input_fps=50, output_fps=self.video_rate)
|
||||
|
||||
z = feat.to(dtype) # Encoding for the motion
|
||||
return z
|
||||
|
||||
def get_audio_embed_bucket(self,
|
||||
audio_embed,
|
||||
stride=2,
|
||||
batch_frames=12,
|
||||
m=2):
|
||||
num_layers, audio_frame_num, audio_dim = audio_embed.shape
|
||||
|
||||
if num_layers > 1:
|
||||
return_all_layers = True
|
||||
else:
|
||||
return_all_layers = False
|
||||
|
||||
min_batch_num = int(audio_frame_num / (batch_frames * stride)) + 1
|
||||
|
||||
bucket_num = min_batch_num * batch_frames
|
||||
batch_idx = [stride * i for i in range(bucket_num)]
|
||||
batch_audio_eb = []
|
||||
for bi in batch_idx:
|
||||
if bi < audio_frame_num:
|
||||
audio_sample_stride = 2
|
||||
chosen_idx = list(
|
||||
range(bi - m * audio_sample_stride,
|
||||
bi + (m + 1) * audio_sample_stride,
|
||||
audio_sample_stride))
|
||||
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
|
||||
chosen_idx = [
|
||||
audio_frame_num - 1 if c >= audio_frame_num else c
|
||||
for c in chosen_idx
|
||||
]
|
||||
|
||||
if return_all_layers:
|
||||
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
|
||||
start_dim=-2, end_dim=-1)
|
||||
else:
|
||||
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
|
||||
else:
|
||||
frame_audio_embed = \
|
||||
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
|
||||
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
|
||||
batch_audio_eb.append(frame_audio_embed)
|
||||
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
|
||||
dim=0)
|
||||
|
||||
return batch_audio_eb, min_batch_num
|
||||
|
||||
def get_audio_embed_bucket_fps(self,
|
||||
audio_embed,
|
||||
fps=16,
|
||||
batch_frames=81,
|
||||
m=0):
|
||||
num_layers, audio_frame_num, audio_dim = audio_embed.shape
|
||||
|
||||
if num_layers > 1:
|
||||
return_all_layers = True
|
||||
else:
|
||||
return_all_layers = False
|
||||
|
||||
scale = self.video_rate / fps
|
||||
|
||||
min_batch_num = int(audio_frame_num / (batch_frames * scale)) + 1
|
||||
|
||||
bucket_num = min_batch_num * batch_frames
|
||||
padd_audio_num = math.ceil(min_batch_num * batch_frames / fps *
|
||||
self.video_rate) - audio_frame_num
|
||||
batch_idx = get_sample_indices(
|
||||
original_fps=self.video_rate,
|
||||
total_frames=audio_frame_num + padd_audio_num,
|
||||
target_fps=fps,
|
||||
num_sample=bucket_num,
|
||||
fixed_start=0)
|
||||
batch_audio_eb = []
|
||||
audio_sample_stride = int(self.video_rate / fps)
|
||||
for bi in batch_idx:
|
||||
if bi < audio_frame_num:
|
||||
|
||||
chosen_idx = list(
|
||||
range(bi - m * audio_sample_stride,
|
||||
bi + (m + 1) * audio_sample_stride,
|
||||
audio_sample_stride))
|
||||
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
|
||||
chosen_idx = [
|
||||
audio_frame_num - 1 if c >= audio_frame_num else c
|
||||
for c in chosen_idx
|
||||
]
|
||||
|
||||
if return_all_layers:
|
||||
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
|
||||
start_dim=-2, end_dim=-1)
|
||||
else:
|
||||
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
|
||||
else:
|
||||
frame_audio_embed = \
|
||||
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
|
||||
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
|
||||
batch_audio_eb.append(frame_audio_embed)
|
||||
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
|
||||
dim=0)
|
||||
|
||||
return batch_audio_eb, min_batch_num
|
||||
@@ -0,0 +1,129 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
class CausalConv1d(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
chan_in,
|
||||
chan_out,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
dilation=1,
|
||||
pad_mode='replicate',
|
||||
**kwargs):
|
||||
super().__init__()
|
||||
|
||||
self.pad_mode = pad_mode
|
||||
padding = (kernel_size - 1, 0) # T
|
||||
self.time_causal_padding = padding
|
||||
|
||||
self.conv = nn.Conv1d(
|
||||
chan_in,
|
||||
chan_out,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
dilation=dilation,
|
||||
**kwargs)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class MotionEncoder_tc(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
in_dim: int,
|
||||
hidden_dim: int,
|
||||
num_heads=int,
|
||||
need_global=True,
|
||||
dtype=None,
|
||||
device=None):
|
||||
factory_kwargs = {"dtype": dtype, "device": device}
|
||||
super().__init__()
|
||||
|
||||
self.num_heads = num_heads
|
||||
self.need_global = need_global
|
||||
self.conv1_local = CausalConv1d(
|
||||
in_dim, hidden_dim // 4 * num_heads, 3, stride=1)
|
||||
if need_global:
|
||||
self.conv1_global = CausalConv1d(
|
||||
in_dim, hidden_dim // 4, 3, stride=1)
|
||||
self.norm1 = nn.LayerNorm(
|
||||
hidden_dim // 4,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
**factory_kwargs)
|
||||
self.act = nn.SiLU()
|
||||
self.conv2 = CausalConv1d(hidden_dim // 4, hidden_dim // 2, 3, stride=2)
|
||||
self.conv3 = CausalConv1d(hidden_dim // 2, hidden_dim, 3, stride=2)
|
||||
|
||||
if need_global:
|
||||
self.final_linear = nn.Linear(hidden_dim, hidden_dim,
|
||||
**factory_kwargs)
|
||||
|
||||
self.norm1 = nn.LayerNorm(
|
||||
hidden_dim // 4,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
**factory_kwargs)
|
||||
|
||||
self.norm2 = nn.LayerNorm(
|
||||
hidden_dim // 2,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6,
|
||||
**factory_kwargs)
|
||||
|
||||
self.norm3 = nn.LayerNorm(
|
||||
hidden_dim, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
|
||||
self.padding_tokens = nn.Parameter(torch.zeros(1, 1, 1, hidden_dim))
|
||||
|
||||
def forward(self, x):
|
||||
x = rearrange(x, 'b t c -> b c t')
|
||||
x_ori = x.clone()
|
||||
b, c, t = x.shape
|
||||
x = self.conv1_local(x)
|
||||
x = rearrange(x, 'b (n c) t -> (b n) t c', n=self.num_heads)
|
||||
x = self.norm1(x)
|
||||
x = self.act(x)
|
||||
x = rearrange(x, 'b t c -> b c t')
|
||||
x = self.conv2(x)
|
||||
x = rearrange(x, 'b c t -> b t c')
|
||||
x = self.norm2(x)
|
||||
x = self.act(x)
|
||||
x = rearrange(x, 'b t c -> b c t')
|
||||
x = self.conv3(x)
|
||||
x = rearrange(x, 'b c t -> b t c')
|
||||
x = self.norm3(x)
|
||||
x = self.act(x)
|
||||
x = rearrange(x, '(b n) t c -> b t n c', b=b)
|
||||
padding = self.padding_tokens.repeat(b, x.shape[1], 1, 1)
|
||||
x = torch.cat([x, padding], dim=-2)
|
||||
x_local = x.clone()
|
||||
|
||||
if not self.need_global:
|
||||
return x_local
|
||||
|
||||
x = self.conv1_global(x_ori)
|
||||
x = rearrange(x, 'b c t -> b t c')
|
||||
x = self.norm1(x)
|
||||
x = self.act(x)
|
||||
x = rearrange(x, 'b t c -> b c t')
|
||||
x = self.conv2(x)
|
||||
x = rearrange(x, 'b c t -> b t c')
|
||||
x = self.norm2(x)
|
||||
x = self.act(x)
|
||||
x = rearrange(x, 'b t c -> b c t')
|
||||
x = self.conv3(x)
|
||||
x = rearrange(x, 'b c t -> b t c')
|
||||
x = self.norm3(x)
|
||||
x = self.act(x)
|
||||
x = self.final_linear(x)
|
||||
x = rearrange(x, '(b n) t c -> b t n c', b=b)
|
||||
|
||||
return x, x_local
|
||||
@@ -0,0 +1,794 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import math
|
||||
from typing import Any, Dict, List, Literal, Optional, Union
|
||||
|
||||
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 einops import rearrange, repeat
|
||||
|
||||
from ..model import flash_attention
|
||||
from .s2v_utils import rope_precompute
|
||||
|
||||
|
||||
def sinusoidal_embedding_1d(dim, position):
|
||||
# preprocess
|
||||
assert dim % 2 == 0
|
||||
half = dim // 2
|
||||
position = position.type(torch.float64)
|
||||
|
||||
# calculation
|
||||
sinusoid = torch.outer(
|
||||
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
|
||||
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
return x
|
||||
|
||||
|
||||
@amp.autocast(enabled=False)
|
||||
def rope_params(max_seq_len, dim, theta=10000):
|
||||
assert dim % 2 == 0
|
||||
freqs = torch.outer(
|
||||
torch.arange(max_seq_len),
|
||||
1.0 / torch.pow(theta,
|
||||
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
|
||||
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
||||
return freqs
|
||||
|
||||
|
||||
@amp.autocast(enabled=False)
|
||||
def rope_apply(x, grid_sizes, freqs, start=None):
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
|
||||
# split freqs
|
||||
if type(freqs) is list:
|
||||
trainable_freqs = freqs[1]
|
||||
freqs = freqs[0]
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
# loop over samples
|
||||
output = []
|
||||
output = x.clone()
|
||||
seq_bucket = [0]
|
||||
if not type(grid_sizes) is list:
|
||||
grid_sizes = [grid_sizes]
|
||||
for g in grid_sizes:
|
||||
if not type(g) is list:
|
||||
g = [torch.zeros_like(g), g]
|
||||
batch_size = g[0].shape[0]
|
||||
for i in range(batch_size):
|
||||
if start is None:
|
||||
f_o, h_o, w_o = g[0][i]
|
||||
else:
|
||||
f_o, h_o, w_o = start[i]
|
||||
|
||||
f, h, w = g[1][i]
|
||||
t_f, t_h, t_w = g[2][i]
|
||||
seq_f, seq_h, seq_w = f - f_o, h - h_o, w - w_o
|
||||
seq_len = int(seq_f * seq_h * seq_w)
|
||||
if seq_len > 0:
|
||||
if t_f > 0:
|
||||
factor_f, factor_h, factor_w = (t_f / seq_f).item(), (
|
||||
t_h / seq_h).item(), (t_w / seq_w).item()
|
||||
|
||||
if f_o >= 0:
|
||||
f_sam = np.linspace(f_o.item(), (t_f + f_o).item() - 1,
|
||||
seq_f).astype(int).tolist()
|
||||
else:
|
||||
f_sam = np.linspace(-f_o.item(),
|
||||
(-t_f - f_o).item() + 1,
|
||||
seq_f).astype(int).tolist()
|
||||
h_sam = np.linspace(h_o.item(), (t_h + h_o).item() - 1,
|
||||
seq_h).astype(int).tolist()
|
||||
w_sam = np.linspace(w_o.item(), (t_w + w_o).item() - 1,
|
||||
seq_w).astype(int).tolist()
|
||||
|
||||
assert f_o * f >= 0 and h_o * h >= 0 and w_o * w >= 0
|
||||
freqs_0 = freqs[0][f_sam] if f_o >= 0 else freqs[0][
|
||||
f_sam].conj()
|
||||
freqs_0 = freqs_0.view(seq_f, 1, 1, -1)
|
||||
|
||||
freqs_i = torch.cat([
|
||||
freqs_0.expand(seq_f, seq_h, seq_w, -1),
|
||||
freqs[1][h_sam].view(1, seq_h, 1, -1).expand(
|
||||
seq_f, seq_h, seq_w, -1),
|
||||
freqs[2][w_sam].view(1, 1, seq_w, -1).expand(
|
||||
seq_f, seq_h, seq_w, -1),
|
||||
],
|
||||
dim=-1).reshape(seq_len, 1, -1)
|
||||
elif t_f < 0:
|
||||
freqs_i = trainable_freqs.unsqueeze(1)
|
||||
# apply rotary embedding
|
||||
# precompute multipliers
|
||||
x_i = torch.view_as_complex(
|
||||
x[i, seq_bucket[-1]:seq_bucket[-1] + seq_len].to(
|
||||
torch.float64).reshape(seq_len, n, -1, 2))
|
||||
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
|
||||
output[i, seq_bucket[-1]:seq_bucket[-1] + seq_len] = x_i
|
||||
seq_bucket.append(seq_bucket[-1] + seq_len)
|
||||
return output.float()
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return self._norm(x.float()).type_as(x) * self.weight
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
|
||||
class LayerNorm(nn.LayerNorm):
|
||||
|
||||
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
||||
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
||||
|
||||
def forward(self, x):
|
||||
return super().forward(x.float()).type_as(x)
|
||||
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
|
||||
# layers
|
||||
self.q = nn.Linear(dim, dim)
|
||||
self.k = nn.Linear(dim, dim)
|
||||
self.v = nn.Linear(dim, dim)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
x = flash_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)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
return x
|
||||
|
||||
|
||||
class SwinSelfAttention(SelfAttention):
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
assert b == 1, 'Only support batch_size 1'
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
q = rope_apply(q, grid_sizes, freqs)
|
||||
k = rope_apply(k, grid_sizes, freqs)
|
||||
T, H, W = grid_sizes[0].tolist()
|
||||
|
||||
q = rearrange(q, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W)
|
||||
k = rearrange(k, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W)
|
||||
v = rearrange(v, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W)
|
||||
|
||||
ref_q = q[-1:]
|
||||
q = q[:-1]
|
||||
|
||||
ref_k = repeat(
|
||||
k[-1:], "1 s n d -> t s n d", t=k.shape[0] - 1) # t hw n d
|
||||
k = k[:-1]
|
||||
k = torch.cat([k[:1], k, k[-1:]])
|
||||
k = torch.cat([k[1:-1], k[2:], k[:-2], ref_k], dim=1) # (bt) (3hw) n d
|
||||
|
||||
ref_v = repeat(v[-1:], "1 s n d -> t s n d", t=v.shape[0] - 1)
|
||||
v = v[:-1]
|
||||
v = torch.cat([v[:1], v, v[-1:]])
|
||||
v = torch.cat([v[1:-1], v[2:], v[:-2], ref_v], dim=1)
|
||||
|
||||
# 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 = 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
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
return x
|
||||
|
||||
|
||||
#Fix the reference frame RoPE to 1,H,W.
|
||||
#Set the current frame RoPE to 1.
|
||||
#Set the previous frame RoPE to 0.
|
||||
class CasualSelfAttention(SelfAttention):
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs):
|
||||
shifting = 3
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
assert b == 1, 'Only support batch_size 1'
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
T, H, W = grid_sizes[0].tolist()
|
||||
|
||||
q = rearrange(q, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W)
|
||||
k = rearrange(k, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W)
|
||||
v = rearrange(v, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W)
|
||||
|
||||
ref_q = q[-1:]
|
||||
q = q[:-1]
|
||||
|
||||
grid_sizes = torch.tensor([[1, H, W]] * q.shape[0], dtype=torch.long)
|
||||
start = [[shifting, 0, 0]] * q.shape[0]
|
||||
q = rope_apply(q, grid_sizes, freqs, start=start)
|
||||
|
||||
ref_k = k[-1:]
|
||||
grid_sizes = torch.tensor([[1, H, W]], dtype=torch.long)
|
||||
# start = [[shifting, H, W]]
|
||||
|
||||
start = [[shifting + 10, 0, 0]]
|
||||
ref_k = rope_apply(ref_k, grid_sizes, freqs, start)
|
||||
ref_k = repeat(
|
||||
ref_k, "1 s n d -> t s n d", t=k.shape[0] - 1) # t hw n d
|
||||
|
||||
k = k[:-1]
|
||||
k = torch.cat([*([k[:1]] * shifting), k])
|
||||
cat_k = []
|
||||
for i in range(shifting):
|
||||
cat_k.append(k[i:i - shifting])
|
||||
cat_k.append(k[shifting:])
|
||||
k = torch.cat(cat_k, dim=1) # (bt) (3hw) n d
|
||||
|
||||
grid_sizes = torch.tensor(
|
||||
[[shifting + 1, H, W]] * q.shape[0], dtype=torch.long)
|
||||
k = rope_apply(k, grid_sizes, freqs)
|
||||
k = torch.cat([k, ref_k], dim=1)
|
||||
|
||||
ref_v = repeat(v[-1:], "1 s n d -> t s n d", t=q.shape[0]) # t hw n d
|
||||
v = v[:-1]
|
||||
v = torch.cat([*([v[:1]] * shifting), v])
|
||||
cat_v = []
|
||||
for i in range(shifting):
|
||||
cat_v.append(v[i:i - shifting])
|
||||
cat_v.append(v[shifting:])
|
||||
v = torch.cat(cat_v, dim=1) # (bt) (3hw) n d
|
||||
v = torch.cat([v, ref_v], dim=1)
|
||||
|
||||
# q: b (t h w) n d
|
||||
# k: b (t h w) n d
|
||||
outs = []
|
||||
for i in range(q.shape[0]):
|
||||
out = flash_attention(
|
||||
q=q[i:i + 1],
|
||||
k=k[i:i + 1],
|
||||
v=v[i:i + 1],
|
||||
window_size=self.window_size)
|
||||
outs.append(out)
|
||||
out = torch.cat(outs, dim=0)
|
||||
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
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
x = self.o(x)
|
||||
return x
|
||||
|
||||
|
||||
class MotionerAttentionBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
ffn_dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
cross_attn_norm=False,
|
||||
eps=1e-6,
|
||||
self_attn_block="SelfAttention"):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
# layers
|
||||
self.norm1 = LayerNorm(dim, eps)
|
||||
if self_attn_block == "SelfAttention":
|
||||
self.self_attn = SelfAttention(dim, num_heads, window_size, qk_norm,
|
||||
eps)
|
||||
elif self_attn_block == "SwinSelfAttention":
|
||||
self.self_attn = SwinSelfAttention(dim, num_heads, window_size,
|
||||
qk_norm, eps)
|
||||
elif self_attn_block == "CasualSelfAttention":
|
||||
self.self_attn = CasualSelfAttention(dim, num_heads, window_size,
|
||||
qk_norm, eps)
|
||||
|
||||
self.norm2 = LayerNorm(dim, eps)
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn_dim, dim))
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
seq_lens,
|
||||
grid_sizes,
|
||||
freqs,
|
||||
):
|
||||
# self-attention
|
||||
y = self.self_attn(self.norm1(x).float(), seq_lens, grid_sizes, freqs)
|
||||
x = x + y
|
||||
y = self.ffn(self.norm2(x).float())
|
||||
x = x + y
|
||||
return x
|
||||
|
||||
|
||||
class Head(nn.Module):
|
||||
|
||||
def __init__(self, dim, out_dim, patch_size, eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.out_dim = out_dim
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
|
||||
# layers
|
||||
out_dim = math.prod(patch_size) * out_dim
|
||||
self.norm = LayerNorm(dim, eps)
|
||||
self.head = nn.Linear(dim, out_dim)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.head(self.norm(x))
|
||||
return x
|
||||
|
||||
|
||||
class MotionerTransformers(nn.Module, PeftAdapterMixin):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=(1, 2, 2),
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
ffn_dim=8192,
|
||||
freq_dim=256,
|
||||
out_dim=16,
|
||||
num_heads=16,
|
||||
num_layers=32,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
cross_attn_norm=False,
|
||||
eps=1e-6,
|
||||
self_attn_block="SelfAttention",
|
||||
motion_token_num=1024,
|
||||
enable_tsm=False,
|
||||
motion_stride=4,
|
||||
expand_ratio=2,
|
||||
trainable_token_pos_emb=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.in_dim = in_dim
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.freq_dim = freq_dim
|
||||
self.out_dim = out_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.window_size = window_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
self.enable_tsm = enable_tsm
|
||||
self.motion_stride = motion_stride
|
||||
self.expand_ratio = expand_ratio
|
||||
self.sample_c = self.patch_size[0]
|
||||
|
||||
# embeddings
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
|
||||
# blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
MotionerAttentionBlock(
|
||||
dim,
|
||||
ffn_dim,
|
||||
num_heads,
|
||||
window_size,
|
||||
qk_norm,
|
||||
cross_attn_norm,
|
||||
eps,
|
||||
self_attn_block=self_attn_block) for _ in range(num_layers)
|
||||
])
|
||||
|
||||
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||
d = dim // num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
],
|
||||
dim=1)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
self.motion_side_len = int(math.sqrt(motion_token_num))
|
||||
assert self.motion_side_len**2 == motion_token_num
|
||||
self.token = nn.Parameter(
|
||||
torch.zeros(1, motion_token_num, dim).contiguous())
|
||||
|
||||
self.trainable_token_pos_emb = trainable_token_pos_emb
|
||||
if trainable_token_pos_emb:
|
||||
x = torch.zeros([1, motion_token_num, num_heads, d])
|
||||
x[..., ::2] = 1
|
||||
|
||||
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),
|
||||
]]
|
||||
token_freqs = rope_apply(x, gride_sizes, self.freqs)
|
||||
token_freqs = token_freqs[0, :, 0].reshape(motion_token_num, -1, 2)
|
||||
token_freqs = token_freqs * 0.01
|
||||
self.token_freqs = torch.nn.Parameter(token_freqs)
|
||||
|
||||
def after_patch_embedding(self, x):
|
||||
return x
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x,
|
||||
):
|
||||
"""
|
||||
x: A list of videos each with shape [C, T, H, W].
|
||||
t: [B].
|
||||
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:
|
||||
freqs = freqs.to(device)
|
||||
|
||||
if self.trainable_token_pos_emb:
|
||||
with amp.autocast(dtype=torch.float64):
|
||||
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 self.enable_tsm:
|
||||
sample_idx = [
|
||||
sample_indices(
|
||||
u.shape[1],
|
||||
stride=self.motion_stride,
|
||||
expand_ratio=self.expand_ratio,
|
||||
c=self.sample_c) for u in x
|
||||
]
|
||||
x = [
|
||||
torch.flip(torch.flip(u, [1])[:, idx], [1])
|
||||
for idx, u in zip(sample_idx, x)
|
||||
]
|
||||
|
||||
# embeddings
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
x = self.after_patch_embedding(x)
|
||||
|
||||
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
|
||||
]]
|
||||
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),
|
||||
]]
|
||||
grid_sizes += tsm_frame_grid_sizes
|
||||
seq_f = sample_idx[0][-1] + 1
|
||||
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
x = torch.cat([u for u in x])
|
||||
|
||||
batch_size = len(x)
|
||||
|
||||
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),
|
||||
] # 第三行代表rope emb的想要覆盖到的范围
|
||||
]
|
||||
|
||||
grid_sizes = grid_sizes + token_grid_sizes
|
||||
token_unpatch_grid_sizes = torch.stack([
|
||||
torch.tensor([1, 32, 32], dtype=torch.long)
|
||||
for b in range(batch_size)
|
||||
])
|
||||
token_len = self.token.shape[1]
|
||||
token = self.token.clone().repeat(x.shape[0], 1, 1).contiguous()
|
||||
seq_lens = seq_lens + torch.tensor([t.size(0) for t in token],
|
||||
dtype=torch.long)
|
||||
x = torch.cat([x, token], dim=1)
|
||||
# arguments
|
||||
kwargs = dict(
|
||||
seq_lens=seq_lens,
|
||||
grid_sizes=grid_sizes,
|
||||
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)
|
||||
# head
|
||||
out = x[:, -token_len:]
|
||||
return out
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
c = self.out_dim
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = torch.einsum('fhwpqrc->cfphqwr', u)
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
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__(
|
||||
self,
|
||||
inner_dim=1024,
|
||||
num_heads=16, # Used to indicate the number of heads in the backbone network; unrelated to this module's design
|
||||
zip_frame_buckets=[
|
||||
1, 2, 16
|
||||
], # Three numbers representing the number of frames sampled for patch operations from the nearest to the farthest frames
|
||||
drop_mode="drop", # If not "drop", it will use "padd", meaning padding instead of deletion
|
||||
*args,
|
||||
**kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.proj = nn.Conv3d(
|
||||
16, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2))
|
||||
self.proj_2x = nn.Conv3d(
|
||||
16, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4))
|
||||
self.proj_4x = nn.Conv3d(
|
||||
16, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8))
|
||||
self.zip_frame_buckets = torch.tensor(
|
||||
zip_frame_buckets, dtype=torch.long)
|
||||
|
||||
self.inner_dim = inner_dim
|
||||
self.num_heads = num_heads
|
||||
|
||||
assert (inner_dim %
|
||||
num_heads) == 0 and (inner_dim // num_heads) % 2 == 0
|
||||
d = inner_dim // num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
],
|
||||
dim=1)
|
||||
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:
|
||||
lat_height, lat_width = m.shape[2], m.shape[3]
|
||||
padd_lat = torch.zeros(16, self.zip_frame_buckets.sum(), lat_height,
|
||||
lat_width).to(
|
||||
device=m.device, dtype=m.dtype)
|
||||
overlap_frame = min(padd_lat.shape[1], m.shape[1])
|
||||
if overlap_frame > 0:
|
||||
padd_lat[:, -overlap_frame:] = m[:, -overlap_frame:]
|
||||
|
||||
if add_last_motion < 2 and self.drop_mode != "drop":
|
||||
zero_end_frame = self.zip_frame_buckets[:self.zip_frame_buckets.
|
||||
__len__() -
|
||||
add_last_motion -
|
||||
1].sum()
|
||||
padd_lat[:, -zero_end_frame:] = 0
|
||||
|
||||
padd_lat = padd_lat.unsqueeze(0)
|
||||
clean_latents_4x, clean_latents_2x, clean_latents_post = padd_lat[:, :, -self.zip_frame_buckets.sum(
|
||||
):, :, :].split(
|
||||
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)
|
||||
|
||||
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
|
||||
|
||||
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())
|
||||
end_time_id = start_time_id + self.zip_frame_buckets[0]
|
||||
grid_sizes = [] if add_last_motion < 2 and self.drop_mode == "drop" else \
|
||||
[
|
||||
[torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1),
|
||||
torch.tensor([end_time_id, lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1),
|
||||
torch.tensor([self.zip_frame_buckets[0], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), ]
|
||||
]
|
||||
|
||||
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([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), ]
|
||||
]
|
||||
|
||||
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
|
||||
]).unsqueeze(0).repeat(1, 1),
|
||||
]]
|
||||
|
||||
grid_sizes = grid_sizes + grid_sizes_2x + grid_sizes_4x
|
||||
|
||||
motion_rope_emb = rope_precompute(
|
||||
motion_lat.detach().view(1, motion_lat.shape[1], self.num_heads,
|
||||
self.inner_dim // self.num_heads),
|
||||
grid_sizes,
|
||||
self.freqs,
|
||||
start=None)
|
||||
|
||||
mot.append(motion_lat)
|
||||
mot_remb.append(motion_rope_emb)
|
||||
return mot, mot_remb
|
||||
|
||||
|
||||
def sample_indices(N, stride, expand_ratio, c):
|
||||
indices = []
|
||||
current_start = 0
|
||||
|
||||
while current_start < N:
|
||||
bucket_width = int(stride * (expand_ratio**(len(indices) / stride)))
|
||||
|
||||
interval = int(bucket_width / stride * c)
|
||||
current_end = min(N, current_start + bucket_width)
|
||||
bucket_samples = []
|
||||
for i in range(current_end - 1, current_start - 1, -interval):
|
||||
for near in range(c):
|
||||
bucket_samples.append(i - near)
|
||||
|
||||
indices += bucket_samples[::-1]
|
||||
current_start += bucket_width
|
||||
|
||||
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)
|
||||
@@ -0,0 +1,70 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def rope_precompute(x, grid_sizes, freqs, start=None):
|
||||
b, s, n, c = x.size(0), x.size(1), x.size(2), x.size(3) // 2
|
||||
|
||||
# split freqs
|
||||
if type(freqs) is list:
|
||||
trainable_freqs = freqs[1]
|
||||
freqs = freqs[0]
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
# loop over samples
|
||||
output = torch.view_as_complex(x.detach().reshape(b, s, n, -1,
|
||||
2).to(torch.float64))
|
||||
seq_bucket = [0]
|
||||
if not type(grid_sizes) is list:
|
||||
grid_sizes = [grid_sizes]
|
||||
for g in grid_sizes:
|
||||
if not type(g) is list:
|
||||
g = [torch.zeros_like(g), g]
|
||||
batch_size = g[0].shape[0]
|
||||
for i in range(batch_size):
|
||||
if start is None:
|
||||
f_o, h_o, w_o = g[0][i]
|
||||
else:
|
||||
f_o, h_o, w_o = start[i]
|
||||
|
||||
f, h, w = g[1][i]
|
||||
t_f, t_h, t_w = g[2][i]
|
||||
seq_f, seq_h, seq_w = f - f_o, h - h_o, w - w_o
|
||||
seq_len = int(seq_f * seq_h * seq_w)
|
||||
if seq_len > 0:
|
||||
if t_f > 0:
|
||||
factor_f, factor_h, factor_w = (t_f / seq_f).item(), (
|
||||
t_h / seq_h).item(), (t_w / seq_w).item()
|
||||
# Generate a list of seq_f integers starting from f_o and ending at math.ceil(factor_f * seq_f.item() + f_o.item())
|
||||
if f_o >= 0:
|
||||
f_sam = np.linspace(f_o.item(), (t_f + f_o).item() - 1,
|
||||
seq_f).astype(int).tolist()
|
||||
else:
|
||||
f_sam = np.linspace(-f_o.item(),
|
||||
(-t_f - f_o).item() + 1,
|
||||
seq_f).astype(int).tolist()
|
||||
h_sam = np.linspace(h_o.item(), (t_h + h_o).item() - 1,
|
||||
seq_h).astype(int).tolist()
|
||||
w_sam = np.linspace(w_o.item(), (t_w + w_o).item() - 1,
|
||||
seq_w).astype(int).tolist()
|
||||
|
||||
assert f_o * f >= 0 and h_o * h >= 0 and w_o * w >= 0
|
||||
freqs_0 = freqs[0][f_sam] if f_o >= 0 else freqs[0][
|
||||
f_sam].conj()
|
||||
freqs_0 = freqs_0.view(seq_f, 1, 1, -1)
|
||||
|
||||
freqs_i = torch.cat([
|
||||
freqs_0.expand(seq_f, seq_h, seq_w, -1),
|
||||
freqs[1][h_sam].view(1, seq_h, 1, -1).expand(
|
||||
seq_f, seq_h, seq_w, -1),
|
||||
freqs[2][w_sam].view(1, 1, seq_w, -1).expand(
|
||||
seq_f, seq_h, seq_w, -1),
|
||||
],
|
||||
dim=-1).reshape(seq_len, 1, -1)
|
||||
elif t_f < 0:
|
||||
freqs_i = trainable_freqs.unsqueeze(1)
|
||||
# apply rotary embedding
|
||||
output[i, seq_bucket[-1]:seq_bucket[-1] + seq_len] = freqs_i
|
||||
seq_bucket.append(seq_bucket[-1] + seq_len)
|
||||
return output
|
||||
Reference in New Issue
Block a user