From 09710f9ca05f10786c23f3ebb3af9690d4fe158a Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 18 Aug 2025 22:42:11 +0300 Subject: [PATCH] Basic MTV Crafter support https://github.com/DINGYANB/MTVCrafter --- MTV/data/mean.npy | Bin 0 -> 416 bytes MTV/data/std.npy | Bin 0 -> 416 bytes MTV/draw_pose.py | 142 +++++++++++++++ MTV/motion4d/__init__.py | 1 + MTV/motion4d/vqvae.py | 329 ++++++++++++++++++++++++++++++++++ MTV/mtv.py | 193 ++++++++++++++++++++ MTV/nlf.py | 0 MTV/nodes.py | 242 +++++++++++++++++++++++++ __init__.py | 12 +- custom_linear.py | 9 +- fp8_optimization.py | 73 -------- nodes.py | 121 ++++++++++--- nodes_model_loading.py | 22 +-- wanvideo/modules/attention.py | 2 + wanvideo/modules/model.py | 96 ++++++++-- 15 files changed, 1111 insertions(+), 131 deletions(-) create mode 100644 MTV/data/mean.npy create mode 100644 MTV/data/std.npy create mode 100644 MTV/draw_pose.py create mode 100644 MTV/motion4d/__init__.py create mode 100644 MTV/motion4d/vqvae.py create mode 100644 MTV/mtv.py create mode 100644 MTV/nlf.py create mode 100644 MTV/nodes.py diff --git a/MTV/data/mean.npy b/MTV/data/mean.npy new file mode 100644 index 0000000000000000000000000000000000000000..001d7e7a450836a286a261bec91b2e72f8f926d2 GIT binary patch literal 416 zcmbR27wQ`j$;eQ~P_3SlTAW;@Zl$1ZlV+l>qoAIaUsO_*m=~X4l#&V(cT3DEP6dh= zXCxM+0{I$7COQhnnmP)#3giN=a1|GOB^58{4;vr5WNnpn+U?Npd@tv*OFXmqA@00R zXZ72UT`V@MJM3co=(H^PnT!6sHBKsfe>mrEz3wt++LlB8Gk-d-UVY2u?T;-EdAB+p z`=&j1`DuB?$zs2oOM2l0mrjB8hYtO8b;&;U$c0mAhQm3}RR@`0KXmzd<(-rFkzf~t z`8QntF28wb!-im&b!oR;`US4-mrb@iTnKbS^SNNBPm7)&x~B5TMJCz$(Bbco4jo{A z;=)iM=OEQyd^j=tu8Y*tmrkuW{vBev^}t1X|J6fX9~lmBy7AcMrJknqtdb7LLzC{f za7$?(?mfljxaG)0mo;MC&av{UPX0<)T>2-A9X4k<>*zb-x(jFbKPR^x_D&~uTy$xy Q6g<5B=4;0@)2_Gx0E$4Y1poj5 literal 0 HcmV?d00001 diff --git a/MTV/data/std.npy b/MTV/data/std.npy new file mode 100644 index 0000000000000000000000000000000000000000..5d1db82923759854dcd53a28beb59abeab80e711 GIT binary patch literal 416 zcmbR27wQ`j$;eQ~P_3SlTAW;@Zl$1ZlV+l>qoAIaUsO_*m=~X4l#&V(cT3DEP6dh= zXCxM+0{I$7COQhnnmP)#3giN=2j@kcSMjwvN0$n^tjm#hzUyA&+;1rCBJL#XY!_7J zoWDxgWumjJ^U;*a&ircpF8tNO&cTZVokPyayZm|<={(0b(D|Q(yi3>tedj)w`OY7F z1zrA|&UXGVtIK)yEhQK6lXILu^L0D3K2~)3^~u87{@EhuYt=$7=VjJ7Kk%F6YJp?B<-9|7qqCx?h|3h!SZBWH8=bk|3cIvU zsB>N%cF@^kjkwF(h+Jp>e_Ne%)WlpQb{9HtkKOM4>6oa?WuaQ)f%BA$YA!8Vvz#p_ LzH`nKR&xOWCFqdU literal 0 HcmV?d00001 diff --git a/MTV/draw_pose.py b/MTV/draw_pose.py new file mode 100644 index 0000000..eb5c606 --- /dev/null +++ b/MTV/draw_pose.py @@ -0,0 +1,142 @@ +import cv2 +import math +import torch +import numpy as np +from PIL import Image +from torchvision import transforms + + +def intrinsic_matrix_from_field_of_view(imshape, fov_degrees:float =55 ): # nlf default fov_degrees 55 + imshape = np.array(imshape) + fov_radians = fov_degrees * np.array(np.pi / 180) + larger_side = np.max(imshape) + focal_length = larger_side / (np.tan(fov_radians / 2) * 2) + # intrinsic_matrix 3*3 + return np.array([ + [focal_length, 0, imshape[1] / 2], + [0, focal_length, imshape[0] / 2], + [0, 0, 1], + ]) + + +def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3 + camera_matrix = intrinsic_matrix_from_field_of_view((height,width)) + camera_matrix = np.expand_dims(camera_matrix, axis=0) + camera_matrix = np.expand_dims(camera_matrix, axis=0) # 1*1*3*3 + point_3d = np.expand_dims(point_3d,axis=-1) # n*1024*3*1 + point_2d = (camera_matrix@point_3d).squeeze(-1) + point_2d[:,:,:2] = point_2d[:,:,:2]/point_2d[:,:,2:3] + return point_2d[:,:,:] # n*1024*2 + + +def get_pose_images(smpl_data, offset): + pose_images = [] + for data in smpl_data: + if isinstance(data, np.ndarray): + joints3d = data + else: + joints3d = data.numpy() + canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8) + joints3d = p3d_to_p2d(joints3d, offset[0], offset[1]) + canvas = draw_3d_points(canvas, joints3d[0], stickwidth=int(offset[1]/350)) + pose_images.append(Image.fromarray(canvas)) + return pose_images + + +def get_control_conditions(poses, h, w): + video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True) + control_images = [] + for idx, pose in enumerate(poses): + canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8) + try: + joints3d = p3d_to_p2d(pose, h, w) + canvas = draw_3d_points( + canvas, + joints3d[0], + stickwidth=int(h / 350), + ) + resized_canvas = cv2.resize(canvas, (w, h)) + # Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg') + control_images.append(resized_canvas) + except Exception as e: + print("wrong:", e) + control_images.append(Image.fromarray(canvas)) + control_pixel_values = np.array(control_images) + control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255. + print("control_pixel_values.shape", control_pixel_values.shape) + #control_pixel_values = video_transforms(control_pixel_values) + return control_pixel_values + + +def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True): + colors = [ + [255, 0, 0], # 0 + [0, 255, 0], # 1 + [0, 0, 255], # 2 + [255, 0, 255], # 3 + [255, 255, 0], # 4 + [85, 255, 0], # 5 + [0, 75, 255], # 6 + [0, 255, 85], # 7 + [0, 255, 170], # 8 + [170, 0, 255], # 9 + [85, 0, 255], # 10 + [0, 85, 255], # 11 + [0, 255, 255], # 12 + [85, 0, 255], # 13 + [170, 0, 255], # 14 + [255, 0, 255], # 15 + [255, 0, 170], # 16 + [255, 0, 85], # 17 + ] + connetions = [ + [15,12],[12, 16],[16, 18],[18, 20],[20, 22], + [12,17],[17,19],[19,21], + [21,23],[12,9],[9,6], + [6,3],[3,0],[0,1], + [1,4],[4,7],[7,10],[0,2],[2,5],[5,8],[8,11] + ] + connection_colors = [ + [255, 0, 0], # 0 + [0, 255, 0], # 1 + [0, 0, 255], # 2 + [255, 255, 0], # 3 + [255, 0, 255], # 4 + [0, 255, 0], # 5 + [0, 85, 255], # 6 + [255, 175, 0], # 7 + [0, 0, 255], # 8 + [255, 85, 0], # 9 + [0, 255, 85], # 10 + [255, 0, 255], # 11 + [255, 0, 0], # 12 + [0, 175, 255], # 13 + [255, 255, 0], # 14 + [0, 0, 255], # 15 + [0, 255, 0], # 16 + ] + + # draw point + for i in range(len(points)): + x,y = points[i][0:2] + x,y = int(x),int(y) + if i==13 or i == 14: + continue + cv2.circle(canvas, (x, y), r, colors[i%17], thickness=-1) + + # draw line + if draw_line: + for i in range(len(connetions)): + point1_idx,point2_idx = connetions[i][0:2] + point1 = points[point1_idx] + point2 = points[point2_idx] + Y = [point2[0],point1[0]] + X = [point2[1],point1[1]] + mX = int(np.mean(X)) + mY = int(np.mean(Y)) + length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5 + angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1])) + polygon = cv2.ellipse2Poly((mY, mX), (int(length / 2), stickwidth), int(angle), 0, 360, 1) + cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17]) + + return canvas diff --git a/MTV/motion4d/__init__.py b/MTV/motion4d/__init__.py new file mode 100644 index 0000000..87ac3b3 --- /dev/null +++ b/MTV/motion4d/__init__.py @@ -0,0 +1 @@ +from .vqvae import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder \ No newline at end of file diff --git a/MTV/motion4d/vqvae.py b/MTV/motion4d/vqvae.py new file mode 100644 index 0000000..30535e0 --- /dev/null +++ b/MTV/motion4d/vqvae.py @@ -0,0 +1,329 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F +import numpy as np + + +class Encoder(nn.Module): + def __init__( + self, + in_channels=3, + mid_channels=[128, 512], + out_channels=3072, + downsample_time=[1, 1], + downsample_joint=[1, 1], + num_attention_heads=8, + attention_head_dim=64, + dim=3072, + ): + super(Encoder, self).__init__() + + self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1) + self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)]) + self.downsample1 = Downsample(mid_channels[0], mid_channels[0], downsample_time[0], downsample_joint[0]) + self.resnet2 = ResBlock(mid_channels[0], mid_channels[1]) + self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)]) + self.downsample2 = Downsample(mid_channels[1], mid_channels[1], downsample_time[1], downsample_joint[1]) + self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1) + + def forward(self, x): + x = self.conv_in(x) + for resnet in self.resnet1: + x = resnet(x) + x = self.downsample1(x) + + x = self.resnet2(x) + for resnet in self.resnet3: + x = resnet(x) + x = self.downsample2(x) + + x = self.conv_out(x) + + return x + + + +class VectorQuantizer(nn.Module): + def __init__(self, nb_code, code_dim): + super().__init__() + self.nb_code = nb_code + self.code_dim = code_dim + self.mu = 0.99 + self.reset_codebook() + self.reset_count = 0 + self.usage = torch.zeros((self.nb_code, 1)) + + def reset_codebook(self): + self.init = False + self.code_sum = None + self.code_count = None + self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).cuda()) + + def _tile(self, x): + nb_code_x, code_dim = x.shape + if nb_code_x < self.nb_code: + n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x + std = 0.01 / np.sqrt(code_dim) + out = x.repeat(n_repeats, 1) + out = out + torch.randn_like(out) * std + else: + out = x + return out + + def preprocess(self, x): + # [bs, c, f, j] -> [bs * f * j, c] + x = x.permute(0, 2, 3, 1).contiguous() + x = x.view(-1, x.shape[-1]) + return x + + def quantize(self, x): + # [bs * f * j, dim=3072] + # Calculate latent code x_l + k_w = self.codebook.t() + distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0, keepdim=True) + _, code_idx = torch.min(distance, dim=-1) + return code_idx + + def dequantize(self, code_idx): + x = F.embedding(code_idx, self.codebook) # indexing: [bs * f * j, 32] + return x + + def forward(self, x, return_vq=False): + bs, c, f, j = x.shape # SMPL data frames: [bs, 3072, f, j] + + # Preprocess + x = self.preprocess(x) + # return x.view(bs, f*j, c).contiguous(), None + assert x.shape[-1] == self.code_dim + + # quantize and dequantize through bottleneck + code_idx = self.quantize(x) + x_d = self.dequantize(code_idx) + + # Loss + commit_loss = F.mse_loss(x, x_d.detach()) + + # Passthrough + x_d = x + (x_d - x).detach() + + if return_vq: + return x_d.view(bs, f*j, c).contiguous(), commit_loss + # return (x_d, x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()), commit_loss, perplexity + + # Postprocess + x_d = x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous() + + return x_d, commit_loss + + + + +class Decoder(nn.Module): + def __init__( + self, + in_channels=3072, + mid_channels=[512, 128], + out_channels=3, + upsample_rate=None, + frame_upsample_rate=[1.0, 1.0], + joint_upsample_rate=[1.0, 1.0], + dim=128, + attention_head_dim=64, + num_attention_heads=8, + ): + super(Decoder, self).__init__() + + self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1) + self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)]) + self.upsample1 = Upsample(mid_channels[0], mid_channels[0], frame_upsample_rate=frame_upsample_rate[0], joint_upsample_rate=joint_upsample_rate[0]) + self.resnet2 = ResBlock(mid_channels[0], mid_channels[1]) + self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)]) + self.upsample2 = Upsample(mid_channels[1], mid_channels[1], frame_upsample_rate=frame_upsample_rate[1], joint_upsample_rate=joint_upsample_rate[1]) + self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1) + + def forward(self, x): + x = self.conv_in(x) + for resnet in self.resnet1: + x = resnet(x) + x = self.upsample1(x) + + x = self.resnet2(x) + for resnet in self.resnet3: + x = resnet(x) + x = self.upsample2(x) + + x = self.conv_out(x) + + return x + + +class Upsample(nn.Module): + def __init__( + self, + in_channels, + out_channels, + upsample_rate=None, + frame_upsample_rate=None, + joint_upsample_rate=None, + ): + super(Upsample, self).__init__() + + self.upsampler = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) + self.upsample_rate = upsample_rate + self.frame_upsample_rate = frame_upsample_rate + self.joint_upsample_rate = joint_upsample_rate + self.upsample_rate = upsample_rate + + def forward(self, inputs): + if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1: + # split first frame + x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:] + + if self.upsample_rate is not None: + # import pdb; pdb.set_trace() + x_first = F.interpolate(x_first, scale_factor=self.upsample_rate) + x_rest = F.interpolate(x_rest, scale_factor=self.upsample_rate) + else: + # import pdb; pdb.set_trace() + # x_first = F.interpolate(x_first, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True) + x_rest = F.interpolate(x_rest, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True) + x_first = x_first[:, :, None, :] + inputs = torch.cat([x_first, x_rest], dim=2) + elif inputs.shape[2] > 1: + if self.upsample_rate is not None: + inputs = F.interpolate(inputs, scale_factor=self.upsample_rate) + else: + inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True) + else: + inputs = inputs.squeeze(2) + if self.upsample_rate is not None: + inputs = F.interpolate(inputs, scale_factor=self.upsample_rate) + else: + inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="linear", align_corners=True) + inputs = inputs[:, :, None, :, :] + + b, c, t, j = inputs.shape + inputs = inputs.permute(0, 2, 1, 3).reshape(b * t, c, j) + inputs = self.upsampler(inputs) + inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3) + + return inputs + + +class Downsample(nn.Module): + def __init__( + self, + in_channels, + out_channels, + frame_downsample_rate, + joint_downsample_rate + ): + super(Downsample, self).__init__() + + self.frame_downsample_rate = frame_downsample_rate + self.joint_downsample_rate = joint_downsample_rate + self.joint_downsample = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=self.joint_downsample_rate, padding=1) + + def forward(self, x): + # (batch_size, channels, frames, joints) -> (batch_size * joints, channels, frames) + if self.frame_downsample_rate > 1: + batch_size, channels, frames, joints = x.shape + x = x.permute(0, 3, 1, 2).reshape(batch_size * joints, channels, frames) + if x.shape[-1] % 2 == 1: + x_first, x_rest = x[..., 0], x[..., 1:] + if x_rest.shape[-1] > 0: + # (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2) + x_rest = F.avg_pool1d(x_rest, kernel_size=self.frame_downsample_rate, stride=self.frame_downsample_rate) + + x = torch.cat([x_first[..., None], x_rest], dim=-1) + # (batch_size * joints, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, joints) + x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1) + else: + # (batch_size * joints, channels, frames) -> (batch_size * joints, channels, frames // 2) + x = F.avg_pool1d(x, kernel_size=2, stride=2) + # (batch_size * joints, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width) + x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1) + + # Pad the tensor + # pad = (0, 1) + # x = F.pad(x, pad, mode="constant", value=0) + batch_size, channels, frames, joints = x.shape + # (batch_size, channels, frames, joints) -> (batch_size * frames, channels, joints) + x = x.permute(0, 2, 1, 3).reshape(batch_size * frames, channels, joints) + x = self.joint_downsample(x) + # (batch_size * frames, channels, joints) -> (batch_size, channels, frames, joints) + x = x.reshape(batch_size, frames, x.shape[1], x.shape[2]).permute(0, 2, 1, 3) + return x + + + +class ResBlock(nn.Module): + def __init__(self, + in_channels, + out_channels, + group_num=32, + max_channels=512): + super(ResBlock, self).__init__() + skip = max(1, max_channels // out_channels - 1) + self.block = nn.Sequential( + nn.GroupNorm(group_num, in_channels, eps=1e-06, affine=True), + nn.SiLU(), + nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=skip, dilation=skip), + nn.GroupNorm(group_num, out_channels, eps=1e-06, affine=True), + nn.SiLU(), + nn.Conv2d(out_channels, out_channels, kernel_size=1, stride=1, padding=0), + ) + self.conv_short = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) if in_channels != out_channels else nn.Identity() + + def forward(self, x): + hidden_states = self.block(x) + if hidden_states.shape != x.shape: + x = self.conv_short(x) + x = x + hidden_states + return x + + + +class SMPL_VQVAE(nn.Module): + def __init__(self, encoder, decoder, vq): + super(SMPL_VQVAE, self).__init__() + + self.encoder = encoder + self.decoder = decoder + self.vq = vq + + def to(self, device): + self.encoder = self.encoder.to(device) + self.decoder = self.decoder.to(device) + self.vq = self.vq.to(device) + self.device = device + return self + + def encdec_slice_frames(self, x, frame_batch_size, encdec, return_vq): + num_frames = x.shape[2] + remaining_frames = num_frames % frame_batch_size + x_output = [] + + for i in range(num_frames // frame_batch_size): + remaining_frames = num_frames % frame_batch_size + start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames) + end_frame = frame_batch_size * (i + 1) + remaining_frames + x_intermediate = x[:, :, start_frame:end_frame] + x_intermediate = encdec(x_intermediate) + x_output.append(x_intermediate) + if encdec == self.encoder and self.vq is not None: + x_output, loss = self.vq(torch.cat(x_output, dim=2), return_vq=return_vq) + return x_output, loss + else: + return torch.cat(x_output, dim=2), None, None + + def forward(self, x, return_vq=False): + x = x.permute(0, 3, 1, 2) + x, loss = self.encdec_slice_frames(x, frame_batch_size=8, encdec=self.encoder, return_vq=return_vq) + + if return_vq: + return x, loss + x, _, _ = self.encdec_slice_frames(x, frame_batch_size=2, encdec=self.decoder, return_vq=return_vq) + x = x.permute(0, 2, 3, 1) + + return x, loss diff --git a/MTV/mtv.py b/MTV/mtv.py new file mode 100644 index 0000000..915773d --- /dev/null +++ b/MTV/mtv.py @@ -0,0 +1,193 @@ +import torch +import numpy as np +from typing import Union, Tuple + + +def get_1d_rotary_pos_embed( + dim: int, + pos: Union[np.ndarray, int], + theta: float = 10000.0, + use_real=False, + linear_factor=1.0, + ntk_factor=1.0, + repeat_interleave_real=True, + freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux) +): + """ + Precompute the frequency tensor for complex exponentials (cis) with given dimensions. + + This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end + index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64 + data type. + + Args: + dim (`int`): Dimension of the frequency tensor. + pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar + theta (`float`, *optional*, defaults to 10000.0): + Scaling factor for frequency computation. Defaults to 10000.0. + use_real (`bool`, *optional*): + If True, return real part and imaginary part separately. Otherwise, return complex numbers. + linear_factor (`float`, *optional*, defaults to 1.0): + Scaling factor for the context extrapolation. Defaults to 1.0. + ntk_factor (`float`, *optional*, defaults to 1.0): + Scaling factor for the NTK-Aware RoPE. Defaults to 1.0. + repeat_interleave_real (`bool`, *optional*, defaults to `True`): + If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`. + Otherwise, they are concateanted with themselves. + freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`): + the dtype of the frequency tensor. + Returns: + `torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2] + """ + assert dim % 2 == 0 + + if isinstance(pos, int): + pos = torch.arange(pos) + if isinstance(pos, np.ndarray): + pos = torch.from_numpy(pos) # type: ignore # [S] + + theta = theta * ntk_factor + freqs = ( + 1.0 + / (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim)) + / linear_factor + ) # [D/2] + freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2] + if use_real and repeat_interleave_real: + freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D] + freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D] + return freqs_cos, freqs_sin + elif use_real: + freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D] + freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D] + return freqs_cos, freqs_sin + else: + freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2] + return freqs_cis + + +def get_3d_rotary_pos_embed( + embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True +) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + """ + RoPE for video tokens with 3D structure. + + Args: + embed_dim: (`int`): + The embedding dimension size, corresponding to hidden_size_head. + crops_coords (`Tuple[int]`): + The top-left and bottom-right coordinates of the crop. + grid_size (`Tuple[int]`): + The grid size of the spatial positional embedding (height, width). + temporal_size (`int`): + The size of the temporal dimension. + theta (`float`): + Scaling factor for frequency computation. + + Returns: + `torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`. + """ + if use_real is not True: + raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed") + start, stop = crops_coords + grid_size_h, grid_size_w = grid_size + grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32) + grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32) + grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32) + + # Compute dimensions for each axis + dim_t = embed_dim // 4 + dim_h = embed_dim // 8 * 3 + dim_w = embed_dim // 8 * 3 + + # Temporal frequencies + freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True) + # Spatial frequencies for height and width + freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True) + freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True) + + # BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor + def combine_time_height_width(freqs_t, freqs_h, freqs_w): + freqs_t = freqs_t[:, None, None, :].expand( + -1, grid_size_h, grid_size_w, -1 + ) # temporal_size, grid_size_h, grid_size_w, dim_t + freqs_h = freqs_h[None, :, None, :].expand( + temporal_size, -1, grid_size_w, -1 + ) # temporal_size, grid_size_h, grid_size_2, dim_h + freqs_w = freqs_w[None, None, :, :].expand( + temporal_size, grid_size_h, -1, -1 + ) # temporal_size, grid_size_h, grid_size_2, dim_w + + freqs = torch.cat( + [freqs_t, freqs_h, freqs_w], dim=-1 + ) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w) + freqs = freqs.view( + temporal_size * grid_size_h * grid_size_w, -1 + ) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w) + return freqs + + t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t + h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h + w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w + cos = combine_time_height_width(t_cos, h_cos, w_cos) + sin = combine_time_height_width(t_sin, h_sin, w_sin) + return cos, sin + + +def get_3d_motion_spatial_embed( + embed_dim: int, num_joints: int, joints_mean: np.ndarray, joints_std: np.ndarray, theta: float = 10000.0 +) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: + assert embed_dim % 2 == 0 and embed_dim % 3 == 0 + + def create_rope_pe(dim, pos, freqs_dtype=torch.float32): + if isinstance(pos, np.ndarray): + pos = torch.from_numpy(pos) + freqs = ( + 1.0 + / (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim)) + ) # [D/2] + freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2] + freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D] + freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D] + return freqs_cos, freqs_sin + + pos_x = joints_mean[:, 0] + pos_y = joints_mean[:, 1] + pos_z = joints_mean[:, 2] + + normalized_pos_x = (pos_x - pos_x.mean()) + normalized_pos_y = (pos_y - pos_y.mean()) + normalized_pos_z = (pos_z - pos_z.mean()) + + freqs_cos_x, freqs_sin_x = create_rope_pe(embed_dim // 3, normalized_pos_x) + freqs_cos_y, freqs_sin_y = create_rope_pe(embed_dim // 3, normalized_pos_y) + freqs_cos_z, freqs_sin_z = create_rope_pe(embed_dim // 3, normalized_pos_z) + + freqs_cos = torch.cat([freqs_cos_x, freqs_cos_y, freqs_cos_z], dim=-1) + freqs_sin = torch.cat([freqs_sin_x, freqs_sin_y, freqs_sin_z], dim=-1) + + return freqs_cos, freqs_sin + +def prepare_motion_embeddings(num_frames, num_joints, joints_mean, joints_std, theta=10000, device='cuda'): + time_embed = get_1d_rotary_pos_embed(44, num_frames, theta, use_real=True) + time_embed_cos = time_embed[0][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1) + time_embed_sin = time_embed[1][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1) + spatial_motion_embed = get_3d_motion_spatial_embed(84, num_joints, joints_mean, joints_std, theta) + spatial_embed_cos = spatial_motion_embed[0][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1) + spatial_embed_sin = spatial_motion_embed[1][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1) + motion_embed_cos = torch.cat([time_embed_cos, spatial_embed_cos], dim=-1).to(device=device) + motion_embed_sin = torch.cat([time_embed_sin, spatial_embed_sin], dim=-1).to(device=device) + return motion_embed_cos, motion_embed_sin + +def apply_rotary_emb(x, freqs_cis): + cos, sin = freqs_cis # [S, D] + cos = cos[None, None] + sin = sin[None, None] + cos, sin = cos.to(x.device), sin.to(x.device) + + x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2] + x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3) + + out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype) + + return out \ No newline at end of file diff --git a/MTV/nlf.py b/MTV/nlf.py new file mode 100644 index 0000000..e69de29 diff --git a/MTV/nodes.py b/MTV/nodes.py new file mode 100644 index 0000000..bb6ae4c --- /dev/null +++ b/MTV/nodes.py @@ -0,0 +1,242 @@ +import os +import torch +import gc +from ..utils import log, dict_to_device +import numpy as np +from accelerate import init_empty_weights +from accelerate.utils import set_module_tensor_to_device + +import comfy.model_management as mm +from comfy.utils import load_torch_file +import folder_paths + +script_directory = os.path.dirname(os.path.abspath(__file__)) +device = mm.get_torch_device() +offload_device = mm.unet_offload_device() + +local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript") + +from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder +from .mtv import prepare_motion_embeddings + +class DownloadAndLoadNLFModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "url": ( + [ + "https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript" + ], + ) + }, + } + + RETURN_TYPES = ("NLFMODEL",) + RETURN_NAMES = ("nlf_model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, url): + + if not os.path.exists(local_model_path): + log.info(f"Downloading NLF model to: {local_model_path}") + import requests + os.makedirs(os.path.dirname(local_model_path), exist_ok=True) + response = requests.get(url) + if response.status_code == 200: + with open(local_model_path, "wb") as f: + f.write(response.content) + else: + print("Failed to download file:", response.status_code) + + model = torch.jit.load(local_model_path).eval() + + return (model,) + +class LoadNLFModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "path": ("STRING", {"default": local_model_path}), + }, + } + + RETURN_TYPES = ("NLFMODEL",) + RETURN_NAMES = ("nlf_model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, path): + model = torch.jit.load(path).eval() + + return model, + +class LoadVQVAE: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}), + }, + } + + RETURN_TYPES = ("VQVAE",) + RETURN_NAMES = ("vqvae", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, model_name): + model_path = folder_paths.get_full_path("vae", model_name) + vae_sd = load_torch_file(model_path, safe_load=True) + + # Get motion tokenizer + motion_encoder = Encoder( + in_channels=3, + mid_channels=[128, 512], + out_channels=3072, + downsample_time=[2, 2], + downsample_joint=[1, 1] + ) + motion_quant = VectorQuantizer(nb_code=8192, code_dim=3072) + motion_decoder = Decoder( + in_channels=3072, + mid_channels=[512, 128], + out_channels=3, + upsample_rate=2.0, + frame_upsample_rate=[2.0, 2.0], + joint_upsample_rate=[1.0, 1.0] + ) + + vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device) + vqvae.load_state_dict(vae_sd, strict=True) + + return vqvae, + +class MTVCrafterEncodePoses: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "vqvae": ("VQVAE", {"tooltip": "VQVAE model"}), + "poses": ("NLFPRED", {"tooltip": "Input poses for the model"}), + }, + } + + RETURN_TYPES = ("MTVCRAFTERMOTION", "NLFPRED") + RETURN_NAMES = ("mtvcrafter_motion", "pose_results") + FUNCTION = "encode" + CATEGORY = "WanVideoWrapper" + + def encode(self, vqvae, poses): + + # import pickle + # with open(os.path.join(script_directory, "data", "sampled_data.pkl"), 'rb') as f: + # data_list = pickle.load(f) + # if not isinstance(data_list, list): + # data_list = [data_list] + # print(data_list) + + # smpl_poses = data_list[1]['pose'] + + global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3) + global_std = np.load(os.path.join(script_directory, "data", "std.npy")) + + smpl_poses = [] + for pose in poses['joints3d_nonparam'][0]: + smpl_poses.append(pose[0].cpu().numpy()) + smpl_poses = np.array(smpl_poses) + + norm_poses = torch.tensor((smpl_poses - global_mean) / global_std).unsqueeze(0) + print(f"norm_poses shape: {norm_poses.shape}, dtype: {norm_poses.dtype}") + + vqvae.to(device) + motion_tokens, vq_loss = vqvae(norm_poses.to(device), return_vq=True) + + recon_motion = vqvae(norm_poses.to(device))[0][0].to(dtype=torch.float32).cpu().detach() * global_std + global_mean + vqvae.to(offload_device) + + poses_dict = { + 'mtv_motion_tokens': motion_tokens, + 'global_mean': global_mean, + 'global_std': global_std + } + + return poses_dict, recon_motion + + +class NLFPredict: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "model": ("NLFMODEL",), + "images": ("IMAGE", {"tooltip": "Input images for the model"}), + }, + } + + RETURN_TYPES = ("NLFPRED", ) + RETURN_NAMES = ("pose_results",) + FUNCTION = "predict" + CATEGORY = "WanVideoWrapper" + + def predict(self, model, images): + + model.to(device) + pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device)) + model.to(offload_device) + + pred = dict_to_device(pred, offload_device) + + pose_results = { + 'joints3d_nonparam': [], + } + # Collect pose data + for key in pose_results.keys(): + if key in pred: + pose_results[key].append(pred[key]) + else: + pose_results[key].append(None) + + return (pose_results,) + +class DrawNLFPoses: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "poses": ("NLFPRED", {"tooltip": "Input poses for the model"}), + "width": ("INT", {"default": 512}), + "height": ("INT", {"default": 512}), + }, + } + + RETURN_TYPES = ("IMAGE", ) + RETURN_NAMES = ("image",) + FUNCTION = "predict" + CATEGORY = "WanVideoWrapper" + + def predict(self, poses, width, height): + from .draw_pose import get_control_conditions + print(type(poses)) + if isinstance(poses, dict): + pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses + else: + pose_input = poses + control_conditions = get_control_conditions(pose_input, height, width) + + return (control_conditions,) + +NODE_CLASS_MAPPINGS = { + "DownloadAndLoadNLFModel": DownloadAndLoadNLFModel, + "NLFPredict": NLFPredict, + "DrawNLFPoses": DrawNLFPoses, + "LoadVQVAE": LoadVQVAE, + "MTVCrafterEncodePoses": MTVCrafterEncodePoses + } +NODE_DISPLAY_NAME_MAPPINGS = { + "DownloadAndLoadNLFModel": "(Download)Load NLF Model", + "NLFPredict": "NLF Predict", + "DrawNLFPoses": "Draw NLF Poses", + "LoadVQVAE": "Load VQVAE", + "MTVCrafterEncodePoses": "MTV Crafter Encode Poses" +} diff --git a/__init__.py b/__init__.py index 112b36e..1f81508 100644 --- a/__init__.py +++ b/__init__.py @@ -35,6 +35,13 @@ except Exception as e: UNIANIMATE_NODE_CLASS_MAPPINGS = {} UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS = {} +try: + from .MTV.nodes import NODE_CLASS_MAPPINGS as MTV_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MTV_NODE_DISPLAY_NAME_MAPPINGS +except Exception as e: + print(f"MTV nodes not available due to error in importing them: {e}") + MTV_NODE_CLASS_MAPPINGS = {} + MTV_NODE_DISPLAY_NAME_MAPPINGS = {} + NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS) @@ -50,6 +57,7 @@ NODE_CLASS_MAPPINGS.update(UTILITY_NODE_CLASS_MAPPINGS) 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_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS) @@ -65,7 +73,7 @@ NODE_DISPLAY_NAME_MAPPINGS.update(MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UTILITY_NODE_DISPLAY_NAME_MAPPINGS) 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(QWEN_NODE_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS) __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/custom_linear.py b/custom_linear.py index 391502b..5236458 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -1,6 +1,7 @@ import torch import torch.nn as nn from accelerate import init_empty_weights +from comfy.ops import cast_bias_weight #based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None): @@ -74,10 +75,12 @@ class CustomLinear(nn.Linear): self.lora = None self.step = 0 self.scale_weight = scale_weight + self.bias_function = [] + self.weight_function = [] def forward(self, input): - weight = self.weight.to(input.dtype) - bias = self.bias.to(input.dtype) if self.bias is not None else None + weight, bias = cast_bias_weight(self, input) + if self.scale_weight is not None: scale_weight = self.scale_weight.to(input.device) if weight.numel() < input.numel(): @@ -86,7 +89,7 @@ class CustomLinear(nn.Linear): input = input * scale_weight if self.lora is not None: - weight = self.apply_lora(weight).to(input.dtype) + weight = self.apply_lora(weight).to(self.compute_dtype) return torch.nn.functional.linear(input, weight, bias) diff --git a/fp8_optimization.py b/fp8_optimization.py index dac0d3d..33d6f60 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -48,26 +48,6 @@ def apply_lora(weight, lora, step=None): weight = weight.add(patch_diff, alpha=scale) return weight - -def linear_with_lora_and_scale_forward(cls, input): - # Handles both scaled and unscaled, with or without LoRA - has_scale = hasattr(cls, "scale_weight") - weight = cls.weight.to(input.dtype) - bias = cls.bias.to(input.dtype) if cls.bias is not None else None - - if has_scale: - scale_weight = cls.scale_weight.to(input.device) - if weight.numel() < input.numel(): - weight = weight * scale_weight - else: - input = input * scale_weight - - lora = getattr(cls, "lora", None) - if lora is not None: - weight = apply_lora(weight, lora, cls.step).to(input.dtype) - - return torch.nn.functional.linear(input, weight, bias) - def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=None): log.info("FP8 matmul enabled") for name, submodule in module.named_modules(): @@ -81,57 +61,4 @@ def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys= original_forward = submodule.forward setattr(submodule, "original_forward", original_forward) setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input)) - -def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=None, params_to_keep={}): - log.info("Patching Linear layers...") - for name, submodule in module.named_modules(): - if not any(keyword in name for keyword in params_to_keep): - # Set scale_weight if present - if scale_weight_keys is not None: - scale_key = f"{name}.scale_weight" - if scale_key in scale_weight_keys: - setattr(submodule, "scale_weight", scale_weight_keys[scale_key]) - # Set LoRA if present - if hasattr(submodule, "lora"): - #print(f"removing old LoRA in {name}" ) - delattr(submodule, "lora") - if patches is not None: - patch_key1 = f"diffusion_model.{name}.weight" - patch_key_compiled = f"diffusion_model.{name.replace('_orig_mod.', '')}.weight" - patch = patches.get(patch_key1, []) or patches.get(patch_key_compiled, []) - if len(patch) != 0: - lora_diffs = [] - for p in patch: - lora_obj = p[1] - if "head" in name: - continue # For now skip LoRA for head layers - elif hasattr(lora_obj, "weights"): - lora_diffs.append(lora_obj.weights) - elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff": - lora_diffs.append(lora_obj[1]) - else: - continue - lora_strengths = [p[0] for p in patch] - lora = (lora_diffs, lora_strengths) - setattr(submodule, "lora", lora) - #print(f"Added LoRA to {name} with {len(lora_diffs)} diffs and strengths {lora_strengths}") - - # Set forward if Linear and has either scale or lora - if isinstance(submodule, nn.Linear): - has_scale = hasattr(submodule, "scale_weight") - has_lora = hasattr(submodule, "lora") - if not hasattr(submodule, "original_forward"): - setattr(submodule, "original_forward", submodule.forward) - if has_scale or has_lora: - setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_and_scale_forward(m, input)) - setattr(submodule, "step", 0) # Initialize step for LoRA scheduling - -def remove_lora_from_module(module): - unloaded = False - for name, submodule in module.named_modules(): - if hasattr(submodule, "lora"): - if not unloaded: - log.info("Unloading all LoRAs") - unloaded = True - delattr(submodule, "lora") diff --git a/nodes.py b/nodes.py index d0059cd..bdac6e4 100644 --- a/nodes.py +++ b/nodes.py @@ -16,9 +16,10 @@ from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fouri add_noise_to_reference_video, optimized_scale, setup_radial_attention, compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance) from .cache_methods.cache_methods import cache_report +from .nodes_model_loading import load_weights, load_weights_gguf from .enhance_a_video.globals import set_enhance_weight, set_num_frames from .taehv import TAEHV - +from contextlib import nullcontext from einops import rearrange from comfy import model_management as mm @@ -796,6 +797,39 @@ class WanVideoAddStandInLatent: updated["standin_input"] = new_entry return (updated,) +class WanVideoAddMTVMotion: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "embeds": ("WANVIDIMAGE_EMBEDS",), + "mtv_crafter_motion": ("MTVCRAFTERMOTION",), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the MTV motion"}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent to apply the ref "}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent to apply the ref "}), + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "add" + CATEGORY = "WanVideoWrapper" + + def add(self, embeds, mtv_crafter_motion, strength, start_percent, end_percent): + # Prepare the new extra latent entry + new_entry = { + "mtv_motion_tokens": mtv_crafter_motion["mtv_motion_tokens"], + "strength": strength, + "start_percent": start_percent, + "end_percent": end_percent, + "global_mean": mtv_crafter_motion["global_mean"], + "global_std": mtv_crafter_motion["global_std"] + } + + # Return a new dict with updated extra_latents + updated = dict(embeds) + updated["mtv_crafter_motion"] = new_entry + return (updated,) + class WanVideoImageToVideoEncode: @classmethod def INPUT_TYPES(s): @@ -1620,17 +1654,13 @@ class WanVideoSampler: is_5b = transformer.out_dim == 48 vae_upscale_factor = 16 if is_5b else 8 - patch_linear = transformer_options.get("patch_linear", False) - from .nodes_model_loading import load_weights, load_weights_gguf - weights_assigned = False - if not merge_loras and gguf_reader is None: + if transformer.patched_linear and gguf_reader is None: load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, dtype, device, block_swap_args=block_swap_args) - weights_assigned = True if gguf_reader is not None: load_weights_gguf(transformer, gguf_reader, patcher.model["sd"], dtype, device, patcher) set_lora_params_gguf(transformer, patcher.patches) - elif len(patcher.patches) != 0 and patch_linear: + elif len(patcher.patches) != 0 and transformer.patched_linear: log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model") if not merge_loras and fp8_matmul: raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported") @@ -1998,7 +2028,7 @@ class WanVideoSampler: "start_percent": fantasy_portrait_embeds.get("start_percent", 0.0), "end_percent": fantasy_portrait_embeds.get("end_percent", 1.0), } - + # MiniMax Remover minimax_latents = minimax_mask_latents = None minimax_latents = image_embeds.get("minimax_latents", None) @@ -2057,6 +2087,30 @@ class WanVideoSampler: self.window_tracker = WindowTracker(verbose=context_options["verbose"]) context = get_context_scheduler(context_schedule) + #MTV Crafter + mtv_input = image_embeds.get("mtv_crafter_motion", None) + if mtv_input is not None: + from .MTV.mtv import prepare_motion_embeddings + log.info("Using MTV Crafter embeddings") + mtv_start_percent = mtv_input.get("start_percent", 0.0) + mtv_end_percent = mtv_input.get("end_percent", 1.0) + mtv_strength = mtv_input.get("strength", 1.0) + mtv_motion_tokens = mtv_input.get("mtv_motion_tokens", None) + if not isinstance(mtv_strength, list): + mtv_strength = [mtv_strength] * (steps + 1) + d = transformer.dim // transformer.num_heads + mtv_freqs = torch.cat([ + rope_params(1024, d - 4 * (d // 6)), + rope_params(1024, 2 * (d // 6)), + rope_params(1024, 2 * (d // 6)) + ], + dim=1) + motion_rotary_emb = prepare_motion_embeddings( + latent_video_length if context_options is None else context_frames, + 24, mtv_input["global_mean"], [mtv_input["global_std"]], device=device) + log.info(f"mtv_motion_rotary_emb: {motion_rotary_emb[0].shape}") + mtv_freqs = mtv_freqs.to(device, dtype) + # vid2vid if samples is not None: saved_generator_state = samples.get("generator_state", None) @@ -2162,12 +2216,8 @@ class WanVideoSampler: mm.soft_empty_cache() gc.collect() - #region transformer settings - if transformer_options is not None: - block_swap_args = transformer_options.get("block_swap_args", None) - #blockswap init - if block_swap_args is not None: + if block_swap_args is not None and not transformer.patched_linear: transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) for name, param in transformer.named_parameters(): if "block" not in name: @@ -2318,10 +2368,12 @@ class WanVideoSampler: #region model pred 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): + add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False, + mtv_motion_tokens=None): nonlocal transformer z = z.to(dtype) - with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): + autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear) + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype) if autocast_enabled else nullcontext(): if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init: return z*0, None @@ -2380,7 +2432,13 @@ class WanVideoSampler: if recammaster is not None: z = torch.cat([z, recam_latents.to(z)], dim=1) - + + if mtv_input is not None: + if ((mtv_start_percent <= current_step_percentage <= mtv_end_percent) or \ + (mtv_end_percent > 0 and idx == 0 and current_step_percentage >= mtv_start_percent)): + mtv_motion_tokens = mtv_motion_tokens.to(z) + mtv_motion_rotary_emb = motion_rotary_emb + use_phantom = False phantom_ref = None if phantom_latents is not None: @@ -2483,7 +2541,11 @@ class WanVideoSampler: "standin_input": standin_input, "fantasy_portrait_input": fantasy_portrait_input, "phantom_ref": phantom_ref, - "reverse_time": reverse_time + "reverse_time": reverse_time, + "mtv_motion_tokens": mtv_motion_tokens if mtv_input is not None else None, + "mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None, + "mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, + "mtv_freqs": mtv_freqs if mtv_input is not None else None, } batch_size = 1 @@ -2994,7 +3056,17 @@ class WanVideoSampler: "start_percent": unianimate_poses["start_percent"], "end_percent": unianimate_poses["end_percent"] } - + + partial_mtv_motion_tokens = None + if mtv_input is not None: + start_token_index = c[0] * 24 + end_token_index = (c[-1] + 1) * 24 + partial_mtv_motion_tokens = mtv_motion_tokens[:, start_token_index:end_token_index, :] + print("mtv_motion_tokens", mtv_motion_tokens.shape) + if context_options["verbose"]: + log.info(f"context window: {c}") + log.info(f"motion_token_indices: {start_token_index}-{end_token_index}") + partial_add_cond = None if add_cond is not None: partial_add_cond = add_cond[:, :, c].to(device, dtype) @@ -3011,7 +3083,8 @@ class WanVideoSampler: cfg[idx], positive, 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) + 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) if cache_args is not None: self.window_tracker.cache_states[window_id] = new_teacache @@ -3282,7 +3355,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) + cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens) if bidirectional_sampling: noise_pred_flipped, self.cache_state = predict_with_cfg( latent_model_input_flipped, @@ -3290,7 +3363,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, reverse_time=True) + cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens,reverse_time=True) if latent_shift_loop: #reverse latent shift @@ -3616,7 +3689,8 @@ NODE_CLASS_MAPPINGS = { "WanVideoLatentReScale": WanVideoLatentReScale, "WanVideoScheduler": WanVideoScheduler, "WanVideoAddStandInLatent": WanVideoAddStandInLatent, - "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds + "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds, + "WanVideoAddMTVMotion": WanVideoAddMTVMotion, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3649,5 +3723,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoAddExtraLatent": "WanVideo Add Extra Latent", "WanVideoLatentReScale": "WanVideo Latent ReScale", "WanVideoAddStandInLatent": "WanVideo Add StandIn Latent", - "WanVideoAddControlEmbeds": "WanVideo Add Control Embeds" + "WanVideoAddControlEmbeds": "WanVideo Add Control Embeds", + "WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion" } diff --git a/nodes_model_loading.py b/nodes_model_loading.py index c1aa61f..d206413 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -701,15 +701,10 @@ class WanVideoSetLoRAs: del lora_sd - if 'transformer_options' not in patcher.model_options: - patcher.model_options['transformer_options'] = {} - - patcher.model_options['transformer_options']["patch_linear"] = True - return (patcher,) def load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device, block_swap_args=None): - params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "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_proj"} log.info("Using accelerate to load and assign model weights to device...") param_count = sum(1 for _ in transformer.named_parameters()) @@ -736,7 +731,7 @@ def load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_dev continue dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else weight_dtype dtype_to_use = weight_dtype if sd[name.replace("_orig_mod.", "")].dtype == weight_dtype else dtype_to_use - if "modulation" in name or "norm" in name or "bias" in name: + if "modulation" in name or "norm" in name or "bias" in name or "img_emb" in name: dtype_to_use = base_dtype if "patch_embedding" in name: dtype_to_use = torch.float32 @@ -1143,7 +1138,8 @@ class WanVideoModelLoader: "inject_sample_info": True if "fps_embedding.weight" in sd else False, "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 + "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 } with init_empty_weights(): @@ -1238,7 +1234,7 @@ class WanVideoModelLoader: if "fp8" in quantization: for k, v in sd.items(): if k.endswith(".scale_weight"): - scale_weights[k] = v + scale_weights[k] = v.to(base_dtype) if "fp8_e4m3fn" in quantization: weight_dtype = torch.float8_e4m3fn @@ -1250,7 +1246,6 @@ class WanVideoModelLoader: params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv"} control_lora = False - patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False) if not merge_loras and control_lora: log.warning("Control-LoRA patching is only supported with merge_loras=True") @@ -1259,7 +1254,7 @@ class WanVideoModelLoader: patcher, control_lora = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras) if not gguf: - if merge_loras and not patch_linear: + if merge_loras and lora is not None: if not lora_low_mem_load: load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device) @@ -1274,9 +1269,11 @@ class WanVideoModelLoader: if not control_lora: scale_weights.clear() patcher.patches.clear() + transformer.patched_linear = False else: from .custom_linear import _replace_linear transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights) + transformer.patched_linear = True if "fast" in quantization: if lora is not None and not merge_loras: @@ -1330,7 +1327,7 @@ class WanVideoModelLoader: compile_args = compile_args, ) - if merge_loras: + if merge_loras and lora is not None: log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}") patcher.model.diffusion_model.to(offload_device) gc.collect() @@ -1354,7 +1351,6 @@ class WanVideoModelLoader: if 'transformer_options' not in patcher.model_options: patcher.model_options['transformer_options'] = {} patcher.model_options["transformer_options"]["block_swap_args"] = block_swap_args - patcher.model_options["transformer_options"]["patch_linear"] = patch_linear patcher.model_options["transformer_options"]["merge_loras"] = merge_loras for model in mm.current_loaded_models: diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 1030a30..b923419 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -189,6 +189,8 @@ def attention( version=fa_version, ) elif attention_mode == 'sdpa': + if not (q.dtype == k.dtype == v.dtype): + return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype)).transpose(1, 2).contiguous() return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous() elif attention_mode == 'sageattn_3': return sageattn_blackwell( diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 1ceb681..07ea579 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -28,6 +28,9 @@ from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCac from ...multitalk.multitalk import get_attn_map_with_target from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot +from ...MTV.mtv import apply_rotary_emb + + __all__ = ['WanModel'] from comfy import model_management as mm @@ -657,6 +660,31 @@ class WanI2VCrossAttention(WanSelfAttention): return self.o(x) +class MTVCrafterMotionAttention(WanSelfAttention): + + def forward(self, x, mo, pe, grid_sizes, freqs): + r""" + Args: + x(Tensor): Shape [B, L1, C] + mo: Motion tokens + pe: 4D RoPE + """ + b, n, d = x.size(0), self.num_heads, self.head_dim + + # compute query, key, value + q = self.norm_q(self.q(x)).view(b, -1, n, d) + k = self.norm_k(self.k(mo)).view(b, n, -1, d) + v = self.v(mo).view(b, -1, n, d) + + # compute attention + x = attention( + q=rope_apply(q, grid_sizes, freqs), + k=apply_rotary_emb(k, pe).transpose(1, 2), + v=v + ) + + return self.o(x.flatten(2)) + WAN_CROSSATTENTION_CLASSES = { 't2v_cross_attn': WanT2VCrossAttention, @@ -678,6 +706,7 @@ class WanAttentionBlock(nn.Module): eps=1e-6, attention_mode='sdpa', rope_func="comfy", + use_motion_attn=False ): super().__init__() self.dim = out_features @@ -694,15 +723,19 @@ class WanAttentionBlock(nn.Module): self.dense_attention_mode = "sageattn" self.kv_cache = None + self.use_motion_attn = use_motion_attn # layers self.norm1 = WanLayerNorm(out_features, eps) - self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, - eps, self.attention_mode) + self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode) + + # MTV Crafter motion attn + if self.use_motion_attn: + self.norm4 = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() + self.motion_attn = MTVCrafterMotionAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode) + if cross_attn_type != "no_cross_attn": - self.norm3 = WanLayerNorm( - out_features, eps, - elementwise_affine=True) if cross_attn_norm else nn.Identity() + self.norm3 = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](in_features, out_features, num_heads, @@ -779,7 +812,11 @@ class WanAttentionBlock(nn.Module): freqs_ip=None, adapter_proj=None, ip_scale=1.0, - reverse_time=False + reverse_time=False, + mtv_motion_tokens=None, + mtv_motion_rotary_emb=None, + mtv_strength=1.0, + mtv_freqs=None ): r""" Args: @@ -921,7 +958,8 @@ class WanAttentionBlock(nn.Module): x = self.cross_attn_ffn(x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed, audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs, - adapter_proj=adapter_proj, ip_scale=ip_scale) + adapter_proj=adapter_proj, ip_scale=ip_scale, + mtv_freqs=mtv_freqs, mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength) else: if self.rope_func == "comfy_chunked": y = self.ffn_chunked(x, shift_mlp, scale_mlp) @@ -940,19 +978,24 @@ class WanAttentionBlock(nn.Module): def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed, audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num, - inner_t, inner_c, cross_freqs, adapter_proj, ip_scale): - - x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed, - audio_proj=audio_proj, audio_scale=audio_scale, - num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond, + inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength): + + x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed, + audio_proj=audio_proj, audio_scale=audio_scale, + num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond, rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs, adapter_proj=adapter_proj, ip_scale=ip_scale) - #multitalk + # MultiTalk if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock): x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding, shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num) x = x + x_audio * audio_scale + # MTV-Crafter Motion Attention + if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None: + x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs) + x = x + x_motion * mtv_strength + if self.rope_func == "comfy_chunked": y = self.ffn_chunked(x, shift_mlp, scale_mlp) else: @@ -1154,6 +1197,7 @@ class WanModel(torch.nn.Module): in_dim_ref_conv=16, add_control_adapter=False, in_dim_control_adapter=24, + use_motion_attn=False ): r""" Initialize the diffusion model backbone. @@ -1215,6 +1259,7 @@ class WanModel(torch.nn.Module): self.offload_device = offload_device self.vace_layers = vace_layers self.device = main_device + self.patched_linear = False self.blocks_to_swap = -1 self.offload_txt_emb = False @@ -1312,9 +1357,12 @@ class WanModel(torch.nn.Module): self.blocks = nn.ModuleList([ WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, - attention_mode=self.attention_mode, rope_func=self.rope_func) - for _ in range(num_layers) + attention_mode=self.attention_mode, rope_func=self.rope_func, use_motion_attn=(i % 4 == 0 and use_motion_attn)) + for i in range(num_layers) ]) + #MTV Crafter + if use_motion_attn: + self.pad_motion_tokens = torch.zeros(1, 1, 2048) # head self.head = Head(dim, out_dim, patch_size, eps) @@ -1543,7 +1591,12 @@ class WanModel(torch.nn.Module): standin_input=None, fantasy_portrait_input=None, phantom_ref=None, - reverse_time=False + reverse_time=False, + mtv_motion_tokens=None, + mtv_motion_rotary_emb=None, + mtv_freqs=None, + mtv_strength=1.0, + ): r""" Forward pass through the diffusion model @@ -1569,6 +1622,11 @@ class WanModel(torch.nn.Module): # Stand-In only used on first positive pass, then cached in kv_cache if is_uncond or current_step > 0: standin_input = None + + # MTV Crafter motion projection + if mtv_motion_tokens is not None: + bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1] + mtv_motion_tokens = torch.cat([mtv_motion_tokens, self.pad_motion_tokens.to(mtv_motion_tokens).expand(bs, motion_seq_len, -1)], dim=-1) # Fantasy Portrait adapter_proj = ip_scale = None @@ -2010,7 +2068,11 @@ class WanModel(torch.nn.Module): e_ip=e0_ip if x_ip is not None else None, adapter_proj=adapter_proj, ip_scale=ip_scale, - reverse_time=reverse_time + reverse_time=reverse_time, + mtv_motion_tokens=mtv_motion_tokens, + mtv_motion_rotary_emb=mtv_motion_rotary_emb, + mtv_strength=mtv_strength, + mtv_freqs=mtv_freqs ) if vace_data is not None: