Merge branch 'dev'

This commit is contained in:
kijai
2025-08-26 16:39:26 +03:00
23 changed files with 4223 additions and 697 deletions
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.
+142
View File
@@ -0,0 +1,142 @@
import cv2
import math
import torch
import numpy as np
from PIL import Image
from torchvision import transforms
def intrinsic_matrix_from_field_of_view(imshape, fov_degrees:float =55 ): # nlf default fov_degrees 55
imshape = np.array(imshape)
fov_radians = fov_degrees * np.array(np.pi / 180)
larger_side = np.max(imshape)
focal_length = larger_side / (np.tan(fov_radians / 2) * 2)
# intrinsic_matrix 3*3
return np.array([
[focal_length, 0, imshape[1] / 2],
[0, focal_length, imshape[0] / 2],
[0, 0, 1],
])
def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3
camera_matrix = intrinsic_matrix_from_field_of_view((height,width))
camera_matrix = np.expand_dims(camera_matrix, axis=0)
camera_matrix = np.expand_dims(camera_matrix, axis=0) # 1*1*3*3
point_3d = np.expand_dims(point_3d,axis=-1) # n*1024*3*1
point_2d = (camera_matrix@point_3d).squeeze(-1)
point_2d[:,:,:2] = point_2d[:,:,:2]/point_2d[:,:,2:3]
return point_2d[:,:,:] # n*1024*2
def get_pose_images(smpl_data, offset):
pose_images = []
for data in smpl_data:
if isinstance(data, np.ndarray):
joints3d = data
else:
joints3d = data.numpy()
canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8)
joints3d = p3d_to_p2d(joints3d, offset[0], offset[1])
canvas = draw_3d_points(canvas, joints3d[0], stickwidth=int(offset[1]/350))
pose_images.append(Image.fromarray(canvas))
return pose_images
def get_control_conditions(poses, h, w):
video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
control_images = []
for idx, pose in enumerate(poses):
canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
try:
joints3d = p3d_to_p2d(pose, h, w)
canvas = draw_3d_points(
canvas,
joints3d[0],
stickwidth=int(h / 350),
)
resized_canvas = cv2.resize(canvas, (w, h))
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
control_images.append(resized_canvas)
except Exception as e:
print("wrong:", e)
control_images.append(Image.fromarray(canvas))
control_pixel_values = np.array(control_images)
control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
print("control_pixel_values.shape", control_pixel_values.shape)
#control_pixel_values = video_transforms(control_pixel_values)
return control_pixel_values
def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True):
colors = [
[255, 0, 0], # 0
[0, 255, 0], # 1
[0, 0, 255], # 2
[255, 0, 255], # 3
[255, 255, 0], # 4
[85, 255, 0], # 5
[0, 75, 255], # 6
[0, 255, 85], # 7
[0, 255, 170], # 8
[170, 0, 255], # 9
[85, 0, 255], # 10
[0, 85, 255], # 11
[0, 255, 255], # 12
[85, 0, 255], # 13
[170, 0, 255], # 14
[255, 0, 255], # 15
[255, 0, 170], # 16
[255, 0, 85], # 17
]
connetions = [
[15,12],[12, 16],[16, 18],[18, 20],[20, 22],
[12,17],[17,19],[19,21],
[21,23],[12,9],[9,6],
[6,3],[3,0],[0,1],
[1,4],[4,7],[7,10],[0,2],[2,5],[5,8],[8,11]
]
connection_colors = [
[255, 0, 0], # 0
[0, 255, 0], # 1
[0, 0, 255], # 2
[255, 255, 0], # 3
[255, 0, 255], # 4
[0, 255, 0], # 5
[0, 85, 255], # 6
[255, 175, 0], # 7
[0, 0, 255], # 8
[255, 85, 0], # 9
[0, 255, 85], # 10
[255, 0, 255], # 11
[255, 0, 0], # 12
[0, 175, 255], # 13
[255, 255, 0], # 14
[0, 0, 255], # 15
[0, 255, 0], # 16
]
# draw point
for i in range(len(points)):
x,y = points[i][0:2]
x,y = int(x),int(y)
if i==13 or i == 14:
continue
cv2.circle(canvas, (x, y), r, colors[i%17], thickness=-1)
# draw line
if draw_line:
for i in range(len(connetions)):
point1_idx,point2_idx = connetions[i][0:2]
point1 = points[point1_idx]
point2 = points[point2_idx]
Y = [point2[0],point1[0]]
X = [point2[1],point1[1]]
mX = int(np.mean(X))
mY = int(np.mean(Y))
length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
polygon = cv2.ellipse2Poly((mY, mX), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
return canvas
+1
View File
@@ -0,0 +1 @@
from .vqvae import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
+329
View File
@@ -0,0 +1,329 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
class Encoder(nn.Module):
def __init__(
self,
in_channels=3,
mid_channels=[128, 512],
out_channels=3072,
downsample_time=[1, 1],
downsample_joint=[1, 1],
num_attention_heads=8,
attention_head_dim=64,
dim=3072,
):
super(Encoder, self).__init__()
self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
self.downsample1 = Downsample(mid_channels[0], mid_channels[0], downsample_time[0], downsample_joint[0])
self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
self.downsample2 = Downsample(mid_channels[1], mid_channels[1], downsample_time[1], downsample_joint[1])
self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = self.conv_in(x)
for resnet in self.resnet1:
x = resnet(x)
x = self.downsample1(x)
x = self.resnet2(x)
for resnet in self.resnet3:
x = resnet(x)
x = self.downsample2(x)
x = self.conv_out(x)
return x
class VectorQuantizer(nn.Module):
def __init__(self, nb_code, code_dim):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.mu = 0.99
self.reset_codebook()
self.reset_count = 0
self.usage = torch.zeros((self.nb_code, 1))
def reset_codebook(self):
self.init = False
self.code_sum = None
self.code_count = None
self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).cuda())
def _tile(self, x):
nb_code_x, code_dim = x.shape
if nb_code_x < self.nb_code:
n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x
std = 0.01 / np.sqrt(code_dim)
out = x.repeat(n_repeats, 1)
out = out + torch.randn_like(out) * std
else:
out = x
return out
def preprocess(self, x):
# [bs, c, f, j] -> [bs * f * j, c]
x = x.permute(0, 2, 3, 1).contiguous()
x = x.view(-1, x.shape[-1])
return x
def quantize(self, x):
# [bs * f * j, dim=3072]
# Calculate latent code x_l
k_w = self.codebook.t()
distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0, keepdim=True)
_, code_idx = torch.min(distance, dim=-1)
return code_idx
def dequantize(self, code_idx):
x = F.embedding(code_idx, self.codebook) # indexing: [bs * f * j, 32]
return x
def forward(self, x, return_vq=False):
bs, c, f, j = x.shape # SMPL data frames: [bs, 3072, f, j]
# Preprocess
x = self.preprocess(x)
# return x.view(bs, f*j, c).contiguous(), None
assert x.shape[-1] == self.code_dim
# quantize and dequantize through bottleneck
code_idx = self.quantize(x)
x_d = self.dequantize(code_idx)
# Loss
commit_loss = F.mse_loss(x, x_d.detach())
# Passthrough
x_d = x + (x_d - x).detach()
if return_vq:
return x_d.view(bs, f*j, c).contiguous(), commit_loss
# return (x_d, x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()), commit_loss, perplexity
# Postprocess
x_d = x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()
return x_d, commit_loss
class Decoder(nn.Module):
def __init__(
self,
in_channels=3072,
mid_channels=[512, 128],
out_channels=3,
upsample_rate=None,
frame_upsample_rate=[1.0, 1.0],
joint_upsample_rate=[1.0, 1.0],
dim=128,
attention_head_dim=64,
num_attention_heads=8,
):
super(Decoder, self).__init__()
self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
self.upsample1 = Upsample(mid_channels[0], mid_channels[0], frame_upsample_rate=frame_upsample_rate[0], joint_upsample_rate=joint_upsample_rate[0])
self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
self.upsample2 = Upsample(mid_channels[1], mid_channels[1], frame_upsample_rate=frame_upsample_rate[1], joint_upsample_rate=joint_upsample_rate[1])
self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = self.conv_in(x)
for resnet in self.resnet1:
x = resnet(x)
x = self.upsample1(x)
x = self.resnet2(x)
for resnet in self.resnet3:
x = resnet(x)
x = self.upsample2(x)
x = self.conv_out(x)
return x
class Upsample(nn.Module):
def __init__(
self,
in_channels,
out_channels,
upsample_rate=None,
frame_upsample_rate=None,
joint_upsample_rate=None,
):
super(Upsample, self).__init__()
self.upsampler = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.upsample_rate = upsample_rate
self.frame_upsample_rate = frame_upsample_rate
self.joint_upsample_rate = joint_upsample_rate
self.upsample_rate = upsample_rate
def forward(self, inputs):
if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1:
# split first frame
x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:]
if self.upsample_rate is not None:
# import pdb; pdb.set_trace()
x_first = F.interpolate(x_first, scale_factor=self.upsample_rate)
x_rest = F.interpolate(x_rest, scale_factor=self.upsample_rate)
else:
# import pdb; pdb.set_trace()
# x_first = F.interpolate(x_first, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
x_rest = F.interpolate(x_rest, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
x_first = x_first[:, :, None, :]
inputs = torch.cat([x_first, x_rest], dim=2)
elif inputs.shape[2] > 1:
if self.upsample_rate is not None:
inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
else:
inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
else:
inputs = inputs.squeeze(2)
if self.upsample_rate is not None:
inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
else:
inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="linear", align_corners=True)
inputs = inputs[:, :, None, :, :]
b, c, t, j = inputs.shape
inputs = inputs.permute(0, 2, 1, 3).reshape(b * t, c, j)
inputs = self.upsampler(inputs)
inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3)
return inputs
class Downsample(nn.Module):
def __init__(
self,
in_channels,
out_channels,
frame_downsample_rate,
joint_downsample_rate
):
super(Downsample, self).__init__()
self.frame_downsample_rate = frame_downsample_rate
self.joint_downsample_rate = joint_downsample_rate
self.joint_downsample = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=self.joint_downsample_rate, padding=1)
def forward(self, x):
# (batch_size, channels, frames, joints) -> (batch_size * joints, channels, frames)
if self.frame_downsample_rate > 1:
batch_size, channels, frames, joints = x.shape
x = x.permute(0, 3, 1, 2).reshape(batch_size * joints, channels, frames)
if x.shape[-1] % 2 == 1:
x_first, x_rest = x[..., 0], x[..., 1:]
if x_rest.shape[-1] > 0:
# (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2)
x_rest = F.avg_pool1d(x_rest, kernel_size=self.frame_downsample_rate, stride=self.frame_downsample_rate)
x = torch.cat([x_first[..., None], x_rest], dim=-1)
# (batch_size * joints, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, joints)
x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
else:
# (batch_size * joints, channels, frames) -> (batch_size * joints, channels, frames // 2)
x = F.avg_pool1d(x, kernel_size=2, stride=2)
# (batch_size * joints, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width)
x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
# Pad the tensor
# pad = (0, 1)
# x = F.pad(x, pad, mode="constant", value=0)
batch_size, channels, frames, joints = x.shape
# (batch_size, channels, frames, joints) -> (batch_size * frames, channels, joints)
x = x.permute(0, 2, 1, 3).reshape(batch_size * frames, channels, joints)
x = self.joint_downsample(x)
# (batch_size * frames, channels, joints) -> (batch_size, channels, frames, joints)
x = x.reshape(batch_size, frames, x.shape[1], x.shape[2]).permute(0, 2, 1, 3)
return x
class ResBlock(nn.Module):
def __init__(self,
in_channels,
out_channels,
group_num=32,
max_channels=512):
super(ResBlock, self).__init__()
skip = max(1, max_channels // out_channels - 1)
self.block = nn.Sequential(
nn.GroupNorm(group_num, in_channels, eps=1e-06, affine=True),
nn.SiLU(),
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=skip, dilation=skip),
nn.GroupNorm(group_num, out_channels, eps=1e-06, affine=True),
nn.SiLU(),
nn.Conv2d(out_channels, out_channels, kernel_size=1, stride=1, padding=0),
)
self.conv_short = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) if in_channels != out_channels else nn.Identity()
def forward(self, x):
hidden_states = self.block(x)
if hidden_states.shape != x.shape:
x = self.conv_short(x)
x = x + hidden_states
return x
class SMPL_VQVAE(nn.Module):
def __init__(self, encoder, decoder, vq):
super(SMPL_VQVAE, self).__init__()
self.encoder = encoder
self.decoder = decoder
self.vq = vq
def to(self, device):
self.encoder = self.encoder.to(device)
self.decoder = self.decoder.to(device)
self.vq = self.vq.to(device)
self.device = device
return self
def encdec_slice_frames(self, x, frame_batch_size, encdec, return_vq):
num_frames = x.shape[2]
remaining_frames = num_frames % frame_batch_size
x_output = []
for i in range(num_frames // frame_batch_size):
remaining_frames = num_frames % frame_batch_size
start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames)
end_frame = frame_batch_size * (i + 1) + remaining_frames
x_intermediate = x[:, :, start_frame:end_frame]
x_intermediate = encdec(x_intermediate)
x_output.append(x_intermediate)
if encdec == self.encoder and self.vq is not None:
x_output, loss = self.vq(torch.cat(x_output, dim=2), return_vq=return_vq)
return x_output, loss
else:
return torch.cat(x_output, dim=2), None, None
def forward(self, x, return_vq=False):
x = x.permute(0, 3, 1, 2)
x, loss = self.encdec_slice_frames(x, frame_batch_size=8, encdec=self.encoder, return_vq=return_vq)
if return_vq:
return x, loss
x, _, _ = self.encdec_slice_frames(x, frame_batch_size=2, encdec=self.decoder, return_vq=return_vq)
x = x.permute(0, 2, 3, 1)
return x, loss
+193
View File
@@ -0,0 +1,193 @@
import torch
import numpy as np
from typing import Union, Tuple
def get_1d_rotary_pos_embed(
dim: int,
pos: Union[np.ndarray, int],
theta: float = 10000.0,
use_real=False,
linear_factor=1.0,
ntk_factor=1.0,
repeat_interleave_real=True,
freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux)
):
"""
Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
data type.
Args:
dim (`int`): Dimension of the frequency tensor.
pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
theta (`float`, *optional*, defaults to 10000.0):
Scaling factor for frequency computation. Defaults to 10000.0.
use_real (`bool`, *optional*):
If True, return real part and imaginary part separately. Otherwise, return complex numbers.
linear_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the context extrapolation. Defaults to 1.0.
ntk_factor (`float`, *optional*, defaults to 1.0):
Scaling factor for the NTK-Aware RoPE. Defaults to 1.0.
repeat_interleave_real (`bool`, *optional*, defaults to `True`):
If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`.
Otherwise, they are concateanted with themselves.
freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`):
the dtype of the frequency tensor.
Returns:
`torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
"""
assert dim % 2 == 0
if isinstance(pos, int):
pos = torch.arange(pos)
if isinstance(pos, np.ndarray):
pos = torch.from_numpy(pos) # type: ignore # [S]
theta = theta * ntk_factor
freqs = (
1.0
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
/ linear_factor
) # [D/2]
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
if use_real and repeat_interleave_real:
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
return freqs_cos, freqs_sin
elif use_real:
freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D]
freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D]
return freqs_cos, freqs_sin
else:
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
return freqs_cis
def get_3d_rotary_pos_embed(
embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
"""
RoPE for video tokens with 3D structure.
Args:
embed_dim: (`int`):
The embedding dimension size, corresponding to hidden_size_head.
crops_coords (`Tuple[int]`):
The top-left and bottom-right coordinates of the crop.
grid_size (`Tuple[int]`):
The grid size of the spatial positional embedding (height, width).
temporal_size (`int`):
The size of the temporal dimension.
theta (`float`):
Scaling factor for frequency computation.
Returns:
`torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
"""
if use_real is not True:
raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
start, stop = crops_coords
grid_size_h, grid_size_w = grid_size
grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
# Compute dimensions for each axis
dim_t = embed_dim // 4
dim_h = embed_dim // 8 * 3
dim_w = embed_dim // 8 * 3
# Temporal frequencies
freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True)
# Spatial frequencies for height and width
freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True)
freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True)
# BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor
def combine_time_height_width(freqs_t, freqs_h, freqs_w):
freqs_t = freqs_t[:, None, None, :].expand(
-1, grid_size_h, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_w, dim_t
freqs_h = freqs_h[None, :, None, :].expand(
temporal_size, -1, grid_size_w, -1
) # temporal_size, grid_size_h, grid_size_2, dim_h
freqs_w = freqs_w[None, None, :, :].expand(
temporal_size, grid_size_h, -1, -1
) # temporal_size, grid_size_h, grid_size_2, dim_w
freqs = torch.cat(
[freqs_t, freqs_h, freqs_w], dim=-1
) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
freqs = freqs.view(
temporal_size * grid_size_h * grid_size_w, -1
) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
return freqs
t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
cos = combine_time_height_width(t_cos, h_cos, w_cos)
sin = combine_time_height_width(t_sin, h_sin, w_sin)
return cos, sin
def get_3d_motion_spatial_embed(
embed_dim: int, num_joints: int, joints_mean: np.ndarray, joints_std: np.ndarray, theta: float = 10000.0
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
assert embed_dim % 2 == 0 and embed_dim % 3 == 0
def create_rope_pe(dim, pos, freqs_dtype=torch.float32):
if isinstance(pos, np.ndarray):
pos = torch.from_numpy(pos)
freqs = (
1.0
/ (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
) # [D/2]
freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
return freqs_cos, freqs_sin
pos_x = joints_mean[:, 0]
pos_y = joints_mean[:, 1]
pos_z = joints_mean[:, 2]
normalized_pos_x = (pos_x - pos_x.mean())
normalized_pos_y = (pos_y - pos_y.mean())
normalized_pos_z = (pos_z - pos_z.mean())
freqs_cos_x, freqs_sin_x = create_rope_pe(embed_dim // 3, normalized_pos_x)
freqs_cos_y, freqs_sin_y = create_rope_pe(embed_dim // 3, normalized_pos_y)
freqs_cos_z, freqs_sin_z = create_rope_pe(embed_dim // 3, normalized_pos_z)
freqs_cos = torch.cat([freqs_cos_x, freqs_cos_y, freqs_cos_z], dim=-1)
freqs_sin = torch.cat([freqs_sin_x, freqs_sin_y, freqs_sin_z], dim=-1)
return freqs_cos, freqs_sin
def prepare_motion_embeddings(num_frames, num_joints, joints_mean, joints_std, theta=10000, device='cuda'):
time_embed = get_1d_rotary_pos_embed(44, num_frames, theta, use_real=True)
time_embed_cos = time_embed[0][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1)
time_embed_sin = time_embed[1][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1)
spatial_motion_embed = get_3d_motion_spatial_embed(84, num_joints, joints_mean, joints_std, theta)
spatial_embed_cos = spatial_motion_embed[0][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1)
spatial_embed_sin = spatial_motion_embed[1][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1)
motion_embed_cos = torch.cat([time_embed_cos, spatial_embed_cos], dim=-1).to(device=device)
motion_embed_sin = torch.cat([time_embed_sin, spatial_embed_sin], dim=-1).to(device=device)
return motion_embed_cos, motion_embed_sin
def apply_rotary_emb(x, freqs_cis):
cos, sin = freqs_cis # [S, D]
cos = cos[None, None]
sin = sin[None, None]
cos, sin = cos.to(x.device), sin.to(x.device)
x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
return out
View File
+242
View File
@@ -0,0 +1,242 @@
import os
import torch
import gc
from ..utils import log, dict_to_device
import numpy as np
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import comfy.model_management as mm
from comfy.utils import load_torch_file
import folder_paths
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript")
from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
from .mtv import prepare_motion_embeddings
class DownloadAndLoadNLFModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"url": (
[
"https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"
],
)
},
}
RETURN_TYPES = ("NLFMODEL",)
RETURN_NAMES = ("nlf_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, url):
if not os.path.exists(local_model_path):
log.info(f"Downloading NLF model to: {local_model_path}")
import requests
os.makedirs(os.path.dirname(local_model_path), exist_ok=True)
response = requests.get(url)
if response.status_code == 200:
with open(local_model_path, "wb") as f:
f.write(response.content)
else:
print("Failed to download file:", response.status_code)
model = torch.jit.load(local_model_path).eval()
return (model,)
class LoadNLFModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"path": ("STRING", {"default": local_model_path}),
},
}
RETURN_TYPES = ("NLFMODEL",)
RETURN_NAMES = ("nlf_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, path):
model = torch.jit.load(path).eval()
return model,
class LoadVQVAE:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
},
}
RETURN_TYPES = ("VQVAE",)
RETURN_NAMES = ("vqvae", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model_name):
model_path = folder_paths.get_full_path("vae", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
# Get motion tokenizer
motion_encoder = Encoder(
in_channels=3,
mid_channels=[128, 512],
out_channels=3072,
downsample_time=[2, 2],
downsample_joint=[1, 1]
)
motion_quant = VectorQuantizer(nb_code=8192, code_dim=3072)
motion_decoder = Decoder(
in_channels=3072,
mid_channels=[512, 128],
out_channels=3,
upsample_rate=2.0,
frame_upsample_rate=[2.0, 2.0],
joint_upsample_rate=[1.0, 1.0]
)
vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device)
vqvae.load_state_dict(vae_sd, strict=True)
return vqvae,
class MTVCrafterEncodePoses:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"vqvae": ("VQVAE", {"tooltip": "VQVAE model"}),
"poses": ("NLFPRED", {"tooltip": "Input poses for the model"}),
},
}
RETURN_TYPES = ("MTVCRAFTERMOTION", "NLFPRED")
RETURN_NAMES = ("mtvcrafter_motion", "pose_results")
FUNCTION = "encode"
CATEGORY = "WanVideoWrapper"
def encode(self, vqvae, poses):
# import pickle
# with open(os.path.join(script_directory, "data", "sampled_data.pkl"), 'rb') as f:
# data_list = pickle.load(f)
# if not isinstance(data_list, list):
# data_list = [data_list]
# print(data_list)
# smpl_poses = data_list[1]['pose']
global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3)
global_std = np.load(os.path.join(script_directory, "data", "std.npy"))
smpl_poses = []
for pose in poses['joints3d_nonparam'][0]:
smpl_poses.append(pose[0].cpu().numpy())
smpl_poses = np.array(smpl_poses)
norm_poses = torch.tensor((smpl_poses - global_mean) / global_std).unsqueeze(0)
print(f"norm_poses shape: {norm_poses.shape}, dtype: {norm_poses.dtype}")
vqvae.to(device)
motion_tokens, vq_loss = vqvae(norm_poses.to(device), return_vq=True)
recon_motion = vqvae(norm_poses.to(device))[0][0].to(dtype=torch.float32).cpu().detach() * global_std + global_mean
vqvae.to(offload_device)
poses_dict = {
'mtv_motion_tokens': motion_tokens,
'global_mean': global_mean,
'global_std': global_std
}
return poses_dict, recon_motion
class NLFPredict:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("NLFMODEL",),
"images": ("IMAGE", {"tooltip": "Input images for the model"}),
},
}
RETURN_TYPES = ("NLFPRED", )
RETURN_NAMES = ("pose_results",)
FUNCTION = "predict"
CATEGORY = "WanVideoWrapper"
def predict(self, model, images):
model.to(device)
pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
model.to(offload_device)
pred = dict_to_device(pred, offload_device)
pose_results = {
'joints3d_nonparam': [],
}
# Collect pose data
for key in pose_results.keys():
if key in pred:
pose_results[key].append(pred[key])
else:
pose_results[key].append(None)
return (pose_results,)
class DrawNLFPoses:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"poses": ("NLFPRED", {"tooltip": "Input poses for the model"}),
"width": ("INT", {"default": 512}),
"height": ("INT", {"default": 512}),
},
}
RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("image",)
FUNCTION = "predict"
CATEGORY = "WanVideoWrapper"
def predict(self, poses, width, height):
from .draw_pose import get_control_conditions
print(type(poses))
if isinstance(poses, dict):
pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
else:
pose_input = poses
control_conditions = get_control_conditions(pose_input, height, width)
return (control_conditions,)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadNLFModel": DownloadAndLoadNLFModel,
"NLFPredict": NLFPredict,
"DrawNLFPoses": DrawNLFPoses,
"LoadVQVAE": LoadVQVAE,
"MTVCrafterEncodePoses": MTVCrafterEncodePoses
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadNLFModel": "(Download)Load NLF Model",
"NLFPredict": "NLF Predict",
"DrawNLFPoses": "Draw NLF Poses",
"LoadVQVAE": "Load VQVAE",
"MTVCrafterEncodePoses": "MTV Crafter Encode Poses"
}
+10 -2
View File
@@ -35,6 +35,13 @@ except Exception as e:
UNIANIMATE_NODE_CLASS_MAPPINGS = {}
UNIANIMATE_NODE_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
View File
@@ -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)
@@ -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
-73
View File
@@ -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
View File
@@ -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=""):
+3 -3
View File
@@ -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
View File
@@ -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",
}
+316 -147
View File
@@ -16,9 +16,10 @@ from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fouri
add_noise_to_reference_video, optimized_scale, setup_radial_attention,
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
View File
@@ -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",
+2 -2
View File
@@ -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):
+2
View File
@@ -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
View File
@@ -28,6 +28,9 @@ from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCac
from ...multitalk.multitalk import get_attn_map_with_target
from ...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]
+3 -9
View File
@@ -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