Basic MTV Crafter support

https://github.com/DINGYANB/MTVCrafter
This commit is contained in:
kijai
2025-08-18 22:42:11 +03:00
parent c16a7b5a7d
commit 09710f9ca0
15 changed files with 1111 additions and 131 deletions
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.
+142
View File
@@ -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
+1
View File
@@ -0,0 +1 @@
from .vqvae import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
+329
View File
@@ -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
+193
View File
@@ -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
View File
+242
View File
@@ -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"
}
+10 -2
View File
@@ -35,6 +35,13 @@ except Exception as e:
UNIANIMATE_NODE_CLASS_MAPPINGS = {} UNIANIMATE_NODE_CLASS_MAPPINGS = {}
UNIANIMATE_NODE_DISPLAY_NAME_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(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_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(NODE_CACHE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(QWEN_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(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_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(UTILITY_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_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(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"] __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+6 -3
View File
@@ -1,6 +1,7 @@
import torch import torch
import torch.nn as nn import torch.nn as nn
from accelerate import init_empty_weights 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 #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): 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.lora = None
self.step = 0 self.step = 0
self.scale_weight = scale_weight self.scale_weight = scale_weight
self.bias_function = []
self.weight_function = []
def forward(self, input): def forward(self, input):
weight = self.weight.to(input.dtype) weight, bias = cast_bias_weight(self, input)
bias = self.bias.to(input.dtype) if self.bias is not None else None
if self.scale_weight is not None: if self.scale_weight is not None:
scale_weight = self.scale_weight.to(input.device) scale_weight = self.scale_weight.to(input.device)
if weight.numel() < input.numel(): if weight.numel() < input.numel():
@@ -86,7 +89,7 @@ class CustomLinear(nn.Linear):
input = input * scale_weight input = input * scale_weight
if self.lora is not None: 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) return torch.nn.functional.linear(input, weight, bias)
-73
View File
@@ -48,26 +48,6 @@ def apply_lora(weight, lora, step=None):
weight = weight.add(patch_diff, alpha=scale) weight = weight.add(patch_diff, alpha=scale)
return weight 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): def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=None):
log.info("FP8 matmul enabled") log.info("FP8 matmul enabled")
for name, submodule in module.named_modules(): 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 original_forward = submodule.forward
setattr(submodule, "original_forward", original_forward) setattr(submodule, "original_forward", original_forward)
setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input)) 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")
+98 -23
View File
@@ -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, 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) compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance)
from .cache_methods.cache_methods import cache_report 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 .enhance_a_video.globals import set_enhance_weight, set_num_frames
from .taehv import TAEHV from .taehv import TAEHV
from contextlib import nullcontext
from einops import rearrange from einops import rearrange
from comfy import model_management as mm from comfy import model_management as mm
@@ -796,6 +797,39 @@ class WanVideoAddStandInLatent:
updated["standin_input"] = new_entry updated["standin_input"] = new_entry
return (updated,) 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: class WanVideoImageToVideoEncode:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -1620,17 +1654,13 @@ class WanVideoSampler:
is_5b = transformer.out_dim == 48 is_5b = transformer.out_dim == 48
vae_upscale_factor = 16 if is_5b else 8 vae_upscale_factor = 16 if is_5b else 8
patch_linear = transformer_options.get("patch_linear", False) if transformer.patched_linear and gguf_reader is None:
from .nodes_model_loading import load_weights, load_weights_gguf
weights_assigned = False
if not merge_loras and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, dtype, device, block_swap_args=block_swap_args) 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: if gguf_reader is not None:
load_weights_gguf(transformer, gguf_reader, patcher.model["sd"], dtype, device, patcher) load_weights_gguf(transformer, gguf_reader, patcher.model["sd"], dtype, device, patcher)
set_lora_params_gguf(transformer, patcher.patches) 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") log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
if not merge_loras and fp8_matmul: if not merge_loras and fp8_matmul:
raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported") 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), "start_percent": fantasy_portrait_embeds.get("start_percent", 0.0),
"end_percent": fantasy_portrait_embeds.get("end_percent", 1.0), "end_percent": fantasy_portrait_embeds.get("end_percent", 1.0),
} }
# MiniMax Remover # MiniMax Remover
minimax_latents = minimax_mask_latents = None minimax_latents = minimax_mask_latents = None
minimax_latents = image_embeds.get("minimax_latents", None) minimax_latents = image_embeds.get("minimax_latents", None)
@@ -2057,6 +2087,30 @@ class WanVideoSampler:
self.window_tracker = WindowTracker(verbose=context_options["verbose"]) self.window_tracker = WindowTracker(verbose=context_options["verbose"])
context = get_context_scheduler(context_schedule) 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 # vid2vid
if samples is not None: if samples is not None:
saved_generator_state = samples.get("generator_state", None) saved_generator_state = samples.get("generator_state", None)
@@ -2162,12 +2216,8 @@ class WanVideoSampler:
mm.soft_empty_cache() mm.soft_empty_cache()
gc.collect() gc.collect()
#region transformer settings
if transformer_options is not None:
block_swap_args = transformer_options.get("block_swap_args", None)
#blockswap init #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) transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
for name, param in transformer.named_parameters(): for name, param in transformer.named_parameters():
if "block" not in name: if "block" not in name:
@@ -2318,10 +2368,12 @@ class WanVideoSampler:
#region model pred #region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, 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, 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 nonlocal transformer
z = z.to(dtype) 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: if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init:
return z*0, None return z*0, None
@@ -2380,7 +2432,13 @@ class WanVideoSampler:
if recammaster is not None: if recammaster is not None:
z = torch.cat([z, recam_latents.to(z)], dim=1) 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 use_phantom = False
phantom_ref = None phantom_ref = None
if phantom_latents is not None: if phantom_latents is not None:
@@ -2483,7 +2541,11 @@ class WanVideoSampler:
"standin_input": standin_input, "standin_input": standin_input,
"fantasy_portrait_input": fantasy_portrait_input, "fantasy_portrait_input": fantasy_portrait_input,
"phantom_ref": phantom_ref, "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 batch_size = 1
@@ -2994,7 +3056,17 @@ class WanVideoSampler:
"start_percent": unianimate_poses["start_percent"], "start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_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 partial_add_cond = None
if add_cond is not None: if add_cond is not None:
partial_add_cond = add_cond[:, :, c].to(device, dtype) partial_add_cond = add_cond[:, :, c].to(device, dtype)
@@ -3011,7 +3083,8 @@ class WanVideoSampler:
cfg[idx], positive, cfg[idx], positive,
text_embeds["negative_prompt_embeds"], 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_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: if cache_args is not None:
self.window_tracker.cache_states[window_id] = new_teacache self.window_tracker.cache_states[window_id] = new_teacache
@@ -3282,7 +3355,7 @@ class WanVideoSampler:
text_embeds["prompt_embeds"], text_embeds["prompt_embeds"],
text_embeds["negative_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, 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: if bidirectional_sampling:
noise_pred_flipped, self.cache_state = predict_with_cfg( noise_pred_flipped, self.cache_state = predict_with_cfg(
latent_model_input_flipped, latent_model_input_flipped,
@@ -3290,7 +3363,7 @@ class WanVideoSampler:
text_embeds["prompt_embeds"], text_embeds["prompt_embeds"],
text_embeds["negative_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, 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: if latent_shift_loop:
#reverse latent shift #reverse latent shift
@@ -3616,7 +3689,8 @@ NODE_CLASS_MAPPINGS = {
"WanVideoLatentReScale": WanVideoLatentReScale, "WanVideoLatentReScale": WanVideoLatentReScale,
"WanVideoScheduler": WanVideoScheduler, "WanVideoScheduler": WanVideoScheduler,
"WanVideoAddStandInLatent": WanVideoAddStandInLatent, "WanVideoAddStandInLatent": WanVideoAddStandInLatent,
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler", "WanVideoSampler": "WanVideo Sampler",
@@ -3649,5 +3723,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddExtraLatent": "WanVideo Add Extra Latent", "WanVideoAddExtraLatent": "WanVideo Add Extra Latent",
"WanVideoLatentReScale": "WanVideo Latent ReScale", "WanVideoLatentReScale": "WanVideo Latent ReScale",
"WanVideoAddStandInLatent": "WanVideo Add StandIn Latent", "WanVideoAddStandInLatent": "WanVideo Add StandIn Latent",
"WanVideoAddControlEmbeds": "WanVideo Add Control Embeds" "WanVideoAddControlEmbeds": "WanVideo Add Control Embeds",
"WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion"
} }
+9 -13
View File
@@ -701,15 +701,10 @@ class WanVideoSetLoRAs:
del lora_sd 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,) return (patcher,)
def load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device, block_swap_args=None): 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...") log.info("Using accelerate to load and assign model weights to device...")
param_count = sum(1 for _ in transformer.named_parameters()) 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 continue
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else weight_dtype 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 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 dtype_to_use = base_dtype
if "patch_embedding" in name: if "patch_embedding" in name:
dtype_to_use = torch.float32 dtype_to_use = torch.float32
@@ -1143,7 +1138,8 @@ class WanVideoModelLoader:
"inject_sample_info": True if "fps_embedding.weight" in sd else False, "inject_sample_info": True if "fps_embedding.weight" in sd else False,
"add_ref_conv": True if "ref_conv.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, "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(): with init_empty_weights():
@@ -1238,7 +1234,7 @@ class WanVideoModelLoader:
if "fp8" in quantization: if "fp8" in quantization:
for k, v in sd.items(): for k, v in sd.items():
if k.endswith(".scale_weight"): if k.endswith(".scale_weight"):
scale_weights[k] = v scale_weights[k] = v.to(base_dtype)
if "fp8_e4m3fn" in quantization: if "fp8_e4m3fn" in quantization:
weight_dtype = torch.float8_e4m3fn 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"} params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv"}
control_lora = False 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: if not merge_loras and control_lora:
log.warning("Control-LoRA patching is only supported with merge_loras=True") 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) patcher, control_lora = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras)
if not gguf: if not gguf:
if merge_loras and not patch_linear: if merge_loras and lora is not None:
if not lora_low_mem_load: if not lora_low_mem_load:
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device) load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
@@ -1274,9 +1269,11 @@ class WanVideoModelLoader:
if not control_lora: if not control_lora:
scale_weights.clear() scale_weights.clear()
patcher.patches.clear() patcher.patches.clear()
transformer.patched_linear = False
else: else:
from .custom_linear import _replace_linear from .custom_linear import _replace_linear
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights) transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
transformer.patched_linear = True
if "fast" in quantization: if "fast" in quantization:
if lora is not None and not merge_loras: if lora is not None and not merge_loras:
@@ -1330,7 +1327,7 @@ class WanVideoModelLoader:
compile_args = compile_args, 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}") log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}")
patcher.model.diffusion_model.to(offload_device) patcher.model.diffusion_model.to(offload_device)
gc.collect() gc.collect()
@@ -1354,7 +1351,6 @@ class WanVideoModelLoader:
if 'transformer_options' not in patcher.model_options: if 'transformer_options' not in patcher.model_options:
patcher.model_options['transformer_options'] = {} patcher.model_options['transformer_options'] = {}
patcher.model_options["transformer_options"]["block_swap_args"] = block_swap_args 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 patcher.model_options["transformer_options"]["merge_loras"] = merge_loras
for model in mm.current_loaded_models: for model in mm.current_loaded_models:
+2
View File
@@ -189,6 +189,8 @@ def attention(
version=fa_version, version=fa_version,
) )
elif attention_mode == 'sdpa': 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() 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': elif attention_mode == 'sageattn_3':
return sageattn_blackwell( return sageattn_blackwell(
+79 -17
View File
@@ -28,6 +28,9 @@ from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCac
from ...multitalk.multitalk import get_attn_map_with_target from ...multitalk.multitalk import get_attn_map_with_target
from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
from ...MTV.mtv import apply_rotary_emb
__all__ = ['WanModel'] __all__ = ['WanModel']
from comfy import model_management as mm from comfy import model_management as mm
@@ -657,6 +660,31 @@ class WanI2VCrossAttention(WanSelfAttention):
return self.o(x) 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 = { WAN_CROSSATTENTION_CLASSES = {
't2v_cross_attn': WanT2VCrossAttention, 't2v_cross_attn': WanT2VCrossAttention,
@@ -678,6 +706,7 @@ class WanAttentionBlock(nn.Module):
eps=1e-6, eps=1e-6,
attention_mode='sdpa', attention_mode='sdpa',
rope_func="comfy", rope_func="comfy",
use_motion_attn=False
): ):
super().__init__() super().__init__()
self.dim = out_features self.dim = out_features
@@ -694,15 +723,19 @@ class WanAttentionBlock(nn.Module):
self.dense_attention_mode = "sageattn" self.dense_attention_mode = "sageattn"
self.kv_cache = None self.kv_cache = None
self.use_motion_attn = use_motion_attn
# layers # layers
self.norm1 = WanLayerNorm(out_features, eps) self.norm1 = WanLayerNorm(out_features, eps)
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode)
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": if cross_attn_type != "no_cross_attn":
self.norm3 = WanLayerNorm( self.norm3 = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
out_features, eps,
elementwise_affine=True) if cross_attn_norm else nn.Identity()
self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](in_features, self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](in_features,
out_features, out_features,
num_heads, num_heads,
@@ -779,7 +812,11 @@ class WanAttentionBlock(nn.Module):
freqs_ip=None, freqs_ip=None,
adapter_proj=None, adapter_proj=None,
ip_scale=1.0, 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""" r"""
Args: 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, 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, 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, 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: else:
if self.rope_func == "comfy_chunked": if self.rope_func == "comfy_chunked":
y = self.ffn_chunked(x, shift_mlp, scale_mlp) 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, 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, audio_proj, audio_scale, num_latent_frames, nag_params,
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num, nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale): 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, x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
audio_proj=audio_proj, audio_scale=audio_scale, 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, 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, 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) adapter_proj=adapter_proj, ip_scale=ip_scale)
#multitalk # MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock): 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, 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) shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
x = x + x_audio * audio_scale 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": if self.rope_func == "comfy_chunked":
y = self.ffn_chunked(x, shift_mlp, scale_mlp) y = self.ffn_chunked(x, shift_mlp, scale_mlp)
else: else:
@@ -1154,6 +1197,7 @@ class WanModel(torch.nn.Module):
in_dim_ref_conv=16, in_dim_ref_conv=16,
add_control_adapter=False, add_control_adapter=False,
in_dim_control_adapter=24, in_dim_control_adapter=24,
use_motion_attn=False
): ):
r""" r"""
Initialize the diffusion model backbone. Initialize the diffusion model backbone.
@@ -1215,6 +1259,7 @@ class WanModel(torch.nn.Module):
self.offload_device = offload_device self.offload_device = offload_device
self.vace_layers = vace_layers self.vace_layers = vace_layers
self.device = main_device self.device = main_device
self.patched_linear = False
self.blocks_to_swap = -1 self.blocks_to_swap = -1
self.offload_txt_emb = False self.offload_txt_emb = False
@@ -1312,9 +1357,12 @@ class WanModel(torch.nn.Module):
self.blocks = nn.ModuleList([ self.blocks = nn.ModuleList([
WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads, WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm, cross_attn_norm, eps, qk_norm, cross_attn_norm, eps,
attention_mode=self.attention_mode, rope_func=self.rope_func) attention_mode=self.attention_mode, rope_func=self.rope_func, use_motion_attn=(i % 4 == 0 and use_motion_attn))
for _ in range(num_layers) for i in range(num_layers)
]) ])
#MTV Crafter
if use_motion_attn:
self.pad_motion_tokens = torch.zeros(1, 1, 2048)
# head # head
self.head = Head(dim, out_dim, patch_size, eps) self.head = Head(dim, out_dim, patch_size, eps)
@@ -1543,7 +1591,12 @@ class WanModel(torch.nn.Module):
standin_input=None, standin_input=None,
fantasy_portrait_input=None, fantasy_portrait_input=None,
phantom_ref=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""" r"""
Forward pass through the diffusion model 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 # Stand-In only used on first positive pass, then cached in kv_cache
if is_uncond or current_step > 0: if is_uncond or current_step > 0:
standin_input = None 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 # Fantasy Portrait
adapter_proj = ip_scale = None 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, e_ip=e0_ip if x_ip is not None else None,
adapter_proj=adapter_proj, adapter_proj=adapter_proj,
ip_scale=ip_scale, 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: if vace_data is not None: