Binary file not shown.
Binary file not shown.
@@ -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
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
from .vqvae import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
|
||||||
@@ -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
@@ -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
|
||||||
+242
@@ -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
@@ -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
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
|
||||||
|
|||||||
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user