Merge branch 'dev'
This commit is contained in:
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_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .MTV.nodes import NODE_CLASS_MAPPINGS as MTV_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MTV_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
print(f"MTV nodes not available due to error in importing them: {e}")
|
||||
MTV_NODE_CLASS_MAPPINGS = {}
|
||||
MTV_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
|
||||
@@ -50,6 +57,7 @@ NODE_CLASS_MAPPINGS.update(UTILITY_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
@@ -65,7 +73,7 @@ NODE_DISPLAY_NAME_MAPPINGS.update(MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UTILITY_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
+7
-4
@@ -1,6 +1,7 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from accelerate import init_empty_weights
|
||||
from comfy.ops import cast_bias_weight
|
||||
|
||||
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
|
||||
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None):
|
||||
@@ -12,7 +13,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
||||
module_prefix = prefix + name + "."
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
|
||||
|
||||
if isinstance(module, nn.Linear):
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
if scale_weights is not None:
|
||||
@@ -74,10 +75,12 @@ class CustomLinear(nn.Linear):
|
||||
self.lora = None
|
||||
self.step = 0
|
||||
self.scale_weight = scale_weight
|
||||
self.bias_function = []
|
||||
self.weight_function = []
|
||||
|
||||
def forward(self, input):
|
||||
weight = self.weight.to(input.dtype)
|
||||
bias = self.bias.to(input.dtype) if self.bias is not None else None
|
||||
weight, bias = cast_bias_weight(self, input)
|
||||
|
||||
if self.scale_weight is not None:
|
||||
scale_weight = self.scale_weight.to(input.device)
|
||||
if weight.numel() < input.numel():
|
||||
@@ -86,7 +89,7 @@ class CustomLinear(nn.Linear):
|
||||
input = input * scale_weight
|
||||
|
||||
if self.lora is not None:
|
||||
weight = self.apply_lora(weight).to(input.dtype)
|
||||
weight = self.apply_lora(weight).to(self.compute_dtype)
|
||||
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
|
||||
Binary file not shown.
@@ -611,6 +611,7 @@
|
||||
},
|
||||
{
|
||||
"name": "image_2",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": 152
|
||||
},
|
||||
@@ -662,6 +663,7 @@
|
||||
},
|
||||
{
|
||||
"name": "image_2",
|
||||
"shape": 7,
|
||||
"type": "IMAGE",
|
||||
"link": 150
|
||||
}
|
||||
@@ -806,7 +808,14 @@
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"inputs": [
|
||||
{
|
||||
"name": "compile_args",
|
||||
"shape": 7,
|
||||
"type": "WANCOMPILEARGS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
@@ -903,7 +912,7 @@
|
||||
],
|
||||
"size": [
|
||||
887.1368408203125,
|
||||
934.646484375
|
||||
334
|
||||
],
|
||||
"flags": {},
|
||||
"order": 38,
|
||||
@@ -988,6 +997,7 @@
|
||||
"inputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "WANVAE",
|
||||
"link": 170
|
||||
},
|
||||
@@ -1095,6 +1105,7 @@
|
||||
"inputs": [
|
||||
{
|
||||
"name": "t5",
|
||||
"shape": 7,
|
||||
"type": "WANTEXTENCODER",
|
||||
"link": 15
|
||||
},
|
||||
@@ -1123,7 +1134,9 @@
|
||||
"widgets_values": [
|
||||
"CG动画风格,一只蓝色的小鸟从地面起飞,煽动翅膀。小鸟羽毛细腻,胸前有独特的花纹,背景是蓝天白云,阳光明媚。镜跟随小鸟向上移动,展现出小鸟飞翔的姿态和天空的广阔。近景,仰视视角",
|
||||
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
||||
true
|
||||
true,
|
||||
false,
|
||||
"gpu"
|
||||
],
|
||||
"color": "#332922",
|
||||
"bgcolor": "#593930"
|
||||
@@ -1219,7 +1232,7 @@
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
154
|
||||
202
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
@@ -1245,7 +1258,9 @@
|
||||
false,
|
||||
false,
|
||||
true,
|
||||
0
|
||||
0,
|
||||
0,
|
||||
false
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
@@ -1477,7 +1492,8 @@
|
||||
"0, 0, 0",
|
||||
"center",
|
||||
16,
|
||||
"cpu"
|
||||
"cpu",
|
||||
"<tr><td>Output: </td><td><b>1</b> x <b>640</b> x <b>640 | 4.69MB</b></td></tr>"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -1489,7 +1505,7 @@
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
286
|
||||
336
|
||||
],
|
||||
"flags": {},
|
||||
"order": 28,
|
||||
@@ -1564,9 +1580,199 @@
|
||||
"0, 0, 0",
|
||||
"center",
|
||||
16,
|
||||
"cpu"
|
||||
"cpu",
|
||||
"<tr><td>Output: </td><td><b>1</b> x <b>640</b> x <b>640 | 4.69MB</b></td></tr>"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 106,
|
||||
"type": "WanVideoLoraSelect",
|
||||
"pos": [
|
||||
-336.7720642089844,
|
||||
-698.3348999023438
|
||||
],
|
||||
"size": [
|
||||
424.9496765136719,
|
||||
150
|
||||
],
|
||||
"flags": {},
|
||||
"order": 17,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "prev_lora",
|
||||
"shape": 7,
|
||||
"type": "WANVIDLORA",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "blocks",
|
||||
"shape": 7,
|
||||
"type": "SELECTEDBLOCKS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "lora",
|
||||
"type": "WANVIDLORA",
|
||||
"links": [
|
||||
179
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "ComfyUI-WanVideoWrapper",
|
||||
"ver": "974dd656dab305f7fa122cca435759105ea44488",
|
||||
"Node name for S&R": "WanVideoLoraSelect"
|
||||
},
|
||||
"widgets_values": [
|
||||
"Wan21_T2V_14B_lightx2v_cfg_step_distill_lora_rank32.safetensors",
|
||||
1.2000000000000002,
|
||||
false,
|
||||
true
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
},
|
||||
{
|
||||
"id": 35,
|
||||
"type": "WanVideoTorchCompileSettings",
|
||||
"pos": [
|
||||
-307.4797058105469,
|
||||
-1197.4749755859375
|
||||
],
|
||||
"size": [
|
||||
421.6000061035156,
|
||||
202
|
||||
],
|
||||
"flags": {},
|
||||
"order": 18,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "torch_compile_args",
|
||||
"type": "WANCOMPILEARGS",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
190
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "ComfyUI-WanVideoWrapper",
|
||||
"ver": "d9b1f4d1a5aea91d101ae97a54714a5861af3f50",
|
||||
"Node name for S&R": "WanVideoTorchCompileSettings"
|
||||
},
|
||||
"widgets_values": [
|
||||
"inductor",
|
||||
false,
|
||||
"default",
|
||||
false,
|
||||
64,
|
||||
true,
|
||||
128
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
},
|
||||
{
|
||||
"id": 22,
|
||||
"type": "WanVideoModelLoader",
|
||||
"pos": [
|
||||
119.37029266357422,
|
||||
-926.8419799804688
|
||||
],
|
||||
"size": [
|
||||
477.4410095214844,
|
||||
314
|
||||
],
|
||||
"flags": {},
|
||||
"order": 24,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "compile_args",
|
||||
"shape": 7,
|
||||
"type": "WANCOMPILEARGS",
|
||||
"link": 190
|
||||
},
|
||||
{
|
||||
"name": "block_swap_args",
|
||||
"shape": 7,
|
||||
"type": "BLOCKSWAPARGS",
|
||||
"link": 174
|
||||
},
|
||||
{
|
||||
"name": "lora",
|
||||
"shape": 7,
|
||||
"type": "WANVIDLORA",
|
||||
"link": 179
|
||||
},
|
||||
{
|
||||
"name": "vram_management_args",
|
||||
"shape": 7,
|
||||
"type": "VRAM_MANAGEMENTARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "extra_model",
|
||||
"shape": 7,
|
||||
"type": "VACEPATH",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "fantasytalking_model",
|
||||
"shape": 7,
|
||||
"type": "FANTASYTALKINGMODEL",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "multitalk_model",
|
||||
"shape": 7,
|
||||
"type": "MULTITALKMODEL",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "fantasyportrait_model",
|
||||
"shape": 7,
|
||||
"type": "FANTASYPORTRAITMODEL",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vace_model",
|
||||
"shape": 7,
|
||||
"type": "VACEPATH",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "WANVIDEOMODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
29,
|
||||
103
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "ComfyUI-WanVideoWrapper",
|
||||
"ver": "d9b1f4d1a5aea91d101ae97a54714a5861af3f50",
|
||||
"Node name for S&R": "WanVideoModelLoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"WanVideo\\Wan2_1-FLF2V-14B-720P_fp8_e4m3fn.safetensors",
|
||||
"fp16_fast",
|
||||
"fp8_e4m3fn",
|
||||
"offload_device",
|
||||
"sageattn"
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
},
|
||||
{
|
||||
"id": 27,
|
||||
"type": "WanVideoSampler",
|
||||
@@ -1675,6 +1881,12 @@
|
||||
"shape": 7,
|
||||
"type": "MULTITALK_EMBEDS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "freeinit_args",
|
||||
"shape": 7,
|
||||
"type": "FREEINITARGS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
@@ -1685,6 +1897,11 @@
|
||||
"links": [
|
||||
166
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "denoised_samples",
|
||||
"type": "LATENT",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
@@ -1704,184 +1921,10 @@
|
||||
1,
|
||||
"",
|
||||
"comfy",
|
||||
""
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 106,
|
||||
"type": "WanVideoLoraSelect",
|
||||
"pos": [
|
||||
-336.7720642089844,
|
||||
-698.3348999023438
|
||||
],
|
||||
"size": [
|
||||
424.9496765136719,
|
||||
126
|
||||
],
|
||||
"flags": {},
|
||||
"order": 17,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "prev_lora",
|
||||
"shape": 7,
|
||||
"type": "WANVIDLORA",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "blocks",
|
||||
"shape": 7,
|
||||
"type": "SELECTEDBLOCKS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "lora",
|
||||
"type": "WANVIDLORA",
|
||||
"links": [
|
||||
179
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "ComfyUI-WanVideoWrapper",
|
||||
"ver": "974dd656dab305f7fa122cca435759105ea44488",
|
||||
"Node name for S&R": "WanVideoLoraSelect"
|
||||
},
|
||||
"widgets_values": [
|
||||
"Wan21_T2V_14B_lightx2v_cfg_step_distill_lora_rank32.safetensors",
|
||||
1.2000000000000002,
|
||||
0,
|
||||
-1,
|
||||
false
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
},
|
||||
{
|
||||
"id": 35,
|
||||
"type": "WanVideoTorchCompileSettings",
|
||||
"pos": [
|
||||
-307.4797058105469,
|
||||
-1197.4749755859375
|
||||
],
|
||||
"size": [
|
||||
421.6000061035156,
|
||||
202
|
||||
],
|
||||
"flags": {},
|
||||
"order": 18,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "torch_compile_args",
|
||||
"type": "WANCOMPILEARGS",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
190
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "ComfyUI-WanVideoWrapper",
|
||||
"ver": "d9b1f4d1a5aea91d101ae97a54714a5861af3f50",
|
||||
"Node name for S&R": "WanVideoTorchCompileSettings"
|
||||
},
|
||||
"widgets_values": [
|
||||
"inductor",
|
||||
false,
|
||||
"default",
|
||||
false,
|
||||
64,
|
||||
true,
|
||||
128
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
},
|
||||
{
|
||||
"id": 22,
|
||||
"type": "WanVideoModelLoader",
|
||||
"pos": [
|
||||
119.37029266357422,
|
||||
-926.8419799804688
|
||||
],
|
||||
"size": [
|
||||
477.4410095214844,
|
||||
274
|
||||
],
|
||||
"flags": {},
|
||||
"order": 24,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "compile_args",
|
||||
"shape": 7,
|
||||
"type": "WANCOMPILEARGS",
|
||||
"link": 190
|
||||
},
|
||||
{
|
||||
"name": "block_swap_args",
|
||||
"shape": 7,
|
||||
"type": "BLOCKSWAPARGS",
|
||||
"link": 174
|
||||
},
|
||||
{
|
||||
"name": "lora",
|
||||
"shape": 7,
|
||||
"type": "WANVIDLORA",
|
||||
"link": 179
|
||||
},
|
||||
{
|
||||
"name": "vram_management_args",
|
||||
"shape": 7,
|
||||
"type": "VRAM_MANAGEMENTARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vace_model",
|
||||
"shape": 7,
|
||||
"type": "VACEPATH",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "fantasytalking_model",
|
||||
"shape": 7,
|
||||
"type": "FANTASYTALKINGMODEL",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "multitalk_model",
|
||||
"shape": 7,
|
||||
"type": "MULTITALKMODEL",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "WANVIDEOMODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
29,
|
||||
103
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"cnr_id": "ComfyUI-WanVideoWrapper",
|
||||
"ver": "d9b1f4d1a5aea91d101ae97a54714a5861af3f50",
|
||||
"Node name for S&R": "WanVideoModelLoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"WanVideo\\Wan2_1-FLF2V-14B-720P_fp8_e4m3fn.safetensors",
|
||||
"fp16_fast",
|
||||
"fp8_e4m3fn",
|
||||
"offload_device",
|
||||
"sageattn"
|
||||
],
|
||||
"color": "#223",
|
||||
"bgcolor": "#335"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
@@ -2216,13 +2259,13 @@
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.6727499949326076,
|
||||
"scale": 0.6115909044841886,
|
||||
"offset": [
|
||||
359.0502881120043,
|
||||
1009.7385911003805
|
||||
710.2610924809787,
|
||||
967.8431584929548
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.23.4",
|
||||
"frontendVersion": "1.26.3",
|
||||
"node_versions": {
|
||||
"ComfyUI-WanVideoWrapper": "f8f423eceeadf2edcb58fab73701333e83ca733e",
|
||||
"comfy-core": "0.3.26",
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -48,26 +48,6 @@ def apply_lora(weight, lora, step=None):
|
||||
weight = weight.add(patch_diff, alpha=scale)
|
||||
return weight
|
||||
|
||||
|
||||
def linear_with_lora_and_scale_forward(cls, input):
|
||||
# Handles both scaled and unscaled, with or without LoRA
|
||||
has_scale = hasattr(cls, "scale_weight")
|
||||
weight = cls.weight.to(input.dtype)
|
||||
bias = cls.bias.to(input.dtype) if cls.bias is not None else None
|
||||
|
||||
if has_scale:
|
||||
scale_weight = cls.scale_weight.to(input.device)
|
||||
if weight.numel() < input.numel():
|
||||
weight = weight * scale_weight
|
||||
else:
|
||||
input = input * scale_weight
|
||||
|
||||
lora = getattr(cls, "lora", None)
|
||||
if lora is not None:
|
||||
weight = apply_lora(weight, lora, cls.step).to(input.dtype)
|
||||
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=None):
|
||||
log.info("FP8 matmul enabled")
|
||||
for name, submodule in module.named_modules():
|
||||
@@ -81,57 +61,4 @@ def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=
|
||||
original_forward = submodule.forward
|
||||
setattr(submodule, "original_forward", original_forward)
|
||||
setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input))
|
||||
|
||||
def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=None, params_to_keep={}):
|
||||
log.info("Patching Linear layers...")
|
||||
for name, submodule in module.named_modules():
|
||||
if not any(keyword in name for keyword in params_to_keep):
|
||||
# Set scale_weight if present
|
||||
if scale_weight_keys is not None:
|
||||
scale_key = f"{name}.scale_weight"
|
||||
if scale_key in scale_weight_keys:
|
||||
setattr(submodule, "scale_weight", scale_weight_keys[scale_key])
|
||||
|
||||
# Set LoRA if present
|
||||
if hasattr(submodule, "lora"):
|
||||
#print(f"removing old LoRA in {name}" )
|
||||
delattr(submodule, "lora")
|
||||
if patches is not None:
|
||||
patch_key1 = f"diffusion_model.{name}.weight"
|
||||
patch_key_compiled = f"diffusion_model.{name.replace('_orig_mod.', '')}.weight"
|
||||
patch = patches.get(patch_key1, []) or patches.get(patch_key_compiled, [])
|
||||
if len(patch) != 0:
|
||||
lora_diffs = []
|
||||
for p in patch:
|
||||
lora_obj = p[1]
|
||||
if "head" in name:
|
||||
continue # For now skip LoRA for head layers
|
||||
elif hasattr(lora_obj, "weights"):
|
||||
lora_diffs.append(lora_obj.weights)
|
||||
elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff":
|
||||
lora_diffs.append(lora_obj[1])
|
||||
else:
|
||||
continue
|
||||
lora_strengths = [p[0] for p in patch]
|
||||
lora = (lora_diffs, lora_strengths)
|
||||
setattr(submodule, "lora", lora)
|
||||
#print(f"Added LoRA to {name} with {len(lora_diffs)} diffs and strengths {lora_strengths}")
|
||||
|
||||
# Set forward if Linear and has either scale or lora
|
||||
if isinstance(submodule, nn.Linear):
|
||||
has_scale = hasattr(submodule, "scale_weight")
|
||||
has_lora = hasattr(submodule, "lora")
|
||||
if not hasattr(submodule, "original_forward"):
|
||||
setattr(submodule, "original_forward", submodule.forward)
|
||||
if has_scale or has_lora:
|
||||
setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_and_scale_forward(m, input))
|
||||
setattr(submodule, "step", 0) # Initialize step for LoRA scheduling
|
||||
|
||||
def remove_lora_from_module(module):
|
||||
unloaded = False
|
||||
for name, submodule in module.named_modules():
|
||||
if hasattr(submodule, "lora"):
|
||||
if not unloaded:
|
||||
log.info("Unloading all LoRAs")
|
||||
unloaded = True
|
||||
delattr(submodule, "lora")
|
||||
|
||||
+15
-3
@@ -1,13 +1,25 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
from diffusers.quantizers.gguf.utils import GGUFParameter, dequantize_gguf_tensor
|
||||
import gguf
|
||||
from diffusers.utils import is_accelerate_available
|
||||
from contextlib import nullcontext
|
||||
|
||||
from ..utils import log
|
||||
if is_accelerate_available():
|
||||
import accelerate
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
def load_gguf(model_path):
|
||||
from gguf import GGUFReader
|
||||
reader = GGUFReader(model_path)
|
||||
parsed_parameters = {}
|
||||
for tensor in reader.tensors:
|
||||
# if the tensor is a torch supported dtype do not use GGUFParameter
|
||||
is_gguf_quant = tensor.tensor_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16]
|
||||
meta_tensor = torch.empty(tensor.data.shape, dtype=torch.from_numpy(np.empty(0, dtype=tensor.data.dtype)).dtype, device='meta')
|
||||
parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor
|
||||
return parsed_parameters, reader
|
||||
|
||||
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
|
||||
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None):
|
||||
def _should_convert_to_gguf(state_dict, prefix):
|
||||
@@ -24,6 +36,7 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
|
||||
if (
|
||||
isinstance(module, nn.Linear)
|
||||
and not isinstance(module, GGUFLinear)
|
||||
and _should_convert_to_gguf(state_dict, module_prefix)
|
||||
and name not in modules_to_not_convert
|
||||
):
|
||||
@@ -42,7 +55,6 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
model._modules[name].source_cls = type(module)
|
||||
# Force requires_grad to False to avoid unexpected errors
|
||||
model._modules[name].requires_grad_(False)
|
||||
|
||||
return model
|
||||
|
||||
def set_lora_params_gguf(module, patches, module_prefix=""):
|
||||
|
||||
@@ -315,13 +315,13 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
return super().forward(x, encoder_hidden_states, shape)
|
||||
|
||||
N_t, N_h, N_w = shape
|
||||
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
|
||||
|
||||
|
||||
x_extra = None
|
||||
if x.shape[0] != encoder_hidden_states.shape[0]:
|
||||
if x.shape[0] * N_t != encoder_hidden_states.shape[0]:
|
||||
x_extra = x[:, -N_h * N_w:, :]
|
||||
x = x[:, :-N_h * N_w, :]
|
||||
N_t = N_t - 1
|
||||
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
|
||||
|
||||
# Query projection
|
||||
B, N, C = x.shape
|
||||
|
||||
+1
-8
@@ -94,14 +94,8 @@ class MultiTalkModelLoader:
|
||||
|
||||
def loadmodel(self, model, base_precision=None):
|
||||
from .multitalk import AudioProjModel
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
|
||||
if model_path.endswith(".gguf"):
|
||||
from diffusers.models.model_loading_utils import load_gguf_checkpoint
|
||||
sd = load_gguf_checkpoint(model_path)
|
||||
else:
|
||||
sd = load_torch_file(model_path, device=offload_device, safe_load=True)
|
||||
|
||||
audio_window=5
|
||||
intermediate_dim=512
|
||||
@@ -122,8 +116,7 @@ class MultiTalkModelLoader:
|
||||
|
||||
multitalk = {
|
||||
"proj_model": multitalk_proj_model,
|
||||
"sd": sd,
|
||||
"is_gguf": model_path.endswith(".gguf"),
|
||||
"model_path": model_path,
|
||||
"model_type": "InfiniteTalk" if "infinite" in model.lower() else "MultiTalk",
|
||||
}
|
||||
|
||||
|
||||
@@ -16,9 +16,10 @@ from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fouri
|
||||
add_noise_to_reference_video, optimized_scale, setup_radial_attention,
|
||||
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance)
|
||||
from .cache_methods.cache_methods import cache_report
|
||||
from .nodes_model_loading import load_weights
|
||||
from .enhance_a_video.globals import set_enhance_weight, set_num_frames
|
||||
from .taehv import TAEHV
|
||||
|
||||
from contextlib import nullcontext
|
||||
from einops import rearrange
|
||||
|
||||
from comfy import model_management as mm
|
||||
@@ -41,7 +42,18 @@ def offload_transformer(transformer):
|
||||
transformer.teacache_state.clear_all()
|
||||
transformer.magcache_state.clear_all()
|
||||
transformer.easycache_state.clear_all()
|
||||
transformer.to(offload_device)
|
||||
#transformer.to(offload_device)
|
||||
for name, param in transformer.named_parameters():
|
||||
module = transformer
|
||||
subnames = name.split('.')
|
||||
for subname in subnames[:-1]:
|
||||
module = getattr(module, subname)
|
||||
attr_name = subnames[-1]
|
||||
if param.data.is_floating_point():
|
||||
meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
|
||||
setattr(module, attr_name, meta_param)
|
||||
else:
|
||||
pass
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
@@ -348,8 +360,11 @@ class WanVideoTextEncode:
|
||||
raise ValueError("No cached text embeds found for prompts, please provide a T5 encoder.")
|
||||
|
||||
if model_to_offload is not None and device == "gpu":
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
try:
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
except:
|
||||
pass
|
||||
|
||||
encoder = t5["model"]
|
||||
dtype = t5["dtype"]
|
||||
@@ -782,6 +797,39 @@ class WanVideoAddStandInLatent:
|
||||
updated["standin_input"] = new_entry
|
||||
return (updated,)
|
||||
|
||||
class WanVideoAddMTVMotion:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"mtv_crafter_motion": ("MTVCRAFTERMOTION",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the MTV motion"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent to apply the ref "}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent to apply the ref "}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, mtv_crafter_motion, strength, start_percent, end_percent):
|
||||
# Prepare the new extra latent entry
|
||||
new_entry = {
|
||||
"mtv_motion_tokens": mtv_crafter_motion["mtv_motion_tokens"],
|
||||
"strength": strength,
|
||||
"start_percent": start_percent,
|
||||
"end_percent": end_percent,
|
||||
"global_mean": mtv_crafter_motion["global_mean"],
|
||||
"global_std": mtv_crafter_motion["global_std"]
|
||||
}
|
||||
|
||||
# Return a new dict with updated extra_latents
|
||||
updated = dict(embeds)
|
||||
updated["mtv_crafter_motion"] = new_entry
|
||||
return (updated,)
|
||||
|
||||
class WanVideoImageToVideoEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -1096,7 +1144,7 @@ class WanVideoPhantomEmbeds:
|
||||
|
||||
log.info(f"Phantom latents shape: {samples.shape}")
|
||||
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1 + T,
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
H * 8 // VAE_STRIDE[1],
|
||||
W * 8 // VAE_STRIDE[2])
|
||||
|
||||
@@ -1534,17 +1582,78 @@ class WanVideoScheduler: #WIP
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"scheduler": (scheduler_list, {"default": "unipc"}),
|
||||
"steps": ("INT", {"default": 30, "min": 1, "tooltip": "Number of steps for the scheduler"}),
|
||||
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||
"start_step": ("INT", {"default": 0, "min": 0, "tooltip": "Starting step for the scheduler"}),
|
||||
"end_step": ("INT", {"default": -1, "min": -1, "tooltip": "Ending step for the scheduler"})
|
||||
},
|
||||
"optional": {
|
||||
"sigmas": ("SIGMAS", ),
|
||||
},
|
||||
"hidden": {
|
||||
"unique_id": "UNIQUE_ID",
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = (scheduler_list, )
|
||||
RETURN_NAMES = ("scheduler",)
|
||||
RETURN_TYPES = ("SIGMAS", "INT", "FLOAT", scheduler_list, "INT", "INT",)
|
||||
RETURN_NAMES = ("sigmas", "steps", "shift", "scheduler", "start_step", "end_step")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def process(self, scheduler):
|
||||
return (scheduler,)
|
||||
def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None):
|
||||
sample_scheduler, timesteps = get_scheduler(
|
||||
scheduler,
|
||||
steps,
|
||||
start_step, end_step, shift,
|
||||
device,
|
||||
sigmas=sigmas)
|
||||
|
||||
scheduler_dict = {
|
||||
"sample_scheduler": sample_scheduler,
|
||||
"timesteps": timesteps,
|
||||
}
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
import io
|
||||
import base64
|
||||
import matplotlib.pyplot as plt
|
||||
except:
|
||||
PromptServer = None
|
||||
if unique_id and PromptServer is not None:
|
||||
try:
|
||||
# Plot sigmas and save to a buffer
|
||||
sigmas_np = sample_scheduler.full_sigmas[:-1].cpu().numpy()
|
||||
buf = io.BytesIO()
|
||||
fig = plt.figure(facecolor='#353535')
|
||||
ax = fig.add_subplot(111)
|
||||
ax.set_facecolor('#353535') # Set axes background color
|
||||
ax.plot(sigmas_np)
|
||||
ax.set_title("Sigmas", color='white') # Title font color
|
||||
ax.set_xlabel("Step", color='white') # X label font color
|
||||
ax.set_ylabel("Sigma Value", color='white') # Y label font color
|
||||
ax.tick_params(axis='x', colors='white') # X tick color
|
||||
ax.tick_params(axis='y', colors='white') # Y tick color
|
||||
# Add split point if end_step is defined
|
||||
if end_step != -1 and 0 <= end_step < len(sigmas_np):
|
||||
ax.axvline(end_step, color='red', linestyle='--', linewidth=2, label='end_step split')
|
||||
ax.legend()
|
||||
plt.tight_layout()
|
||||
plt.savefig(buf, format='png')
|
||||
plt.close(fig)
|
||||
buf.seek(0)
|
||||
img_base64 = base64.b64encode(buf.read()).decode('utf-8')
|
||||
buf.close()
|
||||
|
||||
# Send as HTML img tag with base64 data
|
||||
html_img = f"<img src='data:image/png;base64,{img_base64}' alt='Sigmas Plot' style='max-width:100%; height:100%; overflow:hidden; display:block;'>"
|
||||
PromptServer.instance.send_progress_text(html_img, unique_id)
|
||||
except Exception as e:
|
||||
print("Failed to send sigmas plot:", e)
|
||||
pass
|
||||
|
||||
return (sigmas, steps, shift, scheduler_dict, start_step, end_step)
|
||||
|
||||
rope_functions = ["default", "comfy", "comfy_chunked"]
|
||||
class WanVideoRoPEFunction:
|
||||
@@ -1631,28 +1740,43 @@ class WanVideoSampler:
|
||||
model = model.model
|
||||
transformer = model.diffusion_model
|
||||
|
||||
dtype = model["dtype"]
|
||||
dtype = model["base_dtype"]
|
||||
weight_dtype = model["weight_dtype"]
|
||||
fp8_matmul = model["fp8_matmul"]
|
||||
gguf = model["gguf"]
|
||||
gguf_reader = model["gguf_reader"]
|
||||
control_lora = model["control_lora"]
|
||||
|
||||
transformer_options = patcher.model_options.get("transformer_options", None)
|
||||
merge_loras = transformer_options["merge_loras"]
|
||||
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
transformer.blocks_to_swap = block_swap_args.get("blocks_to_swap", 0)
|
||||
transformer.vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", 0)
|
||||
transformer.prefetch_blocks = block_swap_args.get("prefetch_blocks", 0)
|
||||
transformer.block_swap_debug = block_swap_args.get("block_swap_debug", False)
|
||||
transformer.offload_img_emb = block_swap_args.get("offload_img_emb", False)
|
||||
transformer.offload_txt_emb = block_swap_args.get("offload_txt_emb", False)
|
||||
|
||||
is_5b = transformer.out_dim == 48
|
||||
vae_upscale_factor = 16 if is_5b else 8
|
||||
|
||||
patch_linear = transformer_options.get("patch_linear", False)
|
||||
# Load weights
|
||||
if transformer.patched_linear and gguf_reader is None:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
|
||||
|
||||
if gguf:
|
||||
if gguf_reader is not None: #handle GGUF
|
||||
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
|
||||
set_lora_params_gguf(transformer, patcher.patches)
|
||||
elif len(patcher.patches) != 0 and patch_linear:
|
||||
transformer.patched_linear = True
|
||||
elif len(patcher.patches) != 0 and transformer.patched_linear: #handle patched linear layers (unmerged loras, fp8 scaled)
|
||||
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
|
||||
if not merge_loras and fp8_matmul:
|
||||
raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported")
|
||||
set_lora_params(transformer, patcher.patches)
|
||||
else:
|
||||
remove_lora_from_module(transformer)
|
||||
remove_lora_from_module(transformer) #clear possible unmerged lora weights
|
||||
|
||||
transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False)
|
||||
|
||||
@@ -1681,8 +1805,11 @@ class WanVideoSampler:
|
||||
|
||||
#region Scheduler
|
||||
sample_scheduler = None
|
||||
if scheduler != "multitalk":
|
||||
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
|
||||
if isinstance(scheduler, dict):
|
||||
sample_scheduler = scheduler["sample_scheduler"]
|
||||
timesteps = scheduler["timesteps"]
|
||||
elif scheduler != "multitalk":
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
log.info(f"sigmas: {sample_scheduler.sigmas}")
|
||||
else:
|
||||
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
|
||||
@@ -1698,7 +1825,11 @@ class WanVideoSampler:
|
||||
start_step = steps - int(steps * denoise_strength) - 1
|
||||
add_noise_to_samples = True #for now to not break old workflows
|
||||
|
||||
noise_pred_flipped = None
|
||||
scheduler_step_args = {"generator": seed_g}
|
||||
step_sig = inspect.signature(sample_scheduler.step)
|
||||
for arg in list(scheduler_step_args.keys()):
|
||||
if arg not in step_sig.parameters:
|
||||
scheduler_step_args.pop(arg)
|
||||
|
||||
if isinstance(cfg, list):
|
||||
if steps < len(cfg):
|
||||
@@ -1715,7 +1846,7 @@ class WanVideoSampler:
|
||||
vace_data = vace_context = vace_scale = None
|
||||
fun_or_fl2v_model = has_ref = drop_last = False
|
||||
phantom_latents = fun_ref_image = ATI_tracks = None
|
||||
add_cond = attn_cond = attn_cond_neg = None
|
||||
add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None
|
||||
|
||||
#I2V
|
||||
image_cond = image_embeds.get("image_embeds", None)
|
||||
@@ -1902,8 +2033,6 @@ class WanVideoSampler:
|
||||
phantom_cfg_scale = [phantom_cfg_scale] * (steps +1)
|
||||
phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0)
|
||||
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
|
||||
if phantom_latents is not None:
|
||||
phantom_latents = phantom_latents.to(device)
|
||||
|
||||
latent_video_length = noise.shape[1]
|
||||
|
||||
@@ -2002,7 +2131,7 @@ class WanVideoSampler:
|
||||
"start_percent": fantasy_portrait_embeds.get("start_percent", 0.0),
|
||||
"end_percent": fantasy_portrait_embeds.get("end_percent", 1.0),
|
||||
}
|
||||
|
||||
|
||||
# MiniMax Remover
|
||||
minimax_latents = minimax_mask_latents = None
|
||||
minimax_latents = image_embeds.get("minimax_latents", None)
|
||||
@@ -2058,6 +2187,31 @@ class WanVideoSampler:
|
||||
self.window_tracker = WindowTracker(verbose=context_options["verbose"])
|
||||
context = get_context_scheduler(context_schedule)
|
||||
|
||||
#MTV Crafter
|
||||
mtv_input = image_embeds.get("mtv_crafter_motion", None)
|
||||
mtv_motion_tokens = 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
|
||||
noise_mask=original_image=None
|
||||
if samples is not None and not multitalk_sampling:
|
||||
@@ -2159,43 +2313,39 @@ class WanVideoSampler:
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
#region transformer settings
|
||||
if transformer_options is not None:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
|
||||
#blockswap init
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
if not transformer.patched_linear:
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
prefetch_blocks = block_swap_args.get("prefetch_blocks", 0),
|
||||
block_swap_debug = block_swap_args.get("block_swap_debug", False),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
|
||||
elif model["manual_offloading"]:
|
||||
transformer.to(device)
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
prefetch_blocks = block_swap_args.get("prefetch_blocks", 0),
|
||||
block_swap_debug = block_swap_args.get("block_swap_debug", False),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
else:
|
||||
transformer.to(device)
|
||||
|
||||
# Initialize Cache if enabled
|
||||
previous_cache_states = None
|
||||
@@ -2302,7 +2452,9 @@ class WanVideoSampler:
|
||||
import copy
|
||||
sample_scheduler_flipped = copy.deepcopy(sample_scheduler)
|
||||
|
||||
#rope
|
||||
# Rotary positional embeddings (RoPE)
|
||||
|
||||
# RoPE base freq scaling as used with CineScale
|
||||
ntk_alphas = [1.0, 1.0, 1.0]
|
||||
if isinstance(rope_function, dict):
|
||||
ntk_alphas = rope_function["ntk_scale_f"], rope_function["ntk_scale_h"], rope_function["ntk_scale_w"]
|
||||
@@ -2316,7 +2468,7 @@ class WanVideoSampler:
|
||||
freqs = None
|
||||
transformer.rope_embedder.k = None
|
||||
transformer.rope_embedder.num_frames = None
|
||||
if "default" in rope_function or bidirectional_sampling:
|
||||
if "default" in rope_function or bidirectional_sampling: # original RoPE
|
||||
d = transformer.dim // transformer.num_heads
|
||||
freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
|
||||
@@ -2324,7 +2476,7 @@ class WanVideoSampler:
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
],
|
||||
dim=1)
|
||||
elif "comfy" in rope_function:
|
||||
elif "comfy" in rope_function: # comfy's rope
|
||||
transformer.rope_embedder.k = riflex_freq_index
|
||||
transformer.rope_embedder.num_frames = latent_video_length
|
||||
|
||||
@@ -2338,10 +2490,12 @@ class WanVideoSampler:
|
||||
#region model pred
|
||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
|
||||
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
|
||||
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False):
|
||||
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False,
|
||||
mtv_motion_tokens=None):
|
||||
nonlocal transformer
|
||||
z = z.to(dtype)
|
||||
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
|
||||
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
|
||||
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype) if autocast_enabled else nullcontext():
|
||||
|
||||
if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init:
|
||||
return z*0, None
|
||||
@@ -2400,20 +2554,22 @@ class WanVideoSampler:
|
||||
|
||||
if recammaster is not None:
|
||||
z = torch.cat([z, recam_latents.to(z)], dim=1)
|
||||
|
||||
|
||||
if mtv_input is not None:
|
||||
if ((mtv_start_percent <= current_step_percentage <= mtv_end_percent) or \
|
||||
(mtv_end_percent > 0 and idx == 0 and current_step_percentage >= mtv_start_percent)):
|
||||
mtv_motion_tokens = mtv_motion_tokens.to(z)
|
||||
mtv_motion_rotary_emb = motion_rotary_emb
|
||||
|
||||
use_phantom = False
|
||||
phantom_ref = None
|
||||
if phantom_latents is not None:
|
||||
if (phantom_start_percent <= current_step_percentage <= phantom_end_percent) or \
|
||||
(phantom_end_percent > 0 and idx == 0 and current_step_percentage >= phantom_start_percent):
|
||||
|
||||
z_pos = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1)
|
||||
z_phantom_img = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1)
|
||||
z_neg = torch.cat([z[:,:-phantom_latents.shape[1]], torch.zeros_like(phantom_latents).to(z)], dim=1)
|
||||
phantom_ref = phantom_latents.to(z)
|
||||
use_phantom = True
|
||||
if cache_state is not None and len(cache_state) != 3:
|
||||
cache_state.append(None)
|
||||
if not use_phantom:
|
||||
z_pos = z_neg = z
|
||||
|
||||
if controlnet_latents is not None:
|
||||
if (controlnet_start <= current_step_percentage < controlnet_end):
|
||||
@@ -2439,9 +2595,9 @@ class WanVideoSampler:
|
||||
|
||||
if minimax_latents is not None:
|
||||
if context_window is not None:
|
||||
z_pos = z_neg = torch.cat([z, minimax_latents[:, context_window], minimax_mask_latents[:, context_window]], dim=0)
|
||||
z = torch.cat([z, minimax_latents[:, context_window], minimax_mask_latents[:, context_window]], dim=0)
|
||||
else:
|
||||
z_pos = z_neg = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
|
||||
z = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
|
||||
|
||||
if not multitalk_sampling and multitalk_audio_embedding is not None:
|
||||
audio_embedding = multitalk_audio_embedding
|
||||
@@ -2482,32 +2638,37 @@ class WanVideoSampler:
|
||||
|
||||
|
||||
base_params = {
|
||||
'seq_len': seq_len,
|
||||
'device': device,
|
||||
'freqs': freqs,
|
||||
't': timestep,
|
||||
'current_step': idx,
|
||||
'last_step': len(timesteps) - 1 == idx,
|
||||
'control_lora_enabled': control_lora_enabled,
|
||||
'enhance_enabled': enhance_enabled,
|
||||
'camera_embed': camera_embed,
|
||||
'unianim_data': unianim_data,
|
||||
'fun_ref': fun_ref_input if fun_ref_image is not None else None,
|
||||
'fun_camera': control_camera_input if control_camera_latents is not None else None,
|
||||
'audio_proj': audio_proj if fantasytalking_embeds is not None else None,
|
||||
'audio_scale': audio_scale,
|
||||
"pcd_data": pcd_data_input,
|
||||
"controlnet": controlnet,
|
||||
"add_cond": add_cond_input,
|
||||
"nag_params": text_embeds.get("nag_params", {}),
|
||||
"nag_context": text_embeds.get("nag_prompt_embeds", None),
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None,
|
||||
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None,
|
||||
"inner_t": [shot_len] if shot_len else None,
|
||||
"standin_input": standin_input,
|
||||
"fantasy_portrait_input": fantasy_portrait_input,
|
||||
"reverse_time": reverse_time,
|
||||
"ntk_alphas": ntk_alphas
|
||||
'seq_len': seq_len, # sequence length
|
||||
'device': device, # main device
|
||||
'freqs': freqs, # rope freqs
|
||||
't': timestep, # current timestep
|
||||
'current_step': idx, # current step
|
||||
'last_step': len(timesteps) - 1 == idx, # is last step
|
||||
'control_lora_enabled': control_lora_enabled, # control lora toggle for patch embed selection
|
||||
'enhance_enabled': enhance_enabled, # enhance-a-video toggle
|
||||
'camera_embed': camera_embed, # recammaster embedding
|
||||
'unianim_data': unianim_data, # unianimate input
|
||||
'fun_ref': fun_ref_input if fun_ref_image is not None else None, # Fun model reference latent
|
||||
'fun_camera': control_camera_input if control_camera_latents is not None else None, # Fun model camera embed
|
||||
'audio_proj': audio_proj if fantasytalking_embeds is not None else None, # FantasyTalking audio projection
|
||||
'audio_scale': audio_scale, # FantasyTalking audio scale
|
||||
"pcd_data": pcd_data_input, # Uni3C input
|
||||
"controlnet": controlnet, # TheDenk's controlnet input
|
||||
"add_cond": add_cond_input, # additional conditioning input
|
||||
"nag_params": text_embeds.get("nag_params", {}), # normalized attention guidance
|
||||
"nag_context": text_embeds.get("nag_prompt_embeds", None), # normalized attention guidance context
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk audio input
|
||||
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk reference target masks
|
||||
"inner_t": [shot_len] if shot_len else None, # inner timestep for EchoShot
|
||||
"standin_input": standin_input, # Stand-in reference input
|
||||
"fantasy_portrait_input": fantasy_portrait_input, # Fantasy portrait input
|
||||
"phantom_ref": phantom_ref, # Phantom reference input
|
||||
"reverse_time": reverse_time, # Reverse RoPE toggle
|
||||
"ntk_alphas": ntk_alphas, # RoPE freq scaling values
|
||||
"mtv_motion_tokens": mtv_motion_tokens if mtv_input is not None else None, # MTV-Crafter motion tokens
|
||||
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None, # MTV-Crafter RoPE
|
||||
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, # MTV-Crafter scaling
|
||||
"mtv_freqs": mtv_freqs if mtv_input is not None else None, # MTV-Crafter extra RoPE freqs
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -2522,7 +2683,7 @@ class WanVideoSampler:
|
||||
if not batched_cfg:
|
||||
#cond
|
||||
noise_pred_cond, cache_state_cond = transformer(
|
||||
[z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
[z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
pred_id=cache_state[0] if cache_state else None,
|
||||
vace_data=vace_data, attn_cond=attn_cond,
|
||||
@@ -2543,7 +2704,7 @@ class WanVideoSampler:
|
||||
if not math.isclose(audio_cfg_scale[idx], 1.0):
|
||||
base_params['audio_proj'] = None
|
||||
noise_pred_uncond, cache_state_uncond = transformer(
|
||||
[z_neg], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
|
||||
[z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
|
||||
y=[image_cond_input] if image_cond_input is not None else None,
|
||||
is_uncond=True, current_step_percentage=current_step_percentage,
|
||||
pred_id=cache_state[1] if cache_state else None,
|
||||
@@ -2554,7 +2715,7 @@ class WanVideoSampler:
|
||||
#phantom
|
||||
if use_phantom and not math.isclose(phantom_cfg_scale[idx], 1.0):
|
||||
noise_pred_phantom, cache_state_phantom = transformer(
|
||||
[z_phantom_img], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
|
||||
[z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea,
|
||||
y=[image_cond_input] if image_cond_input is not None else None,
|
||||
is_uncond=True, current_step_percentage=current_step_percentage,
|
||||
pred_id=cache_state[2] if cache_state else None,
|
||||
@@ -2572,7 +2733,7 @@ class WanVideoSampler:
|
||||
cache_state.append(None)
|
||||
base_params['audio_proj'] = None
|
||||
noise_pred_no_audio, cache_state_audio = transformer(
|
||||
[z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
[z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
pred_id=cache_state[2] if cache_state else None,
|
||||
vace_data=vace_data,
|
||||
@@ -2591,7 +2752,7 @@ class WanVideoSampler:
|
||||
cache_state.append(None)
|
||||
base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:]
|
||||
noise_pred_no_audio, cache_state_audio = transformer(
|
||||
[z_pos], context=negative_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
[z], context=negative_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
pred_id=cache_state[2] if cache_state else None,
|
||||
vace_data=vace_data,
|
||||
@@ -2618,7 +2779,7 @@ class WanVideoSampler:
|
||||
except Exception as e:
|
||||
log.error(f"Error during model prediction: {e}")
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
if not model["auto_cpu_offload"]:
|
||||
offload_transformer(transformer)
|
||||
raise e
|
||||
|
||||
@@ -2701,7 +2862,7 @@ class WanVideoSampler:
|
||||
# FreeInit noise reinitialization (after first iteration)
|
||||
if freeinit_args is not None and iter_idx > 0:
|
||||
# restart scheduler for each iteration
|
||||
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
|
||||
# Re-apply start_step and end_step logic to timesteps and sigmas
|
||||
if end_step != -1:
|
||||
@@ -3010,7 +3171,16 @@ class WanVideoSampler:
|
||||
"start_percent": unianimate_poses["start_percent"],
|
||||
"end_percent": unianimate_poses["end_percent"]
|
||||
}
|
||||
|
||||
|
||||
partial_mtv_motion_tokens = None
|
||||
if mtv_input is not None:
|
||||
start_token_index = c[0] * 24
|
||||
end_token_index = (c[-1] + 1) * 24
|
||||
partial_mtv_motion_tokens = mtv_motion_tokens[:, start_token_index:end_token_index, :]
|
||||
if context_options["verbose"]:
|
||||
log.info(f"context window: {c}")
|
||||
log.info(f"motion_token_indices: {start_token_index}-{end_token_index}")
|
||||
|
||||
partial_add_cond = None
|
||||
if add_cond is not None:
|
||||
partial_add_cond = add_cond[:, :, c].to(device, dtype)
|
||||
@@ -3027,7 +3197,8 @@ class WanVideoSampler:
|
||||
cfg[idx], positive,
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
partial_timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
|
||||
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input)
|
||||
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input,
|
||||
mtv_motion_tokens=partial_mtv_motion_tokens)
|
||||
|
||||
if cache_args is not None:
|
||||
self.window_tracker.cache_states[window_id] = new_teacache
|
||||
@@ -3208,7 +3379,7 @@ class WanVideoSampler:
|
||||
timesteps = [torch.tensor([t], device=device) for t in timesteps]
|
||||
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
|
||||
else:
|
||||
sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g)
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
|
||||
|
||||
# sample videos
|
||||
@@ -3223,36 +3394,36 @@ class WanVideoSampler:
|
||||
|
||||
if offload:
|
||||
#blockswap init
|
||||
if transformer_options is not None:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
if not transformer.patched_linear:
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
)
|
||||
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
elif model["manual_offloading"]:
|
||||
transformer.to(device)
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
else:
|
||||
transformer.to(device)
|
||||
|
||||
# Use the appropriate prompt for this section
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
@@ -3373,11 +3544,7 @@ class WanVideoSampler:
|
||||
|
||||
del noise, latent_motion_frames
|
||||
if offload:
|
||||
transformer.to(offload_device)
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
offload_transformer(transformer)
|
||||
vae.to(device)
|
||||
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
|
||||
vae.model.clear_cache()
|
||||
@@ -3465,7 +3632,7 @@ class WanVideoSampler:
|
||||
text_embeds["prompt_embeds"],
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input)
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens)
|
||||
if bidirectional_sampling:
|
||||
noise_pred_flipped, self.cache_state = predict_with_cfg(
|
||||
latent_model_input_flipped,
|
||||
@@ -3473,7 +3640,7 @@ class WanVideoSampler:
|
||||
text_embeds["prompt_embeds"],
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, reverse_time=True)
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens,reverse_time=True)
|
||||
|
||||
if latent_shift_loop:
|
||||
#reverse latent shift
|
||||
@@ -3547,8 +3714,8 @@ class WanVideoSampler:
|
||||
if callback is not None:
|
||||
if recammaster is not None:
|
||||
callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach()
|
||||
elif phantom_latents is not None:
|
||||
callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
|
||||
#elif phantom_latents is not None:
|
||||
# callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
|
||||
else:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
|
||||
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
|
||||
@@ -3563,7 +3730,7 @@ class WanVideoSampler:
|
||||
except Exception as e:
|
||||
log.error(f"Error during sampling: {e}")
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
if not model["auto_cpu_offload"]:
|
||||
offload_transformer(transformer)
|
||||
raise e
|
||||
|
||||
@@ -3582,7 +3749,7 @@ class WanVideoSampler:
|
||||
}
|
||||
|
||||
if force_offload:
|
||||
if model["manual_offloading"]:
|
||||
if not model["auto_cpu_offload"]:
|
||||
offload_transformer(transformer)
|
||||
|
||||
try:
|
||||
@@ -3792,6 +3959,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoScheduler": WanVideoScheduler,
|
||||
"WanVideoAddStandInLatent": WanVideoAddStandInLatent,
|
||||
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
|
||||
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,
|
||||
"WanVideoRoPEFunction": WanVideoRoPEFunction,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -3825,5 +3993,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddExtraLatent": "WanVideo Add Extra Latent",
|
||||
"WanVideoAddStandInLatent": "WanVideo Add StandIn Latent",
|
||||
"WanVideoAddControlEmbeds": "WanVideo Add Control Embeds",
|
||||
"WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion",
|
||||
"WanVideoRoPEFunction": "WanVideo RoPE Function",
|
||||
}
|
||||
|
||||
+332
-226
@@ -18,6 +18,11 @@ import comfy.model_management as mm
|
||||
from comfy.utils import load_torch_file, ProgressBar
|
||||
import comfy.model_base
|
||||
from comfy.sd import load_lora_for_models
|
||||
try:
|
||||
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
|
||||
from gguf import GGMLQuantizationType
|
||||
except:
|
||||
pass
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@@ -392,7 +397,7 @@ class WanVideoLoraSelect:
|
||||
with safe_open(lora_path, framework="pt", device="cpu") as f:
|
||||
metadata = f.metadata()
|
||||
except Exception as e:
|
||||
print(f"Could not load metadata from {lora}: {e}")
|
||||
log.info(f"Could not load metadata from {lora}: {e}")
|
||||
|
||||
if unique_id and PromptServer is not None:
|
||||
try:
|
||||
@@ -419,7 +424,7 @@ class WanVideoLoraSelect:
|
||||
unique_id
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Error displaying metadata: {e}")
|
||||
log.warning(f"Error displaying metadata: {e}")
|
||||
pass
|
||||
|
||||
lora = {
|
||||
@@ -508,7 +513,7 @@ class WanVideoVACEModelSelect:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VACEPATH",)
|
||||
RETURN_NAMES = ("vace_model", )
|
||||
RETURN_NAMES = ("extra_model", )
|
||||
FUNCTION = "getvacepath"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "VACE model to use when not using model that has it included, loaded from 'ComfyUI/models/diffusion_models'"
|
||||
@@ -518,6 +523,27 @@ class WanVideoVACEModelSelect:
|
||||
"path": folder_paths.get_full_path("diffusion_models", vace_model),
|
||||
}
|
||||
return (vace_model,)
|
||||
|
||||
class WanVideoExtraModelSelect:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"extra_model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' path to extra state dict to add to the main model"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VACEPATH",)
|
||||
RETURN_NAMES = ("extra_model", )
|
||||
FUNCTION = "getvacepath"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Extra model to load and add to the main model, ie. VACE or MTV Crafter 'ComfyUI/models/diffusion_models'"
|
||||
|
||||
def getvacepath(self, extra_model):
|
||||
extra_model = {
|
||||
"path": folder_paths.get_full_path("diffusion_models", extra_model),
|
||||
}
|
||||
return (extra_model,)
|
||||
|
||||
class WanVideoLoraBlockEdit:
|
||||
def __init__(self):
|
||||
@@ -701,13 +727,196 @@ class WanVideoSetLoRAs:
|
||||
|
||||
del lora_sd
|
||||
|
||||
if 'transformer_options' not in patcher.model_options:
|
||||
patcher.model_options['transformer_options'] = {}
|
||||
|
||||
patcher.model_options['transformer_options']["patch_linear"] = True
|
||||
|
||||
return (patcher,)
|
||||
|
||||
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
|
||||
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
block_idx = vace_block_idx = None
|
||||
|
||||
if gguf:
|
||||
log.info("Using GGUF to load and assign model weights to device...")
|
||||
|
||||
# Prepare sd from GGUF readers
|
||||
|
||||
# UniAnimate embedding weight workaround
|
||||
unianimate_sd = {}
|
||||
for key in sd.keys():
|
||||
if "dwpose_embedding" in key or "randomref_embedding_pose" in key:
|
||||
unianimate_sd[key] = sd[key]
|
||||
|
||||
sd = {}
|
||||
all_tensors = []
|
||||
for r in reader:
|
||||
all_tensors.extend(r.tensors)
|
||||
for tensor in all_tensors:
|
||||
load_device = device
|
||||
if "vace_blocks." in tensor.name:
|
||||
try:
|
||||
vace_block_idx = int(tensor.name.split("vace_blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
vace_block_idx = None
|
||||
elif "blocks." in tensor.name:
|
||||
try:
|
||||
block_idx = int(tensor.name.split("blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
block_idx = None
|
||||
|
||||
if block_swap_args is not None:
|
||||
if block_idx is not None:
|
||||
if block_idx >= len(transformer.blocks) - block_swap_args.get("blocks_to_swap", 0):
|
||||
load_device = offload_device
|
||||
elif vace_block_idx is not None:
|
||||
if vace_block_idx >= len(transformer.vace_blocks) - block_swap_args.get("vace_blocks_to_swap", 0):
|
||||
load_device = offload_device
|
||||
|
||||
is_gguf_quant = tensor.tensor_type not in [GGMLQuantizationType.F32, GGMLQuantizationType.F16]
|
||||
weights = torch.from_numpy(tensor.data.copy()).to(load_device)
|
||||
sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
|
||||
sd.update(unianimate_sd)
|
||||
del unianimate_sd
|
||||
|
||||
if not getattr(transformer, "gguf_patched", False):
|
||||
transformer = _replace_with_gguf_linear(
|
||||
transformer, base_dtype, sd, patches=patcher.patches
|
||||
)
|
||||
transformer.gguf_patched = True
|
||||
else:
|
||||
log.info("Using accelerate to load and assign model weights to device...")
|
||||
named_params = transformer.named_parameters()
|
||||
|
||||
for name, param in tqdm(named_params,
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
total=param_count,
|
||||
leave=True):
|
||||
block_idx = vace_block_idx = None
|
||||
if "vace_blocks." in name:
|
||||
try:
|
||||
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
vace_block_idx = None
|
||||
elif "blocks." in name:
|
||||
try:
|
||||
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
block_idx = None
|
||||
|
||||
if "loras" in name:
|
||||
continue
|
||||
|
||||
# GGUF: skip GGUFParameter params
|
||||
if gguf and isinstance(param, GGUFParameter):
|
||||
continue
|
||||
|
||||
if gguf:
|
||||
dtype_to_use = torch.float32 if "patch_embedding" in name else base_dtype
|
||||
else:
|
||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else weight_dtype
|
||||
dtype_to_use = weight_dtype if sd[name.replace("_orig_mod.", "")].dtype == weight_dtype else dtype_to_use
|
||||
if "modulation" in name or "norm" in name or "bias" in name or "img_emb" in name:
|
||||
dtype_to_use = base_dtype
|
||||
if "patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
|
||||
load_device = device
|
||||
if block_swap_args is not None:
|
||||
if block_idx is not None:
|
||||
if block_idx >= len(transformer.blocks) - block_swap_args.get("blocks_to_swap", 0):
|
||||
load_device = offload_device
|
||||
elif vace_block_idx is not None:
|
||||
if vace_block_idx >= len(transformer.vace_blocks) - block_swap_args.get("vace_blocks_to_swap", 0):
|
||||
load_device = offload_device
|
||||
# Set tensor to device
|
||||
set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=sd[name.replace("_orig_mod.", "")])
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
|
||||
pbar.update_absolute(param_count)
|
||||
pbar.update_absolute(0)
|
||||
|
||||
def patch_control_lora(transformer, device):
|
||||
log.info("Control-LoRA detected, patching model...")
|
||||
|
||||
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
|
||||
old_in_dim = transformer.in_dim # 16
|
||||
new_in_dim = 32
|
||||
|
||||
new_in = in_cls(
|
||||
new_in_dim,
|
||||
transformer.patch_embedding.out_channels,
|
||||
transformer.patch_embedding.kernel_size,
|
||||
transformer.patch_embedding.stride,
|
||||
transformer.patch_embedding.padding,
|
||||
).to(device=device, dtype=torch.float32)
|
||||
|
||||
new_in.weight.zero_()
|
||||
new_in.bias.zero_()
|
||||
|
||||
new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight)
|
||||
new_in.bias.copy_(transformer.patch_embedding.bias)
|
||||
|
||||
transformer.patch_embedding = new_in
|
||||
transformer.expanded_patch_embedding = new_in
|
||||
|
||||
def patch_stand_in_lora(transformer, lora_sd, transformer_load_device, base_dtype, lora_strength):
|
||||
if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
|
||||
log.info("Stand-In LoRA detected")
|
||||
for block in transformer.blocks:
|
||||
block.self_attn.q_loras = LoRALinearLayer(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.k_loras = LoRALinearLayer(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.v_loras = LoRALinearLayer(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
for lora in [block.self_attn.q_loras, block.self_attn.k_loras, block.self_attn.v_loras]:
|
||||
for param in lora.parameters():
|
||||
param.requires_grad = False
|
||||
for name, param in transformer.named_parameters():
|
||||
if "lora" in name:
|
||||
param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype))
|
||||
|
||||
def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
unianimate_sd = None
|
||||
#spacepxl's control LoRA patch
|
||||
for l in lora:
|
||||
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
|
||||
lora_path = l["path"]
|
||||
lora_strength = l["strength"]
|
||||
if isinstance(lora_strength, list):
|
||||
if merge_loras:
|
||||
raise ValueError("LoRA strength should be a single value when merge_loras=True")
|
||||
patcher.model.diffusion_model.lora_scheduling_enabled = True
|
||||
if lora_strength == 0:
|
||||
log.warning(f"LoRA {lora_path} has strength 0, skipping...")
|
||||
continue
|
||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
||||
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
|
||||
from .unianimate.nodes import update_transformer
|
||||
log.info("Unianimate LoRA detected, patching model...")
|
||||
patcher.model.diffusion_model, unianimate_sd = update_transformer(patcher.model.diffusion_model, lora_sd)
|
||||
|
||||
lora_sd = standardize_lora_key_format(lora_sd)
|
||||
|
||||
if l["blocks"]:
|
||||
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
|
||||
|
||||
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
|
||||
#if not any('img' in k for k in sd.keys()):
|
||||
# lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
|
||||
control_lora=False
|
||||
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
|
||||
control_lora = True
|
||||
#stand-in LoRA patch
|
||||
if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
|
||||
patch_stand_in_lora(patcher.model.diffusion_model, lora_sd, device, base_dtype, lora_strength)
|
||||
# normal LoRA patch
|
||||
else:
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
del lora_sd
|
||||
return patcher, control_lora, unianimate_sd
|
||||
|
||||
#region Model loading
|
||||
class WanVideoModelLoader:
|
||||
@classmethod
|
||||
@@ -718,7 +927,7 @@ class WanVideoModelLoader:
|
||||
|
||||
"base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}),
|
||||
"quantization": (["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e4m3fn_scaled", "fp8_e4m3fn_scaled_fast", "fp8_e5m2", "fp8_e5m2_fast", "fp8_e5m2_scaled", "fp8_e5m2_scaled_fast"], {"default": "disabled", "tooltip": "optional quantization method"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
},
|
||||
"optional": {
|
||||
"attention_mode": ([
|
||||
@@ -734,7 +943,7 @@ class WanVideoModelLoader:
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
"lora": ("WANVIDLORA", {"default": None}),
|
||||
"vram_management_args": ("VRAM_MANAGEMENTARGS", {"default": None, "tooltip": "Alternative offloading method from DiffSynth-Studio, more aggressive in reducing memory use than block swapping, but can be slower"}),
|
||||
"vace_model": ("VACEPATH", {"default": None, "tooltip": "VACE model to use when not using model that has it included"}),
|
||||
"extra_model": ("VACEPATH", {"default": None, "tooltip": "Extra model to add to the main model, ie. VACE or MTV Crafter"}),
|
||||
"fantasytalking_model": ("FANTASYTALKINGMODEL", {"default": None, "tooltip": "FantasyTalking model https://github.com/Fantasy-AMAP"}),
|
||||
"multitalk_model": ("MULTITALKMODEL", {"default": None, "tooltip": "Multitalk model"}),
|
||||
"fantasyportrait_model": ("FANTASYPORTRAITMODEL", {"default": None, "tooltip": "FantasyPortrait model"}),
|
||||
@@ -747,10 +956,11 @@ class WanVideoModelLoader:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, model, base_precision, load_device, quantization,
|
||||
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None,
|
||||
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, extra_model=None, vace_model=None,
|
||||
fantasytalking_model=None, multitalk_model=None, fantasyportrait_model=None):
|
||||
assert not (vram_management_args is not None and block_swap_args is not None), "Can't use both block_swap_args and vram_management_args at the same time"
|
||||
|
||||
if vace_model is not None:
|
||||
extra_model = vace_model
|
||||
lora_low_mem_load = merge_loras = False
|
||||
if lora is not None:
|
||||
for l in lora:
|
||||
@@ -761,7 +971,7 @@ class WanVideoModelLoader:
|
||||
mm.unload_all_models()
|
||||
mm.cleanup_models()
|
||||
mm.soft_empty_cache()
|
||||
manual_offloading = True
|
||||
|
||||
if "sage" in attention_mode:
|
||||
try:
|
||||
from sageattention import sageattn
|
||||
@@ -777,8 +987,6 @@ class WanVideoModelLoader:
|
||||
if merge_loras is True:
|
||||
raise ValueError("GGUF models do not support LoRA merging, please disable merge_loras in the LoRA select node.")
|
||||
|
||||
|
||||
manual_offloading = True
|
||||
transformer_load_device = device if load_device == "main_device" else offload_device
|
||||
|
||||
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
|
||||
@@ -797,13 +1005,16 @@ class WanVideoModelLoader:
|
||||
|
||||
|
||||
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
|
||||
|
||||
|
||||
gguf_reader = None
|
||||
if not gguf:
|
||||
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
|
||||
else:
|
||||
from diffusers.models.model_loading_utils import load_gguf_checkpoint
|
||||
sd = load_gguf_checkpoint(model_path)
|
||||
|
||||
gguf_reader=[]
|
||||
from .gguf.gguf import load_gguf
|
||||
sd, reader = load_gguf(model_path)
|
||||
gguf_reader.append(reader)
|
||||
|
||||
if quantization == "disabled":
|
||||
for k, v in sd.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
@@ -831,14 +1042,18 @@ class WanVideoModelLoader:
|
||||
if "vace_blocks.0.after_proj.weight" in sd and not "patch_embedding.weight" in sd:
|
||||
raise ValueError("You are attempting to load a VACE module as a WanVideo model, instead you should use the vace_model input and matching T2V base model")
|
||||
|
||||
if vace_model is not None:
|
||||
# currently this can be VAE or MTV-Crafter weights
|
||||
if extra_model is not None:
|
||||
if gguf:
|
||||
if not vace_model["path"].endswith(".gguf"):
|
||||
raise ValueError("With GGUF main model the VACE module must also be a GGUF quantized, if the main model already has VACE included, you can disconnect the VACE module loader")
|
||||
vace_sd = load_gguf_checkpoint(vace_model["path"])
|
||||
if not extra_model["path"].endswith(".gguf"):
|
||||
raise ValueError("With GGUF main model the extra model must also be a GGUF quantized, if the main model already has extra included, you can disconnect the extra module loader")
|
||||
extra_sd, extra_reader = load_gguf(extra_model["path"])
|
||||
gguf_reader.append(extra_reader)
|
||||
del extra_reader
|
||||
else:
|
||||
vace_sd = load_torch_file(vace_model["path"], device=transformer_load_device, safe_load=True)
|
||||
sd.update(vace_sd)
|
||||
extra_sd = load_torch_file(extra_model["path"], device=transformer_load_device, safe_load=True)
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
|
||||
first_key = next(iter(sd))
|
||||
if first_key.startswith("model.diffusion_model."):
|
||||
@@ -971,6 +1186,7 @@ class WanVideoModelLoader:
|
||||
"add_ref_conv": True if "ref_conv.weight" in sd else False,
|
||||
"in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None,
|
||||
"add_control_adapter": True if "control_adapter.conv.weight" in sd else False,
|
||||
"use_motion_attn": True if "blocks.0.motion_attn.k.weight" in sd else False
|
||||
}
|
||||
|
||||
with init_empty_weights():
|
||||
@@ -1002,40 +1218,54 @@ class WanVideoModelLoader:
|
||||
log.info("FantasyPortrait model detected, patching model...")
|
||||
context_dim = fantasyportrait_model["sd"]["ip_adapter.blocks.0.cross_attn.ip_adapter_single_stream_k_proj.weight"].shape[1]
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
with init_empty_weights():
|
||||
for block in transformer.blocks:
|
||||
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
ip_adapter_sd = {}
|
||||
for k, v in fantasyportrait_model["sd"].items():
|
||||
if k.startswith("ip_adapter."):
|
||||
ip_adapter_sd[k.replace("ip_adapter.", "")] = v
|
||||
sd.update(ip_adapter_sd)
|
||||
del ip_adapter_sd
|
||||
|
||||
if multitalk_model is not None:
|
||||
if multitalk_model["is_gguf"] and not gguf:
|
||||
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
|
||||
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
|
||||
log.info(f"{multitalk_model_type} detected, patching model...")
|
||||
|
||||
multitalk_model_path = multitalk_model["model_path"]
|
||||
if multitalk_model_path.endswith(".gguf") and not gguf:
|
||||
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
|
||||
|
||||
# init audio module
|
||||
from .multitalk.multitalk import SingleStreamMultiAttention
|
||||
from .wanvideo.modules.model import WanLayerNorm
|
||||
|
||||
with init_empty_weights():
|
||||
for block in transformer.blocks:
|
||||
|
||||
for block in transformer.blocks:
|
||||
with init_empty_weights():
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
|
||||
block.audio_cross_attn = SingleStreamMultiAttention(
|
||||
dim=dim,
|
||||
encoder_hidden_states_dim=768,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
class_range=24,
|
||||
class_interval=4,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
|
||||
log.info(f"{multitalk_model_type} detected, patching model...")
|
||||
qkv_bias=True,
|
||||
class_range=24,
|
||||
class_interval=4,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
transformer.multitalk_model_type = multitalk_model_type
|
||||
sd.update(multitalk_model["sd"])
|
||||
|
||||
|
||||
extra_model_path = multitalk_model["model_path"]
|
||||
if gguf:
|
||||
extra_sd, extra_reader = load_gguf(extra_model_path)
|
||||
gguf_reader.append(extra_reader)
|
||||
del extra_reader
|
||||
else:
|
||||
extra_sd = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
|
||||
# Additional cond latents
|
||||
if "add_conv_in.weight" in sd:
|
||||
def zero_module(module):
|
||||
@@ -1055,195 +1285,69 @@ class WanVideoModelLoader:
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
device=device,
|
||||
)
|
||||
scale_weights = {}
|
||||
if not gguf:
|
||||
if "fp8" in quantization:
|
||||
for k, v in sd.items():
|
||||
if k.endswith(".scale_weight"):
|
||||
scale_weights[k] = v
|
||||
if not merge_loras:
|
||||
from .custom_linear import _replace_linear
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
|
||||
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
dtype = torch.float8_e4m3fn
|
||||
elif "fp8_e5m2" in quantization:
|
||||
dtype = torch.float8_e5m2
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
|
||||
if not lora_low_mem_load:
|
||||
log.info("Using accelerate to load and assign model weights to device...")
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
for name, param in tqdm(transformer.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
total=param_count,
|
||||
leave=True):
|
||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
||||
dtype_to_use = dtype if sd[name].dtype == dtype else dtype_to_use
|
||||
if "modulation" in name or "norm" in name or "bias" in name:
|
||||
dtype_to_use = base_dtype
|
||||
if "patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.dtype, param.device, param.shape)
|
||||
pbar.update_absolute(param_count)
|
||||
pbar.update_absolute(0)
|
||||
|
||||
comfy_model.diffusion_model = transformer
|
||||
comfy_model.load_device = transformer_load_device
|
||||
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
patcher.model.is_patched = False
|
||||
|
||||
scale_weights = {}
|
||||
if "fp8" in quantization:
|
||||
for k, v in sd.items():
|
||||
if k.endswith(".scale_weight"):
|
||||
scale_weights[k] = v.to(base_dtype)
|
||||
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
weight_dtype = torch.float8_e4m3fn
|
||||
elif "fp8_e5m2" in quantization:
|
||||
weight_dtype = torch.float8_e5m2
|
||||
else:
|
||||
weight_dtype = base_dtype
|
||||
|
||||
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
|
||||
|
||||
control_lora = False
|
||||
|
||||
if not merge_loras and control_lora:
|
||||
log.warning("Control-LoRA patching is only supported with merge_loras=True")
|
||||
|
||||
unianimate_sd = None
|
||||
control_lora = False
|
||||
if lora is not None:
|
||||
for l in lora:
|
||||
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
|
||||
lora_path = l["path"]
|
||||
lora_strength = l["strength"]
|
||||
if isinstance(lora_strength, list):
|
||||
if merge_loras:
|
||||
raise ValueError("LoRA strength should be a single value when merge_loras=True")
|
||||
transformer.lora_scheduling_enabled = True
|
||||
if lora_strength == 0:
|
||||
log.warning(f"LoRA {lora_path} has strength 0, skipping...")
|
||||
continue
|
||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
||||
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
|
||||
from .unianimate.nodes import update_transformer
|
||||
log.info("Unianimate LoRA detected, patching model...")
|
||||
transformer, unianimate_sd = update_transformer(transformer, lora_sd)
|
||||
|
||||
lora_sd = standardize_lora_key_format(lora_sd)
|
||||
|
||||
if l["blocks"]:
|
||||
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
|
||||
|
||||
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
|
||||
if not any('img' in k for k in sd.keys()):
|
||||
lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
|
||||
|
||||
#spacepxl's control LoRA patch
|
||||
# for key in lora_sd.keys():
|
||||
# print(key)
|
||||
patcher, control_lora, unianimate_sd = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras)
|
||||
if unianimate_sd is not None:
|
||||
log.info("Merging UniAnimate weights to the model...")
|
||||
sd.update(unianimate_sd)
|
||||
del unianimate_sd
|
||||
|
||||
if not gguf:
|
||||
if merge_loras and lora is not None:
|
||||
if not lora_low_mem_load:
|
||||
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
|
||||
|
||||
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
|
||||
log.info("Control-LoRA detected, patching model...")
|
||||
if not merge_loras:
|
||||
log.warning("Control-LoRA patching is only supported with merge_loras=True, setting it to True")
|
||||
merge_loras = True
|
||||
control_lora = True
|
||||
|
||||
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
|
||||
old_in_dim = transformer.in_dim # 16
|
||||
new_in_dim = lora_sd["diffusion_model.patch_embedding.lora_A.weight"].shape[1]
|
||||
assert new_in_dim == 32
|
||||
if control_lora:
|
||||
patch_control_lora(patcher.model.diffusion_model, device)
|
||||
patcher.model.is_patched = True
|
||||
|
||||
new_in = in_cls(
|
||||
new_in_dim,
|
||||
transformer.patch_embedding.out_channels,
|
||||
transformer.patch_embedding.kernel_size,
|
||||
transformer.patch_embedding.stride,
|
||||
transformer.patch_embedding.padding,
|
||||
).to(device=device, dtype=torch.float32)
|
||||
|
||||
new_in.weight.zero_()
|
||||
new_in.bias.zero_()
|
||||
|
||||
new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight)
|
||||
new_in.bias.copy_(transformer.patch_embedding.bias)
|
||||
|
||||
transformer.patch_embedding = new_in
|
||||
transformer.expanded_patch_embedding = new_in
|
||||
|
||||
if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
|
||||
log.info("Stand-In LoRA detected")
|
||||
for block in transformer.blocks:
|
||||
block.self_attn.q_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.k_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
block.self_attn.v_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
|
||||
for lora in [block.self_attn.q_loras, block.self_attn.k_loras, block.self_attn.v_loras]:
|
||||
for param in lora.parameters():
|
||||
param.requires_grad = False
|
||||
for name, param in transformer.named_parameters():
|
||||
if "lora" in name:
|
||||
param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype))
|
||||
else:
|
||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||
|
||||
del lora_sd
|
||||
|
||||
if not gguf and merge_loras:
|
||||
log.info("Patching LoRA to the model...")
|
||||
log.info("Merging LoRA to the model...")
|
||||
patcher = apply_lora(
|
||||
patcher, device, transformer_load_device,
|
||||
params_to_keep=params_to_keep, dtype=dtype, base_dtype=base_dtype, state_dict=sd,
|
||||
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights)
|
||||
scale_weights.clear()
|
||||
patcher.patches.clear()
|
||||
patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
|
||||
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights,)
|
||||
if not control_lora:
|
||||
scale_weights.clear()
|
||||
patcher.patches.clear()
|
||||
transformer.patched_linear = False
|
||||
else:
|
||||
from .custom_linear import _replace_linear
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
|
||||
transformer.patched_linear = True
|
||||
|
||||
if unianimate_sd is not None:
|
||||
sd.update(unianimate_sd)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "dwpose_embedding" in name or "randomref_embedding_pose" in name:
|
||||
dtype_to_use = base_dtype
|
||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
|
||||
if gguf:
|
||||
#from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter
|
||||
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
|
||||
log.info("Using GGUF to load and assign model weights to device...")
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
|
||||
out_features = sd["blocks.0.self_attn.k.weight"].shape[1]
|
||||
|
||||
patcher.model.diffusion_model = _replace_with_gguf_linear(patcher.model.diffusion_model, base_dtype, sd, patches=patcher.patches)
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
for name, param in tqdm(patcher.model.diffusion_model.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
total=param_count,
|
||||
leave=True):
|
||||
if "loras" in name:
|
||||
continue
|
||||
#print(name, param.dtype, param.device, param.shape)
|
||||
if isinstance(param, GGUFParameter):
|
||||
dtype_to_use = torch.uint8
|
||||
elif "patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
else:
|
||||
dtype_to_use = base_dtype
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.dtype, param.device, param.shape)
|
||||
#patcher.load(device, full_load=True)
|
||||
pbar.update_absolute(param_count)
|
||||
|
||||
patcher.model.is_patched = True
|
||||
|
||||
patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False)
|
||||
|
||||
if "fast" in quantization:
|
||||
if lora is not None and not merge_loras:
|
||||
raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs")
|
||||
from .fp8_optimization import convert_fp8_linear
|
||||
convert_fp8_linear(transformer, base_dtype, params_to_keep, scale_weight_keys=scale_weights)
|
||||
patch_linear = False
|
||||
|
||||
del sd
|
||||
if multitalk_model is not None:
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
|
||||
if vram_management_args is not None:
|
||||
if gguf:
|
||||
@@ -1269,18 +1373,18 @@ class WanVideoModelLoader:
|
||||
WanRMSNorm: AutoWrappedModule,
|
||||
},
|
||||
module_config = dict(
|
||||
offload_dtype=dtype,
|
||||
offload_dtype=weight_dtype,
|
||||
offload_device=offload_device,
|
||||
onload_dtype=dtype,
|
||||
onload_dtype=weight_dtype,
|
||||
onload_device=device,
|
||||
computation_dtype=base_dtype,
|
||||
computation_device=device,
|
||||
),
|
||||
max_num_param=params_to_keep,
|
||||
overflow_module_config = dict(
|
||||
offload_dtype=dtype,
|
||||
offload_dtype=weight_dtype,
|
||||
offload_device=offload_device,
|
||||
onload_dtype=dtype,
|
||||
onload_dtype=weight_dtype,
|
||||
onload_device=offload_device,
|
||||
computation_dtype=base_dtype,
|
||||
computation_device=device,
|
||||
@@ -1288,28 +1392,29 @@ class WanVideoModelLoader:
|
||||
compile_args = compile_args,
|
||||
)
|
||||
|
||||
if load_device == "offload_device" and patcher.model.diffusion_model.device != offload_device:
|
||||
if merge_loras and lora is not None:
|
||||
log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}")
|
||||
patcher.model.diffusion_model.to(offload_device)
|
||||
gc.collect()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
patcher.model["dtype"] = base_dtype
|
||||
patcher.model["base_dtype"] = base_dtype
|
||||
patcher.model["weight_dtype"] = weight_dtype
|
||||
patcher.model["base_path"] = model_path
|
||||
patcher.model["model_name"] = model
|
||||
patcher.model["manual_offloading"] = manual_offloading
|
||||
patcher.model["quantization"] = quantization
|
||||
patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False
|
||||
patcher.model["control_lora"] = control_lora
|
||||
patcher.model["compile_args"] = compile_args
|
||||
patcher.model["gguf"] = gguf
|
||||
patcher.model["gguf_reader"] = gguf_reader
|
||||
patcher.model["fp8_matmul"] = "fast" in quantization
|
||||
patcher.model["scale_weights"] = scale_weights
|
||||
patcher.model["sd"] = sd
|
||||
patcher.model["lora"] = lora
|
||||
|
||||
if 'transformer_options' not in patcher.model_options:
|
||||
patcher.model_options['transformer_options'] = {}
|
||||
patcher.model_options["transformer_options"]["block_swap_args"] = block_swap_args
|
||||
patcher.model_options["transformer_options"]["patch_linear"] = patch_linear
|
||||
patcher.model_options["transformer_options"]["merge_loras"] = merge_loras
|
||||
|
||||
for model in mm.current_loaded_models:
|
||||
@@ -1374,8 +1479,6 @@ class WanVideoVAELoader:
|
||||
|
||||
def loadmodel(self, model_name, precision, compile_args=None):
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
#with open(os.path.join(script_directory, 'configs', 'hy_vae_config.json')) as f:
|
||||
# vae_config = json.load(f)
|
||||
model_path = folder_paths.get_full_path("vae", model_name)
|
||||
vae_sd = load_torch_file(model_path, safe_load=True)
|
||||
|
||||
@@ -1389,8 +1492,9 @@ class WanVideoVAELoader:
|
||||
vae = WanVideoVAE38(dtype=dtype)
|
||||
|
||||
vae.load_state_dict(vae_sd)
|
||||
del vae_sd
|
||||
vae.eval()
|
||||
vae.to(device = offload_device, dtype = dtype)
|
||||
vae.to(device=offload_device, dtype=dtype)
|
||||
if compile_args is not None:
|
||||
vae.model.decoder = torch.compile(vae.model.decoder, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
@@ -1589,6 +1693,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoLoraBlockEdit": WanVideoLoraBlockEdit,
|
||||
"WanVideoTinyVAELoader": WanVideoTinyVAELoader,
|
||||
"WanVideoVACEModelSelect": WanVideoVACEModelSelect,
|
||||
"WanVideoExtraModelSelect": WanVideoExtraModelSelect,
|
||||
"WanVideoLoraSelectMulti": WanVideoLoraSelectMulti,
|
||||
"WanVideoBlockSwap": WanVideoBlockSwap,
|
||||
"WanVideoVRAMManagement": WanVideoVRAMManagement,
|
||||
@@ -1605,6 +1710,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoLoraBlockEdit": "WanVideo Lora Block Edit",
|
||||
"WanVideoTinyVAELoader": "WanVideo Tiny VAE Loader",
|
||||
"WanVideoVACEModelSelect": "WanVideo VACE Module Select",
|
||||
"WanVideoExtraModelSelect": "WanVideo Extra Model Select",
|
||||
"WanVideoLoraSelectMulti": "WanVideo Lora Select Multi",
|
||||
"WanVideoBlockSwap": "WanVideo Block Swap",
|
||||
"WanVideoVRAMManagement": "WanVideo VRAM Management",
|
||||
|
||||
@@ -173,7 +173,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
to_load.append((n, m, params))
|
||||
|
||||
to_load.sort(reverse=True)
|
||||
pbar = ProgressBar(len(to_load))
|
||||
#pbar = ProgressBar(len(to_load))
|
||||
for x in tqdm(to_load, desc="Loading model and applying LoRA weights:", leave=True):
|
||||
name = x[0]
|
||||
m = x[1]
|
||||
@@ -207,7 +207,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
except:
|
||||
continue
|
||||
m.comfy_patched_weights = True
|
||||
pbar.update(1)
|
||||
#pbar.update(1)
|
||||
|
||||
# After LoRA patching, scale weights that have scale_weight but are NOT LoRA patched
|
||||
if len(scale_weights) > 0 and not getattr(model, "scale_weights_applied", False):
|
||||
|
||||
@@ -189,6 +189,8 @@ def attention(
|
||||
version=fa_version,
|
||||
)
|
||||
elif attention_mode == 'sdpa':
|
||||
if not (q.dtype == k.dtype == v.dtype):
|
||||
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype)).transpose(1, 2).contiguous()
|
||||
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous()
|
||||
elif attention_mode == 'sageattn_3':
|
||||
return sageattn_blackwell(
|
||||
|
||||
+118
-31
@@ -28,6 +28,9 @@ from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCac
|
||||
from ...multitalk.multitalk import get_attn_map_with_target
|
||||
from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
|
||||
|
||||
from ...MTV.mtv import apply_rotary_emb
|
||||
|
||||
|
||||
__all__ = ['WanModel']
|
||||
|
||||
from comfy import model_management as mm
|
||||
@@ -668,6 +671,31 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
return self.o(x)
|
||||
|
||||
class MTVCrafterMotionAttention(WanSelfAttention):
|
||||
|
||||
def forward(self, x, mo, pe, grid_sizes, freqs):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
mo: Motion tokens
|
||||
pe: 4D RoPE
|
||||
"""
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
k = self.norm_k(self.k(mo)).view(b, n, -1, d)
|
||||
v = self.v(mo).view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
x = attention(
|
||||
q=rope_apply(q, grid_sizes, freqs),
|
||||
k=apply_rotary_emb(k, pe).transpose(1, 2),
|
||||
v=v
|
||||
)
|
||||
|
||||
return self.o(x.flatten(2))
|
||||
|
||||
|
||||
WAN_CROSSATTENTION_CLASSES = {
|
||||
't2v_cross_attn': WanT2VCrossAttention,
|
||||
@@ -689,6 +717,7 @@ class WanAttentionBlock(nn.Module):
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func="comfy",
|
||||
use_motion_attn=False
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
@@ -705,15 +734,19 @@ class WanAttentionBlock(nn.Module):
|
||||
self.dense_attention_mode = "sageattn"
|
||||
|
||||
self.kv_cache = None
|
||||
self.use_motion_attn = use_motion_attn
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(out_features, eps)
|
||||
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm,
|
||||
eps, self.attention_mode)
|
||||
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode)
|
||||
|
||||
# MTV Crafter motion attn
|
||||
if self.use_motion_attn:
|
||||
self.norm4 = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.motion_attn = MTVCrafterMotionAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode)
|
||||
|
||||
if cross_attn_type != "no_cross_attn":
|
||||
self.norm3 = WanLayerNorm(
|
||||
out_features, eps,
|
||||
elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.norm3 = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](in_features,
|
||||
out_features,
|
||||
num_heads,
|
||||
@@ -790,7 +823,11 @@ class WanAttentionBlock(nn.Module):
|
||||
freqs_ip=None,
|
||||
adapter_proj=None,
|
||||
ip_scale=1.0,
|
||||
reverse_time=False
|
||||
reverse_time=False,
|
||||
mtv_motion_tokens=None,
|
||||
mtv_motion_rotary_emb=None,
|
||||
mtv_strength=1.0,
|
||||
mtv_freqs=None
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
@@ -932,7 +969,8 @@ class WanAttentionBlock(nn.Module):
|
||||
x = self.cross_attn_ffn(x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
||||
audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond,
|
||||
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale)
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale,
|
||||
mtv_freqs=mtv_freqs, mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength)
|
||||
else:
|
||||
if self.rope_func == "comfy_chunked":
|
||||
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
@@ -951,19 +989,24 @@ class WanAttentionBlock(nn.Module):
|
||||
def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
||||
audio_proj, audio_scale, num_latent_frames, nag_params,
|
||||
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
|
||||
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale):
|
||||
|
||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
|
||||
audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength):
|
||||
|
||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
|
||||
audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale)
|
||||
#multitalk
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
|
||||
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
x = x + x_audio * audio_scale
|
||||
|
||||
# MTV-Crafter Motion Attention
|
||||
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
|
||||
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
|
||||
x = x + x_motion * mtv_strength
|
||||
|
||||
if self.rope_func == "comfy_chunked":
|
||||
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
else:
|
||||
@@ -1165,6 +1208,7 @@ class WanModel(torch.nn.Module):
|
||||
in_dim_ref_conv=16,
|
||||
add_control_adapter=False,
|
||||
in_dim_control_adapter=24,
|
||||
use_motion_attn=False
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
@@ -1226,6 +1270,7 @@ class WanModel(torch.nn.Module):
|
||||
self.offload_device = offload_device
|
||||
self.vace_layers = vace_layers
|
||||
self.device = main_device
|
||||
self.patched_linear = False
|
||||
|
||||
self.blocks_to_swap = -1
|
||||
self.offload_txt_emb = False
|
||||
@@ -1325,9 +1370,12 @@ class WanModel(torch.nn.Module):
|
||||
self.blocks = nn.ModuleList([
|
||||
WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads,
|
||||
qk_norm, cross_attn_norm, eps,
|
||||
attention_mode=self.attention_mode, rope_func=self.rope_func)
|
||||
for _ in range(num_layers)
|
||||
attention_mode=self.attention_mode, rope_func=self.rope_func, use_motion_attn=(i % 4 == 0 and use_motion_attn))
|
||||
for i in range(num_layers)
|
||||
])
|
||||
#MTV Crafter
|
||||
if use_motion_attn:
|
||||
self.pad_motion_tokens = torch.zeros(1, 1, 2048)
|
||||
|
||||
# head
|
||||
self.head = Head(dim, out_dim, patch_size, eps)
|
||||
@@ -1410,7 +1458,10 @@ class WanModel(torch.nn.Module):
|
||||
return block_mask
|
||||
|
||||
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False):
|
||||
log.info(f"Swapping {blocks_to_swap + 1} transformer blocks")
|
||||
# Clamp blocks_to_swap to valid range
|
||||
blocks_to_swap = max(0, min(blocks_to_swap, len(self.blocks)))
|
||||
|
||||
log.info(f"Swapping {blocks_to_swap} transformer blocks")
|
||||
self.blocks_to_swap = blocks_to_swap
|
||||
self.prefetch_blocks = prefetch_blocks
|
||||
self.block_swap_debug = block_swap_debug
|
||||
@@ -1420,11 +1471,14 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
total_offload_memory = 0
|
||||
total_main_memory = 0
|
||||
|
||||
# Calculate the index where swapping starts
|
||||
swap_start_idx = len(self.blocks) - blocks_to_swap
|
||||
|
||||
for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
if b > self.blocks_to_swap:
|
||||
if b < swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
else:
|
||||
@@ -1435,12 +1489,17 @@ class WanModel(torch.nn.Module):
|
||||
vace_blocks_to_swap = 1
|
||||
|
||||
if vace_blocks_to_swap > 0 and self.vace_layers is not None:
|
||||
# Clamp vace_blocks_to_swap to valid range
|
||||
vace_blocks_to_swap = max(0, min(vace_blocks_to_swap, len(self.vace_blocks)))
|
||||
self.vace_blocks_to_swap = vace_blocks_to_swap
|
||||
|
||||
# Calculate the index where VACE swapping starts
|
||||
vace_swap_start_idx = len(self.vace_blocks) - vace_blocks_to_swap
|
||||
|
||||
for b, block in tqdm(enumerate(self.vace_blocks), total=len(self.vace_blocks), desc="Initializing vace block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
if b > self.vace_blocks_to_swap:
|
||||
if b < vace_swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
else:
|
||||
@@ -1480,9 +1539,10 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
hints = []
|
||||
current_c = c
|
||||
vace_swap_start_idx = len(self.vace_blocks) - self.vace_blocks_to_swap if self.vace_blocks_to_swap > 0 else len(self.vace_blocks)
|
||||
|
||||
for b, block in enumerate(self.vace_blocks):
|
||||
if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0:
|
||||
if b >= vace_swap_start_idx and self.vace_blocks_to_swap > 0:
|
||||
block.to(self.main_device)
|
||||
|
||||
if b == 0:
|
||||
@@ -1495,13 +1555,13 @@ class WanModel(torch.nn.Module):
|
||||
# Store skip connection
|
||||
c_skip = block.after_proj(c_processed)
|
||||
hints.append(c_skip.to(
|
||||
self.offload_device if self.vace_blocks_to_swap != -1 else self.main_device,
|
||||
self.offload_device if self.vace_blocks_to_swap > 0 else self.main_device,
|
||||
non_blocking=self.use_non_blocking
|
||||
))
|
||||
|
||||
current_c = c_processed
|
||||
|
||||
if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0:
|
||||
if b >= vace_swap_start_idx and self.vace_blocks_to_swap > 0:
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
return hints
|
||||
@@ -1543,8 +1603,14 @@ class WanModel(torch.nn.Module):
|
||||
inner_t=None,
|
||||
standin_input=None,
|
||||
fantasy_portrait_input=None,
|
||||
phantom_ref=None,
|
||||
reverse_time=False,
|
||||
ntk_alphas = [1.0, 1.0, 1.0]
|
||||
ntk_alphas = [1.0, 1.0, 1.0],
|
||||
mtv_motion_tokens=None,
|
||||
mtv_motion_rotary_emb=None,
|
||||
mtv_freqs=None,
|
||||
mtv_strength=1.0,
|
||||
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -1570,6 +1636,11 @@ class WanModel(torch.nn.Module):
|
||||
# Stand-In only used on first positive pass, then cached in kv_cache
|
||||
if is_uncond or current_step > 0:
|
||||
standin_input = None
|
||||
|
||||
# MTV Crafter motion projection
|
||||
if mtv_motion_tokens is not None:
|
||||
bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1]
|
||||
mtv_motion_tokens = torch.cat([mtv_motion_tokens, self.pad_motion_tokens.to(mtv_motion_tokens).expand(bs, motion_seq_len, -1)], dim=-1)
|
||||
|
||||
# Fantasy Portrait
|
||||
adapter_proj = ip_scale = None
|
||||
@@ -1657,6 +1728,15 @@ class WanModel(torch.nn.Module):
|
||||
F += 1
|
||||
x = [torch.concat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)]
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_frames = phantom_ref.size(1)
|
||||
phantom_ref = self.original_patch_embedding(phantom_ref.unsqueeze(0).to(torch.float32)).flatten(2).transpose(1, 2).to(x[0].dtype)
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
phantom_ref_seq_len = phantom_ref.size(1)
|
||||
seq_len += phantom_ref_seq_len
|
||||
F += phantom_ref_frames
|
||||
x = [torch.concat([u, phantom_ref.unsqueeze(0)], dim=1) for phantom_ref, u in zip(phantom_ref, x)]
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat([
|
||||
@@ -2010,7 +2090,11 @@ class WanModel(torch.nn.Module):
|
||||
e_ip=e0_ip if x_ip is not None else None,
|
||||
adapter_proj=adapter_proj,
|
||||
ip_scale=ip_scale,
|
||||
reverse_time=reverse_time
|
||||
reverse_time=reverse_time,
|
||||
mtv_motion_tokens=mtv_motion_tokens,
|
||||
mtv_motion_rotary_emb=mtv_motion_rotary_emb,
|
||||
mtv_strength=mtv_strength,
|
||||
mtv_freqs=mtv_freqs
|
||||
)
|
||||
|
||||
if vace_data is not None:
|
||||
@@ -2050,20 +2134,21 @@ class WanModel(torch.nn.Module):
|
||||
# Asynchronous block offloading with CUDA streams and events
|
||||
cuda_stream = mm.get_offload_stream(device)
|
||||
events = [torch.cuda.Event() for _ in self.blocks]
|
||||
swap_start_idx = len(self.blocks) - self.blocks_to_swap if self.blocks_to_swap > 0 else len(self.blocks)
|
||||
|
||||
for b, block in enumerate(self.blocks):
|
||||
# Prefetch blocks if enabled
|
||||
if self.prefetch_blocks > 0:
|
||||
for prefetch_offset in range(1, self.prefetch_blocks + 1):
|
||||
prefetch_idx = b + prefetch_offset
|
||||
if prefetch_idx < len(self.blocks) and self.blocks_to_swap >= 0 and prefetch_idx <= self.blocks_to_swap:
|
||||
if prefetch_idx < len(self.blocks) and self.blocks_to_swap > 0 and prefetch_idx >= swap_start_idx:
|
||||
with torch.cuda.stream(cuda_stream):
|
||||
self.blocks[prefetch_idx].to(self.main_device, non_blocking=self.use_non_blocking)
|
||||
events[prefetch_idx].record(cuda_stream)
|
||||
if self.block_swap_debug:
|
||||
transfer_start = time.perf_counter()
|
||||
# Wait for block to be ready
|
||||
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
|
||||
if b >= swap_start_idx and self.blocks_to_swap > 0:
|
||||
if self.prefetch_blocks > 0:
|
||||
if not events[b].query():
|
||||
events[b].synchronize()
|
||||
@@ -2082,7 +2167,7 @@ class WanModel(torch.nn.Module):
|
||||
compute_end = time.perf_counter()
|
||||
compute_time = compute_end - compute_start
|
||||
to_cpu_transfer_start = time.perf_counter()
|
||||
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
|
||||
if b >= swap_start_idx and self.blocks_to_swap > 0:
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
if self.block_swap_debug:
|
||||
to_cpu_transfer_end = time.perf_counter()
|
||||
@@ -2096,9 +2181,6 @@ class WanModel(torch.nn.Module):
|
||||
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
|
||||
x[:, :x_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
|
||||
|
||||
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
if self.enable_teacache and (self.teacache_start_step <= current_step <= self.teacache_end_step) and pred_id is not None:
|
||||
self.teacache_state.update(
|
||||
pred_id,
|
||||
@@ -2131,9 +2213,14 @@ class WanModel(torch.nn.Module):
|
||||
)
|
||||
|
||||
if self.ref_conv is not None and fun_ref is not None:
|
||||
full_ref_length = fun_ref.size(1)
|
||||
x = x[:, full_ref_length:]
|
||||
fun_ref_length = fun_ref.size(1)
|
||||
x = x[:, fun_ref_length:]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_length = phantom_ref.size(1)
|
||||
x = x[:, :-phantom_ref_length]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if attn_cond is not None:
|
||||
x = x[:, :x_len]
|
||||
|
||||
@@ -23,7 +23,7 @@ scheduler_list = [
|
||||
"multitalk"
|
||||
]
|
||||
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None, seed_g=None):
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None):
|
||||
timesteps = None
|
||||
if 'unipc' in scheduler:
|
||||
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
|
||||
@@ -130,6 +130,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
|
||||
# Slice timesteps and sigmas once, based on indices
|
||||
timesteps = timesteps[start_idx:end_idx+1]
|
||||
sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone()
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
|
||||
|
||||
|
||||
@@ -138,11 +139,4 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
|
||||
if seed_g is not None:
|
||||
scheduler_step_args = {"generator": seed_g}
|
||||
step_sig = inspect.signature(sample_scheduler.step)
|
||||
for arg in list(scheduler_step_args.keys()):
|
||||
if arg not in step_sig.parameters:
|
||||
scheduler_step_args.pop(arg)
|
||||
|
||||
return sample_scheduler, timesteps, scheduler_step_args
|
||||
return sample_scheduler, timesteps
|
||||
Reference in New Issue
Block a user