diff --git a/MTV/data/mean.npy b/MTV/data/mean.npy
new file mode 100644
index 0000000..001d7e7
Binary files /dev/null and b/MTV/data/mean.npy differ
diff --git a/MTV/data/std.npy b/MTV/data/std.npy
new file mode 100644
index 0000000..5d1db82
Binary files /dev/null and b/MTV/data/std.npy differ
diff --git a/MTV/draw_pose.py b/MTV/draw_pose.py
new file mode 100644
index 0000000..eb5c606
--- /dev/null
+++ b/MTV/draw_pose.py
@@ -0,0 +1,142 @@
+import cv2
+import math
+import torch
+import numpy as np
+from PIL import Image
+from torchvision import transforms
+
+
+def intrinsic_matrix_from_field_of_view(imshape, fov_degrees:float =55 ): # nlf default fov_degrees 55
+ imshape = np.array(imshape)
+ fov_radians = fov_degrees * np.array(np.pi / 180)
+ larger_side = np.max(imshape)
+ focal_length = larger_side / (np.tan(fov_radians / 2) * 2)
+ # intrinsic_matrix 3*3
+ return np.array([
+ [focal_length, 0, imshape[1] / 2],
+ [0, focal_length, imshape[0] / 2],
+ [0, 0, 1],
+ ])
+
+
+def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3
+ camera_matrix = intrinsic_matrix_from_field_of_view((height,width))
+ camera_matrix = np.expand_dims(camera_matrix, axis=0)
+ camera_matrix = np.expand_dims(camera_matrix, axis=0) # 1*1*3*3
+ point_3d = np.expand_dims(point_3d,axis=-1) # n*1024*3*1
+ point_2d = (camera_matrix@point_3d).squeeze(-1)
+ point_2d[:,:,:2] = point_2d[:,:,:2]/point_2d[:,:,2:3]
+ return point_2d[:,:,:] # n*1024*2
+
+
+def get_pose_images(smpl_data, offset):
+ pose_images = []
+ for data in smpl_data:
+ if isinstance(data, np.ndarray):
+ joints3d = data
+ else:
+ joints3d = data.numpy()
+ canvas = np.zeros(shape=(offset[0], offset[1], 3), dtype=np.uint8)
+ joints3d = p3d_to_p2d(joints3d, offset[0], offset[1])
+ canvas = draw_3d_points(canvas, joints3d[0], stickwidth=int(offset[1]/350))
+ pose_images.append(Image.fromarray(canvas))
+ return pose_images
+
+
+def get_control_conditions(poses, h, w):
+ video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
+ control_images = []
+ for idx, pose in enumerate(poses):
+ canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
+ try:
+ joints3d = p3d_to_p2d(pose, h, w)
+ canvas = draw_3d_points(
+ canvas,
+ joints3d[0],
+ stickwidth=int(h / 350),
+ )
+ resized_canvas = cv2.resize(canvas, (w, h))
+ # Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
+ control_images.append(resized_canvas)
+ except Exception as e:
+ print("wrong:", e)
+ control_images.append(Image.fromarray(canvas))
+ control_pixel_values = np.array(control_images)
+ control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
+ print("control_pixel_values.shape", control_pixel_values.shape)
+ #control_pixel_values = video_transforms(control_pixel_values)
+ return control_pixel_values
+
+
+def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True):
+ colors = [
+ [255, 0, 0], # 0
+ [0, 255, 0], # 1
+ [0, 0, 255], # 2
+ [255, 0, 255], # 3
+ [255, 255, 0], # 4
+ [85, 255, 0], # 5
+ [0, 75, 255], # 6
+ [0, 255, 85], # 7
+ [0, 255, 170], # 8
+ [170, 0, 255], # 9
+ [85, 0, 255], # 10
+ [0, 85, 255], # 11
+ [0, 255, 255], # 12
+ [85, 0, 255], # 13
+ [170, 0, 255], # 14
+ [255, 0, 255], # 15
+ [255, 0, 170], # 16
+ [255, 0, 85], # 17
+ ]
+ connetions = [
+ [15,12],[12, 16],[16, 18],[18, 20],[20, 22],
+ [12,17],[17,19],[19,21],
+ [21,23],[12,9],[9,6],
+ [6,3],[3,0],[0,1],
+ [1,4],[4,7],[7,10],[0,2],[2,5],[5,8],[8,11]
+ ]
+ connection_colors = [
+ [255, 0, 0], # 0
+ [0, 255, 0], # 1
+ [0, 0, 255], # 2
+ [255, 255, 0], # 3
+ [255, 0, 255], # 4
+ [0, 255, 0], # 5
+ [0, 85, 255], # 6
+ [255, 175, 0], # 7
+ [0, 0, 255], # 8
+ [255, 85, 0], # 9
+ [0, 255, 85], # 10
+ [255, 0, 255], # 11
+ [255, 0, 0], # 12
+ [0, 175, 255], # 13
+ [255, 255, 0], # 14
+ [0, 0, 255], # 15
+ [0, 255, 0], # 16
+ ]
+
+ # draw point
+ for i in range(len(points)):
+ x,y = points[i][0:2]
+ x,y = int(x),int(y)
+ if i==13 or i == 14:
+ continue
+ cv2.circle(canvas, (x, y), r, colors[i%17], thickness=-1)
+
+ # draw line
+ if draw_line:
+ for i in range(len(connetions)):
+ point1_idx,point2_idx = connetions[i][0:2]
+ point1 = points[point1_idx]
+ point2 = points[point2_idx]
+ Y = [point2[0],point1[0]]
+ X = [point2[1],point1[1]]
+ mX = int(np.mean(X))
+ mY = int(np.mean(Y))
+ length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5
+ angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
+ polygon = cv2.ellipse2Poly((mY, mX), (int(length / 2), stickwidth), int(angle), 0, 360, 1)
+ cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
+
+ return canvas
diff --git a/MTV/motion4d/__init__.py b/MTV/motion4d/__init__.py
new file mode 100644
index 0000000..87ac3b3
--- /dev/null
+++ b/MTV/motion4d/__init__.py
@@ -0,0 +1 @@
+from .vqvae import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
\ No newline at end of file
diff --git a/MTV/motion4d/vqvae.py b/MTV/motion4d/vqvae.py
new file mode 100644
index 0000000..30535e0
--- /dev/null
+++ b/MTV/motion4d/vqvae.py
@@ -0,0 +1,329 @@
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+import numpy as np
+
+
+class Encoder(nn.Module):
+ def __init__(
+ self,
+ in_channels=3,
+ mid_channels=[128, 512],
+ out_channels=3072,
+ downsample_time=[1, 1],
+ downsample_joint=[1, 1],
+ num_attention_heads=8,
+ attention_head_dim=64,
+ dim=3072,
+ ):
+ super(Encoder, self).__init__()
+
+ self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
+ self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
+ self.downsample1 = Downsample(mid_channels[0], mid_channels[0], downsample_time[0], downsample_joint[0])
+ self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
+ self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
+ self.downsample2 = Downsample(mid_channels[1], mid_channels[1], downsample_time[1], downsample_joint[1])
+ self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
+
+ def forward(self, x):
+ x = self.conv_in(x)
+ for resnet in self.resnet1:
+ x = resnet(x)
+ x = self.downsample1(x)
+
+ x = self.resnet2(x)
+ for resnet in self.resnet3:
+ x = resnet(x)
+ x = self.downsample2(x)
+
+ x = self.conv_out(x)
+
+ return x
+
+
+
+class VectorQuantizer(nn.Module):
+ def __init__(self, nb_code, code_dim):
+ super().__init__()
+ self.nb_code = nb_code
+ self.code_dim = code_dim
+ self.mu = 0.99
+ self.reset_codebook()
+ self.reset_count = 0
+ self.usage = torch.zeros((self.nb_code, 1))
+
+ def reset_codebook(self):
+ self.init = False
+ self.code_sum = None
+ self.code_count = None
+ self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).cuda())
+
+ def _tile(self, x):
+ nb_code_x, code_dim = x.shape
+ if nb_code_x < self.nb_code:
+ n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x
+ std = 0.01 / np.sqrt(code_dim)
+ out = x.repeat(n_repeats, 1)
+ out = out + torch.randn_like(out) * std
+ else:
+ out = x
+ return out
+
+ def preprocess(self, x):
+ # [bs, c, f, j] -> [bs * f * j, c]
+ x = x.permute(0, 2, 3, 1).contiguous()
+ x = x.view(-1, x.shape[-1])
+ return x
+
+ def quantize(self, x):
+ # [bs * f * j, dim=3072]
+ # Calculate latent code x_l
+ k_w = self.codebook.t()
+ distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0, keepdim=True)
+ _, code_idx = torch.min(distance, dim=-1)
+ return code_idx
+
+ def dequantize(self, code_idx):
+ x = F.embedding(code_idx, self.codebook) # indexing: [bs * f * j, 32]
+ return x
+
+ def forward(self, x, return_vq=False):
+ bs, c, f, j = x.shape # SMPL data frames: [bs, 3072, f, j]
+
+ # Preprocess
+ x = self.preprocess(x)
+ # return x.view(bs, f*j, c).contiguous(), None
+ assert x.shape[-1] == self.code_dim
+
+ # quantize and dequantize through bottleneck
+ code_idx = self.quantize(x)
+ x_d = self.dequantize(code_idx)
+
+ # Loss
+ commit_loss = F.mse_loss(x, x_d.detach())
+
+ # Passthrough
+ x_d = x + (x_d - x).detach()
+
+ if return_vq:
+ return x_d.view(bs, f*j, c).contiguous(), commit_loss
+ # return (x_d, x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()), commit_loss, perplexity
+
+ # Postprocess
+ x_d = x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()
+
+ return x_d, commit_loss
+
+
+
+
+class Decoder(nn.Module):
+ def __init__(
+ self,
+ in_channels=3072,
+ mid_channels=[512, 128],
+ out_channels=3,
+ upsample_rate=None,
+ frame_upsample_rate=[1.0, 1.0],
+ joint_upsample_rate=[1.0, 1.0],
+ dim=128,
+ attention_head_dim=64,
+ num_attention_heads=8,
+ ):
+ super(Decoder, self).__init__()
+
+ self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
+ self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
+ self.upsample1 = Upsample(mid_channels[0], mid_channels[0], frame_upsample_rate=frame_upsample_rate[0], joint_upsample_rate=joint_upsample_rate[0])
+ self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
+ self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
+ self.upsample2 = Upsample(mid_channels[1], mid_channels[1], frame_upsample_rate=frame_upsample_rate[1], joint_upsample_rate=joint_upsample_rate[1])
+ self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
+
+ def forward(self, x):
+ x = self.conv_in(x)
+ for resnet in self.resnet1:
+ x = resnet(x)
+ x = self.upsample1(x)
+
+ x = self.resnet2(x)
+ for resnet in self.resnet3:
+ x = resnet(x)
+ x = self.upsample2(x)
+
+ x = self.conv_out(x)
+
+ return x
+
+
+class Upsample(nn.Module):
+ def __init__(
+ self,
+ in_channels,
+ out_channels,
+ upsample_rate=None,
+ frame_upsample_rate=None,
+ joint_upsample_rate=None,
+ ):
+ super(Upsample, self).__init__()
+
+ self.upsampler = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
+ self.upsample_rate = upsample_rate
+ self.frame_upsample_rate = frame_upsample_rate
+ self.joint_upsample_rate = joint_upsample_rate
+ self.upsample_rate = upsample_rate
+
+ def forward(self, inputs):
+ if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1:
+ # split first frame
+ x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:]
+
+ if self.upsample_rate is not None:
+ # import pdb; pdb.set_trace()
+ x_first = F.interpolate(x_first, scale_factor=self.upsample_rate)
+ x_rest = F.interpolate(x_rest, scale_factor=self.upsample_rate)
+ else:
+ # import pdb; pdb.set_trace()
+ # x_first = F.interpolate(x_first, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
+ x_rest = F.interpolate(x_rest, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
+ x_first = x_first[:, :, None, :]
+ inputs = torch.cat([x_first, x_rest], dim=2)
+ elif inputs.shape[2] > 1:
+ if self.upsample_rate is not None:
+ inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
+ else:
+ inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
+ else:
+ inputs = inputs.squeeze(2)
+ if self.upsample_rate is not None:
+ inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
+ else:
+ inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="linear", align_corners=True)
+ inputs = inputs[:, :, None, :, :]
+
+ b, c, t, j = inputs.shape
+ inputs = inputs.permute(0, 2, 1, 3).reshape(b * t, c, j)
+ inputs = self.upsampler(inputs)
+ inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3)
+
+ return inputs
+
+
+class Downsample(nn.Module):
+ def __init__(
+ self,
+ in_channels,
+ out_channels,
+ frame_downsample_rate,
+ joint_downsample_rate
+ ):
+ super(Downsample, self).__init__()
+
+ self.frame_downsample_rate = frame_downsample_rate
+ self.joint_downsample_rate = joint_downsample_rate
+ self.joint_downsample = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=self.joint_downsample_rate, padding=1)
+
+ def forward(self, x):
+ # (batch_size, channels, frames, joints) -> (batch_size * joints, channels, frames)
+ if self.frame_downsample_rate > 1:
+ batch_size, channels, frames, joints = x.shape
+ x = x.permute(0, 3, 1, 2).reshape(batch_size * joints, channels, frames)
+ if x.shape[-1] % 2 == 1:
+ x_first, x_rest = x[..., 0], x[..., 1:]
+ if x_rest.shape[-1] > 0:
+ # (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2)
+ x_rest = F.avg_pool1d(x_rest, kernel_size=self.frame_downsample_rate, stride=self.frame_downsample_rate)
+
+ x = torch.cat([x_first[..., None], x_rest], dim=-1)
+ # (batch_size * joints, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, joints)
+ x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
+ else:
+ # (batch_size * joints, channels, frames) -> (batch_size * joints, channels, frames // 2)
+ x = F.avg_pool1d(x, kernel_size=2, stride=2)
+ # (batch_size * joints, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width)
+ x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
+
+ # Pad the tensor
+ # pad = (0, 1)
+ # x = F.pad(x, pad, mode="constant", value=0)
+ batch_size, channels, frames, joints = x.shape
+ # (batch_size, channels, frames, joints) -> (batch_size * frames, channels, joints)
+ x = x.permute(0, 2, 1, 3).reshape(batch_size * frames, channels, joints)
+ x = self.joint_downsample(x)
+ # (batch_size * frames, channels, joints) -> (batch_size, channels, frames, joints)
+ x = x.reshape(batch_size, frames, x.shape[1], x.shape[2]).permute(0, 2, 1, 3)
+ return x
+
+
+
+class ResBlock(nn.Module):
+ def __init__(self,
+ in_channels,
+ out_channels,
+ group_num=32,
+ max_channels=512):
+ super(ResBlock, self).__init__()
+ skip = max(1, max_channels // out_channels - 1)
+ self.block = nn.Sequential(
+ nn.GroupNorm(group_num, in_channels, eps=1e-06, affine=True),
+ nn.SiLU(),
+ nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=skip, dilation=skip),
+ nn.GroupNorm(group_num, out_channels, eps=1e-06, affine=True),
+ nn.SiLU(),
+ nn.Conv2d(out_channels, out_channels, kernel_size=1, stride=1, padding=0),
+ )
+ self.conv_short = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) if in_channels != out_channels else nn.Identity()
+
+ def forward(self, x):
+ hidden_states = self.block(x)
+ if hidden_states.shape != x.shape:
+ x = self.conv_short(x)
+ x = x + hidden_states
+ return x
+
+
+
+class SMPL_VQVAE(nn.Module):
+ def __init__(self, encoder, decoder, vq):
+ super(SMPL_VQVAE, self).__init__()
+
+ self.encoder = encoder
+ self.decoder = decoder
+ self.vq = vq
+
+ def to(self, device):
+ self.encoder = self.encoder.to(device)
+ self.decoder = self.decoder.to(device)
+ self.vq = self.vq.to(device)
+ self.device = device
+ return self
+
+ def encdec_slice_frames(self, x, frame_batch_size, encdec, return_vq):
+ num_frames = x.shape[2]
+ remaining_frames = num_frames % frame_batch_size
+ x_output = []
+
+ for i in range(num_frames // frame_batch_size):
+ remaining_frames = num_frames % frame_batch_size
+ start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames)
+ end_frame = frame_batch_size * (i + 1) + remaining_frames
+ x_intermediate = x[:, :, start_frame:end_frame]
+ x_intermediate = encdec(x_intermediate)
+ x_output.append(x_intermediate)
+ if encdec == self.encoder and self.vq is not None:
+ x_output, loss = self.vq(torch.cat(x_output, dim=2), return_vq=return_vq)
+ return x_output, loss
+ else:
+ return torch.cat(x_output, dim=2), None, None
+
+ def forward(self, x, return_vq=False):
+ x = x.permute(0, 3, 1, 2)
+ x, loss = self.encdec_slice_frames(x, frame_batch_size=8, encdec=self.encoder, return_vq=return_vq)
+
+ if return_vq:
+ return x, loss
+ x, _, _ = self.encdec_slice_frames(x, frame_batch_size=2, encdec=self.decoder, return_vq=return_vq)
+ x = x.permute(0, 2, 3, 1)
+
+ return x, loss
diff --git a/MTV/mtv.py b/MTV/mtv.py
new file mode 100644
index 0000000..915773d
--- /dev/null
+++ b/MTV/mtv.py
@@ -0,0 +1,193 @@
+import torch
+import numpy as np
+from typing import Union, Tuple
+
+
+def get_1d_rotary_pos_embed(
+ dim: int,
+ pos: Union[np.ndarray, int],
+ theta: float = 10000.0,
+ use_real=False,
+ linear_factor=1.0,
+ ntk_factor=1.0,
+ repeat_interleave_real=True,
+ freqs_dtype=torch.float32, # torch.float32, torch.float64 (flux)
+):
+ """
+ Precompute the frequency tensor for complex exponentials (cis) with given dimensions.
+
+ This function calculates a frequency tensor with complex exponentials using the given dimension 'dim' and the end
+ index 'end'. The 'theta' parameter scales the frequencies. The returned tensor contains complex values in complex64
+ data type.
+
+ Args:
+ dim (`int`): Dimension of the frequency tensor.
+ pos (`np.ndarray` or `int`): Position indices for the frequency tensor. [S] or scalar
+ theta (`float`, *optional*, defaults to 10000.0):
+ Scaling factor for frequency computation. Defaults to 10000.0.
+ use_real (`bool`, *optional*):
+ If True, return real part and imaginary part separately. Otherwise, return complex numbers.
+ linear_factor (`float`, *optional*, defaults to 1.0):
+ Scaling factor for the context extrapolation. Defaults to 1.0.
+ ntk_factor (`float`, *optional*, defaults to 1.0):
+ Scaling factor for the NTK-Aware RoPE. Defaults to 1.0.
+ repeat_interleave_real (`bool`, *optional*, defaults to `True`):
+ If `True` and `use_real`, real part and imaginary part are each interleaved with themselves to reach `dim`.
+ Otherwise, they are concateanted with themselves.
+ freqs_dtype (`torch.float32` or `torch.float64`, *optional*, defaults to `torch.float32`):
+ the dtype of the frequency tensor.
+ Returns:
+ `torch.Tensor`: Precomputed frequency tensor with complex exponentials. [S, D/2]
+ """
+ assert dim % 2 == 0
+
+ if isinstance(pos, int):
+ pos = torch.arange(pos)
+ if isinstance(pos, np.ndarray):
+ pos = torch.from_numpy(pos) # type: ignore # [S]
+
+ theta = theta * ntk_factor
+ freqs = (
+ 1.0
+ / (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
+ / linear_factor
+ ) # [D/2]
+ freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
+ if use_real and repeat_interleave_real:
+ freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
+ freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
+ return freqs_cos, freqs_sin
+ elif use_real:
+ freqs_cos = torch.cat([freqs.cos(), freqs.cos()], dim=-1).float() # [S, D]
+ freqs_sin = torch.cat([freqs.sin(), freqs.sin()], dim=-1).float() # [S, D]
+ return freqs_cos, freqs_sin
+ else:
+ freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
+ return freqs_cis
+
+
+def get_3d_rotary_pos_embed(
+ embed_dim, crops_coords, grid_size, temporal_size, theta: int = 10000, use_real: bool = True
+) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
+ """
+ RoPE for video tokens with 3D structure.
+
+ Args:
+ embed_dim: (`int`):
+ The embedding dimension size, corresponding to hidden_size_head.
+ crops_coords (`Tuple[int]`):
+ The top-left and bottom-right coordinates of the crop.
+ grid_size (`Tuple[int]`):
+ The grid size of the spatial positional embedding (height, width).
+ temporal_size (`int`):
+ The size of the temporal dimension.
+ theta (`float`):
+ Scaling factor for frequency computation.
+
+ Returns:
+ `torch.Tensor`: positional embedding with shape `(temporal_size * grid_size[0] * grid_size[1], embed_dim/2)`.
+ """
+ if use_real is not True:
+ raise ValueError(" `use_real = False` is not currently supported for get_3d_rotary_pos_embed")
+ start, stop = crops_coords
+ grid_size_h, grid_size_w = grid_size
+ grid_h = np.linspace(start[0], stop[0], grid_size_h, endpoint=False, dtype=np.float32)
+ grid_w = np.linspace(start[1], stop[1], grid_size_w, endpoint=False, dtype=np.float32)
+ grid_t = np.linspace(0, temporal_size, temporal_size, endpoint=False, dtype=np.float32)
+
+ # Compute dimensions for each axis
+ dim_t = embed_dim // 4
+ dim_h = embed_dim // 8 * 3
+ dim_w = embed_dim // 8 * 3
+
+ # Temporal frequencies
+ freqs_t = get_1d_rotary_pos_embed(dim_t, grid_t, use_real=True)
+ # Spatial frequencies for height and width
+ freqs_h = get_1d_rotary_pos_embed(dim_h, grid_h, use_real=True)
+ freqs_w = get_1d_rotary_pos_embed(dim_w, grid_w, use_real=True)
+
+ # BroadCast and concatenate temporal and spaial frequencie (height and width) into a 3d tensor
+ def combine_time_height_width(freqs_t, freqs_h, freqs_w):
+ freqs_t = freqs_t[:, None, None, :].expand(
+ -1, grid_size_h, grid_size_w, -1
+ ) # temporal_size, grid_size_h, grid_size_w, dim_t
+ freqs_h = freqs_h[None, :, None, :].expand(
+ temporal_size, -1, grid_size_w, -1
+ ) # temporal_size, grid_size_h, grid_size_2, dim_h
+ freqs_w = freqs_w[None, None, :, :].expand(
+ temporal_size, grid_size_h, -1, -1
+ ) # temporal_size, grid_size_h, grid_size_2, dim_w
+
+ freqs = torch.cat(
+ [freqs_t, freqs_h, freqs_w], dim=-1
+ ) # temporal_size, grid_size_h, grid_size_w, (dim_t + dim_h + dim_w)
+ freqs = freqs.view(
+ temporal_size * grid_size_h * grid_size_w, -1
+ ) # (temporal_size * grid_size_h * grid_size_w), (dim_t + dim_h + dim_w)
+ return freqs
+
+ t_cos, t_sin = freqs_t # both t_cos and t_sin has shape: temporal_size, dim_t
+ h_cos, h_sin = freqs_h # both h_cos and h_sin has shape: grid_size_h, dim_h
+ w_cos, w_sin = freqs_w # both w_cos and w_sin has shape: grid_size_w, dim_w
+ cos = combine_time_height_width(t_cos, h_cos, w_cos)
+ sin = combine_time_height_width(t_sin, h_sin, w_sin)
+ return cos, sin
+
+
+def get_3d_motion_spatial_embed(
+ embed_dim: int, num_joints: int, joints_mean: np.ndarray, joints_std: np.ndarray, theta: float = 10000.0
+) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
+ assert embed_dim % 2 == 0 and embed_dim % 3 == 0
+
+ def create_rope_pe(dim, pos, freqs_dtype=torch.float32):
+ if isinstance(pos, np.ndarray):
+ pos = torch.from_numpy(pos)
+ freqs = (
+ 1.0
+ / (theta ** (torch.arange(0, dim, 2, dtype=freqs_dtype, device=pos.device)[: (dim // 2)] / dim))
+ ) # [D/2]
+ freqs = torch.outer(pos, freqs) # type: ignore # [S, D/2]
+ freqs_cos = freqs.cos().repeat_interleave(2, dim=1).float() # [S, D]
+ freqs_sin = freqs.sin().repeat_interleave(2, dim=1).float() # [S, D]
+ return freqs_cos, freqs_sin
+
+ pos_x = joints_mean[:, 0]
+ pos_y = joints_mean[:, 1]
+ pos_z = joints_mean[:, 2]
+
+ normalized_pos_x = (pos_x - pos_x.mean())
+ normalized_pos_y = (pos_y - pos_y.mean())
+ normalized_pos_z = (pos_z - pos_z.mean())
+
+ freqs_cos_x, freqs_sin_x = create_rope_pe(embed_dim // 3, normalized_pos_x)
+ freqs_cos_y, freqs_sin_y = create_rope_pe(embed_dim // 3, normalized_pos_y)
+ freqs_cos_z, freqs_sin_z = create_rope_pe(embed_dim // 3, normalized_pos_z)
+
+ freqs_cos = torch.cat([freqs_cos_x, freqs_cos_y, freqs_cos_z], dim=-1)
+ freqs_sin = torch.cat([freqs_sin_x, freqs_sin_y, freqs_sin_z], dim=-1)
+
+ return freqs_cos, freqs_sin
+
+def prepare_motion_embeddings(num_frames, num_joints, joints_mean, joints_std, theta=10000, device='cuda'):
+ time_embed = get_1d_rotary_pos_embed(44, num_frames, theta, use_real=True)
+ time_embed_cos = time_embed[0][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1)
+ time_embed_sin = time_embed[1][:, None, :].expand(-1, num_joints, -1).reshape(num_frames*num_joints, -1)
+ spatial_motion_embed = get_3d_motion_spatial_embed(84, num_joints, joints_mean, joints_std, theta)
+ spatial_embed_cos = spatial_motion_embed[0][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1)
+ spatial_embed_sin = spatial_motion_embed[1][None, :, :].expand(num_frames, -1, -1).reshape(num_frames*num_joints, -1)
+ motion_embed_cos = torch.cat([time_embed_cos, spatial_embed_cos], dim=-1).to(device=device)
+ motion_embed_sin = torch.cat([time_embed_sin, spatial_embed_sin], dim=-1).to(device=device)
+ return motion_embed_cos, motion_embed_sin
+
+def apply_rotary_emb(x, freqs_cis):
+ cos, sin = freqs_cis # [S, D]
+ cos = cos[None, None]
+ sin = sin[None, None]
+ cos, sin = cos.to(x.device), sin.to(x.device)
+
+ x_real, x_imag = x.reshape(*x.shape[:-1], -1, 2).unbind(-1) # [B, S, H, D//2]
+ x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
+
+ out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
+
+ return out
\ No newline at end of file
diff --git a/MTV/nlf.py b/MTV/nlf.py
new file mode 100644
index 0000000..e69de29
diff --git a/MTV/nodes.py b/MTV/nodes.py
new file mode 100644
index 0000000..bb6ae4c
--- /dev/null
+++ b/MTV/nodes.py
@@ -0,0 +1,242 @@
+import os
+import torch
+import gc
+from ..utils import log, dict_to_device
+import numpy as np
+from accelerate import init_empty_weights
+from accelerate.utils import set_module_tensor_to_device
+
+import comfy.model_management as mm
+from comfy.utils import load_torch_file
+import folder_paths
+
+script_directory = os.path.dirname(os.path.abspath(__file__))
+device = mm.get_torch_device()
+offload_device = mm.unet_offload_device()
+
+local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript")
+
+from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
+from .mtv import prepare_motion_embeddings
+
+class DownloadAndLoadNLFModel:
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ "required": {
+ "url": (
+ [
+ "https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"
+ ],
+ )
+ },
+ }
+
+ RETURN_TYPES = ("NLFMODEL",)
+ RETURN_NAMES = ("nlf_model", )
+ FUNCTION = "loadmodel"
+ CATEGORY = "WanVideoWrapper"
+
+ def loadmodel(self, url):
+
+ if not os.path.exists(local_model_path):
+ log.info(f"Downloading NLF model to: {local_model_path}")
+ import requests
+ os.makedirs(os.path.dirname(local_model_path), exist_ok=True)
+ response = requests.get(url)
+ if response.status_code == 200:
+ with open(local_model_path, "wb") as f:
+ f.write(response.content)
+ else:
+ print("Failed to download file:", response.status_code)
+
+ model = torch.jit.load(local_model_path).eval()
+
+ return (model,)
+
+class LoadNLFModel:
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ "required": {
+ "path": ("STRING", {"default": local_model_path}),
+ },
+ }
+
+ RETURN_TYPES = ("NLFMODEL",)
+ RETURN_NAMES = ("nlf_model", )
+ FUNCTION = "loadmodel"
+ CATEGORY = "WanVideoWrapper"
+
+ def loadmodel(self, path):
+ model = torch.jit.load(path).eval()
+
+ return model,
+
+class LoadVQVAE:
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ "required": {
+ "model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
+ },
+ }
+
+ RETURN_TYPES = ("VQVAE",)
+ RETURN_NAMES = ("vqvae", )
+ FUNCTION = "loadmodel"
+ CATEGORY = "WanVideoWrapper"
+
+ def loadmodel(self, model_name):
+ model_path = folder_paths.get_full_path("vae", model_name)
+ vae_sd = load_torch_file(model_path, safe_load=True)
+
+ # Get motion tokenizer
+ motion_encoder = Encoder(
+ in_channels=3,
+ mid_channels=[128, 512],
+ out_channels=3072,
+ downsample_time=[2, 2],
+ downsample_joint=[1, 1]
+ )
+ motion_quant = VectorQuantizer(nb_code=8192, code_dim=3072)
+ motion_decoder = Decoder(
+ in_channels=3072,
+ mid_channels=[512, 128],
+ out_channels=3,
+ upsample_rate=2.0,
+ frame_upsample_rate=[2.0, 2.0],
+ joint_upsample_rate=[1.0, 1.0]
+ )
+
+ vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device)
+ vqvae.load_state_dict(vae_sd, strict=True)
+
+ return vqvae,
+
+class MTVCrafterEncodePoses:
+ @classmethod
+ def INPUT_TYPES(s):
+ return {
+ "required": {
+ "vqvae": ("VQVAE", {"tooltip": "VQVAE model"}),
+ "poses": ("NLFPRED", {"tooltip": "Input poses for the model"}),
+ },
+ }
+
+ RETURN_TYPES = ("MTVCRAFTERMOTION", "NLFPRED")
+ RETURN_NAMES = ("mtvcrafter_motion", "pose_results")
+ FUNCTION = "encode"
+ CATEGORY = "WanVideoWrapper"
+
+ def encode(self, vqvae, poses):
+
+ # import pickle
+ # with open(os.path.join(script_directory, "data", "sampled_data.pkl"), 'rb') as f:
+ # data_list = pickle.load(f)
+ # if not isinstance(data_list, list):
+ # data_list = [data_list]
+ # print(data_list)
+
+ # smpl_poses = data_list[1]['pose']
+
+ global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3)
+ global_std = np.load(os.path.join(script_directory, "data", "std.npy"))
+
+ smpl_poses = []
+ for pose in poses['joints3d_nonparam'][0]:
+ smpl_poses.append(pose[0].cpu().numpy())
+ smpl_poses = np.array(smpl_poses)
+
+ norm_poses = torch.tensor((smpl_poses - global_mean) / global_std).unsqueeze(0)
+ print(f"norm_poses shape: {norm_poses.shape}, dtype: {norm_poses.dtype}")
+
+ vqvae.to(device)
+ motion_tokens, vq_loss = vqvae(norm_poses.to(device), return_vq=True)
+
+ recon_motion = vqvae(norm_poses.to(device))[0][0].to(dtype=torch.float32).cpu().detach() * global_std + global_mean
+ vqvae.to(offload_device)
+
+ poses_dict = {
+ 'mtv_motion_tokens': motion_tokens,
+ 'global_mean': global_mean,
+ 'global_std': global_std
+ }
+
+ return poses_dict, recon_motion
+
+
+class NLFPredict:
+ @classmethod
+ def INPUT_TYPES(s):
+ return {"required": {
+ "model": ("NLFMODEL",),
+ "images": ("IMAGE", {"tooltip": "Input images for the model"}),
+ },
+ }
+
+ RETURN_TYPES = ("NLFPRED", )
+ RETURN_NAMES = ("pose_results",)
+ FUNCTION = "predict"
+ CATEGORY = "WanVideoWrapper"
+
+ def predict(self, model, images):
+
+ model.to(device)
+ pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
+ model.to(offload_device)
+
+ pred = dict_to_device(pred, offload_device)
+
+ pose_results = {
+ 'joints3d_nonparam': [],
+ }
+ # Collect pose data
+ for key in pose_results.keys():
+ if key in pred:
+ pose_results[key].append(pred[key])
+ else:
+ pose_results[key].append(None)
+
+ return (pose_results,)
+
+class DrawNLFPoses:
+ @classmethod
+ def INPUT_TYPES(s):
+ return {"required": {
+ "poses": ("NLFPRED", {"tooltip": "Input poses for the model"}),
+ "width": ("INT", {"default": 512}),
+ "height": ("INT", {"default": 512}),
+ },
+ }
+
+ RETURN_TYPES = ("IMAGE", )
+ RETURN_NAMES = ("image",)
+ FUNCTION = "predict"
+ CATEGORY = "WanVideoWrapper"
+
+ def predict(self, poses, width, height):
+ from .draw_pose import get_control_conditions
+ print(type(poses))
+ if isinstance(poses, dict):
+ pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
+ else:
+ pose_input = poses
+ control_conditions = get_control_conditions(pose_input, height, width)
+
+ return (control_conditions,)
+
+NODE_CLASS_MAPPINGS = {
+ "DownloadAndLoadNLFModel": DownloadAndLoadNLFModel,
+ "NLFPredict": NLFPredict,
+ "DrawNLFPoses": DrawNLFPoses,
+ "LoadVQVAE": LoadVQVAE,
+ "MTVCrafterEncodePoses": MTVCrafterEncodePoses
+ }
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "DownloadAndLoadNLFModel": "(Download)Load NLF Model",
+ "NLFPredict": "NLF Predict",
+ "DrawNLFPoses": "Draw NLF Poses",
+ "LoadVQVAE": "Load VQVAE",
+ "MTVCrafterEncodePoses": "MTV Crafter Encode Poses"
+}
diff --git a/__init__.py b/__init__.py
index 112b36e..1f81508 100644
--- a/__init__.py
+++ b/__init__.py
@@ -35,6 +35,13 @@ except Exception as e:
UNIANIMATE_NODE_CLASS_MAPPINGS = {}
UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS = {}
+try:
+ from .MTV.nodes import NODE_CLASS_MAPPINGS as MTV_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MTV_NODE_DISPLAY_NAME_MAPPINGS
+except Exception as e:
+ print(f"MTV nodes not available due to error in importing them: {e}")
+ MTV_NODE_CLASS_MAPPINGS = {}
+ MTV_NODE_DISPLAY_NAME_MAPPINGS = {}
+
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
@@ -50,6 +57,7 @@ NODE_CLASS_MAPPINGS.update(UTILITY_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
+NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
@@ -65,7 +73,7 @@ NODE_DISPLAY_NAME_MAPPINGS.update(MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UTILITY_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
-NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
-
+NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
+NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
\ No newline at end of file
diff --git a/custom_linear.py b/custom_linear.py
index ae5e81d..5236458 100644
--- a/custom_linear.py
+++ b/custom_linear.py
@@ -1,6 +1,7 @@
import torch
import torch.nn as nn
from accelerate import init_empty_weights
+from comfy.ops import cast_bias_weight
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None):
@@ -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)
diff --git a/example_workflows/example_inputs/MTV_crafter_example_pose.mp4 b/example_workflows/example_inputs/MTV_crafter_example_pose.mp4
new file mode 100644
index 0000000..9776fc1
Binary files /dev/null and b/example_workflows/example_inputs/MTV_crafter_example_pose.mp4 differ
diff --git a/example_workflows/wanvideo_FLF2V_720P_example_02.json b/example_workflows/wanvideo_FLF2V_720P_example_02.json
index 8d37929..a2fd652 100644
--- a/example_workflows/wanvideo_FLF2V_720P_example_02.json
+++ b/example_workflows/wanvideo_FLF2V_720P_example_02.json
@@ -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",
+ "
| Output: | 1 x 640 x 640 | 4.69MB |
"
]
},
{
@@ -1489,7 +1505,7 @@
],
"size": [
270,
- 286
+ 336
],
"flags": {},
"order": 28,
@@ -1564,9 +1580,199 @@
"0, 0, 0",
"center",
16,
- "cpu"
+ "cpu",
+ "| Output: | 1 x 640 x 640 | 4.69MB |
"
]
},
+ {
+ "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",
diff --git a/example_workflows/wanvideo_MTV_Crafter_example_WIP.json b/example_workflows/wanvideo_MTV_Crafter_example_WIP.json
new file mode 100644
index 0000000..98b8ff2
--- /dev/null
+++ b/example_workflows/wanvideo_MTV_Crafter_example_WIP.json
@@ -0,0 +1,2275 @@
+{
+ "id": "7c6603fc-8871-4e88-a3dd-ac43e7650d8e",
+ "revision": 0,
+ "last_node_id": 198,
+ "last_link_id": 353,
+ "nodes": [
+ {
+ "id": 158,
+ "type": "WanVideoVAELoader",
+ "pos": [
+ -498.183349609375,
+ -1420.46484375
+ ],
+ "size": [
+ 270,
+ 82
+ ],
+ "flags": {},
+ "order": 0,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "compile_args",
+ "shape": 7,
+ "type": "WANCOMPILEARGS",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "vae",
+ "type": "WANVAE",
+ "links": [
+ 279,
+ 284
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoVAELoader"
+ },
+ "widgets_values": [
+ "wanvideo\\Wan2_1_VAE_bf16.safetensors",
+ "bf16"
+ ],
+ "color": "#322",
+ "bgcolor": "#533"
+ },
+ {
+ "id": 172,
+ "type": "CLIPVisionLoader",
+ "pos": [
+ -169.79689025878906,
+ -1124.9453125
+ ],
+ "size": [
+ 270,
+ 58
+ ],
+ "flags": {},
+ "order": 1,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "CLIP_VISION",
+ "type": "CLIP_VISION",
+ "links": [
+ 301
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfy-core",
+ "ver": "0.3.50",
+ "Node name for S&R": "CLIPVisionLoader"
+ },
+ "widgets_values": [
+ "clip_vision_h.safetensors"
+ ],
+ "color": "#233",
+ "bgcolor": "#355"
+ },
+ {
+ "id": 171,
+ "type": "WanVideoClipVisionEncode",
+ "pos": [
+ -171.7789764404297,
+ -988.2925415039062
+ ],
+ "size": [
+ 280.9771423339844,
+ 262
+ ],
+ "flags": {},
+ "order": 25,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "clip_vision",
+ "type": "CLIP_VISION",
+ "link": 301
+ },
+ {
+ "name": "image_1",
+ "type": "IMAGE",
+ "link": 344
+ },
+ {
+ "name": "image_2",
+ "shape": 7,
+ "type": "IMAGE",
+ "link": null
+ },
+ {
+ "name": "negative_image",
+ "shape": 7,
+ "type": "IMAGE",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "image_embeds",
+ "type": "WANVIDIMAGE_CLIPEMBEDS",
+ "links": []
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoClipVisionEncode"
+ },
+ "widgets_values": [
+ 1,
+ 1,
+ "center",
+ "average",
+ true,
+ 0,
+ 0.5
+ ],
+ "color": "#2a363b",
+ "bgcolor": "#3f5159"
+ },
+ {
+ "id": 140,
+ "type": "NLFPredict",
+ "pos": [
+ -722.2062377929688,
+ -439.0661926269531
+ ],
+ "size": [
+ 186.79513549804688,
+ 46
+ ],
+ "flags": {},
+ "order": 19,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "model",
+ "type": "NLFMODEL",
+ "link": 245
+ },
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": 294
+ }
+ ],
+ "outputs": [
+ {
+ "name": "pose_results",
+ "type": "NLFPRED",
+ "links": [
+ 260
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "9a7dba06c13c79cf35d76a1837fd115fc247df43",
+ "Node name for S&R": "NLFPredict"
+ },
+ "widgets_values": []
+ },
+ {
+ "id": 146,
+ "type": "LoadVQVAE",
+ "pos": [
+ -494.146240234375,
+ -596.1441650390625
+ ],
+ "size": [
+ 324.502685546875,
+ 58
+ ],
+ "flags": {},
+ "order": 2,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "vqvae",
+ "type": "VQVAE",
+ "links": [
+ 259
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "9a7dba06c13c79cf35d76a1837fd115fc247df43",
+ "Node name for S&R": "LoadVQVAE"
+ },
+ "widgets_values": [
+ "wanvideo\\MTV_Crafter_4DMoT_VQVAE_fp32.safetensors"
+ ],
+ "color": "#322",
+ "bgcolor": "#533"
+ },
+ {
+ "id": 147,
+ "type": "MTVCrafterEncodePoses",
+ "pos": [
+ -446.90484619140625,
+ -449.3239440917969
+ ],
+ "size": [
+ 245.973876953125,
+ 78
+ ],
+ "flags": {},
+ "order": 22,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "vqvae",
+ "type": "VQVAE",
+ "link": 259
+ },
+ {
+ "name": "poses",
+ "type": "NLFPRED",
+ "link": 260
+ }
+ ],
+ "outputs": [
+ {
+ "name": "mtvcrafter_motion",
+ "type": "MTVCRAFTERMOTION",
+ "links": [
+ 293
+ ]
+ },
+ {
+ "name": "pose_results",
+ "type": "NLFPRED",
+ "links": [
+ 296
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "MTVCrafterEncodePoses"
+ },
+ "widgets_values": [],
+ "color": "#322",
+ "bgcolor": "#533"
+ },
+ {
+ "id": 139,
+ "type": "DownloadAndLoadNLFModel",
+ "pos": [
+ -1100.665771484375,
+ -593.3814086914062
+ ],
+ "size": [
+ 550.6461791992188,
+ 58
+ ],
+ "flags": {},
+ "order": 3,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "nlf_model",
+ "type": "NLFMODEL",
+ "links": [
+ 245
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "9a7dba06c13c79cf35d76a1837fd115fc247df43",
+ "Node name for S&R": "DownloadAndLoadNLFModel"
+ },
+ "widgets_values": [
+ "https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"
+ ]
+ },
+ {
+ "id": 175,
+ "type": "WanVideoContextOptions",
+ "pos": [
+ 1080.3250732421875,
+ -1423.6617431640625
+ ],
+ "size": [
+ 275.783203125,
+ 202
+ ],
+ "flags": {},
+ "order": 4,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "reference_latent",
+ "shape": 7,
+ "type": "LATENT",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "context_options",
+ "type": "WANVIDCONTEXT",
+ "links": []
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoContextOptions"
+ },
+ "widgets_values": [
+ "uniform_standard",
+ 49,
+ 4,
+ 24,
+ true,
+ true,
+ "linear"
+ ]
+ },
+ {
+ "id": 170,
+ "type": "VHS_VideoCombine",
+ "pos": [
+ -98.18639373779297,
+ -390.89251708984375
+ ],
+ "size": [
+ 237.98948669433594,
+ 565.989501953125
+ ],
+ "flags": {},
+ "order": 27,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": 297
+ },
+ {
+ "name": "audio",
+ "shape": 7,
+ "type": "AUDIO",
+ "link": null
+ },
+ {
+ "name": "meta_batch",
+ "shape": 7,
+ "type": "VHS_BatchManager",
+ "link": null
+ },
+ {
+ "name": "vae",
+ "shape": 7,
+ "type": "VAE",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "Filenames",
+ "type": "VHS_FILENAMES",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-videohelpersuite",
+ "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523",
+ "Node name for S&R": "VHS_VideoCombine"
+ },
+ "widgets_values": {
+ "frame_rate": 16,
+ "loop_count": 0,
+ "filename_prefix": "NLFtest",
+ "format": "video/h264-mp4",
+ "pix_fmt": "yuv420p",
+ "crf": 19,
+ "save_metadata": true,
+ "trim_to_audio": false,
+ "pingpong": false,
+ "save_output": false,
+ "videopreview": {
+ "hidden": false,
+ "paused": false,
+ "params": {
+ "filename": "NLFtest_00010.mp4",
+ "subfolder": "",
+ "type": "temp",
+ "format": "video/h264-mp4",
+ "frame_rate": 16,
+ "workflow": "NLFtest_00010.png",
+ "fullpath": "N:\\AI\\ComfyUI\\temp\\NLFtest_00010.mp4"
+ }
+ }
+ }
+ },
+ {
+ "id": 169,
+ "type": "DrawNLFPoses",
+ "pos": [
+ -500.3387756347656,
+ -263.48956298828125
+ ],
+ "size": [
+ 270,
+ 82
+ ],
+ "flags": {},
+ "order": 24,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "poses",
+ "type": "NLFPRED",
+ "link": 296
+ },
+ {
+ "name": "width",
+ "type": "INT",
+ "widget": {
+ "name": "width"
+ },
+ "link": 322
+ },
+ {
+ "name": "height",
+ "type": "INT",
+ "widget": {
+ "name": "height"
+ },
+ "link": 323
+ }
+ ],
+ "outputs": [
+ {
+ "name": "image",
+ "type": "IMAGE",
+ "links": [
+ 297,
+ 336
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "DrawNLFPoses"
+ },
+ "widgets_values": [
+ 1024,
+ 1024
+ ]
+ },
+ {
+ "id": 174,
+ "type": "ImageConcatMulti",
+ "pos": [
+ 1450.1494140625,
+ -865.9163818359375
+ ],
+ "size": [
+ 270,
+ 150
+ ],
+ "flags": {},
+ "order": 28,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "image_1",
+ "type": "IMAGE",
+ "link": 346
+ },
+ {
+ "name": "image_2",
+ "shape": 7,
+ "type": "IMAGE",
+ "link": 336
+ }
+ ],
+ "outputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "links": [
+ 339
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "2f7300dc546ec2d36fa8b0feebe493d41026524c"
+ },
+ "widgets_values": [
+ 2,
+ "down",
+ true,
+ null
+ ]
+ },
+ {
+ "id": 186,
+ "type": "ImageConcatMulti",
+ "pos": [
+ 1457.7677001953125,
+ -646.75439453125
+ ],
+ "size": [
+ 270,
+ 150
+ ],
+ "flags": {},
+ "order": 32,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "image_1",
+ "type": "IMAGE",
+ "link": 338
+ },
+ {
+ "name": "image_2",
+ "shape": 7,
+ "type": "IMAGE",
+ "link": 339
+ }
+ ],
+ "outputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "links": [
+ 337
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "2f7300dc546ec2d36fa8b0feebe493d41026524c"
+ },
+ "widgets_values": [
+ 2,
+ "left",
+ true,
+ null
+ ]
+ },
+ {
+ "id": 164,
+ "type": "WanVideoSetBlockSwap",
+ "pos": [
+ -60.36056900024414,
+ -1537.5782470703125
+ ],
+ "size": [
+ 242.3073272705078,
+ 46
+ ],
+ "flags": {},
+ "order": 18,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "model",
+ "type": "WANVIDEOMODEL",
+ "link": 351
+ },
+ {
+ "name": "block_swap_args",
+ "shape": 7,
+ "type": "BLOCKSWAPARGS",
+ "link": 290
+ }
+ ],
+ "outputs": [
+ {
+ "name": "model",
+ "type": "WANVIDEOMODEL",
+ "links": [
+ 289
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoSetBlockSwap"
+ },
+ "widgets_values": [],
+ "color": "#223",
+ "bgcolor": "#335"
+ },
+ {
+ "id": 165,
+ "type": "WanVideoBlockSwap",
+ "pos": [
+ -77.65693664550781,
+ -1804.4154052734375
+ ],
+ "size": [
+ 281.404296875,
+ 202
+ ],
+ "flags": {},
+ "order": 5,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "block_swap_args",
+ "type": "BLOCKSWAPARGS",
+ "links": [
+ 290
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoBlockSwap"
+ },
+ "widgets_values": [
+ 10,
+ false,
+ false,
+ false,
+ 0,
+ 1,
+ false
+ ],
+ "color": "#223",
+ "bgcolor": "#335"
+ },
+ {
+ "id": 190,
+ "type": "GetLatentRangeFromBatch",
+ "pos": [
+ 625.1795043945312,
+ -400.8965148925781
+ ],
+ "size": [
+ 286.6646423339844,
+ 82
+ ],
+ "flags": {},
+ "order": 6,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "latents",
+ "type": "LATENT",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "LATENT",
+ "type": "LATENT",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "876a6dd2929d88ec35c09c7cddc6c360f4b27013",
+ "Node name for S&R": "GetLatentRangeFromBatch"
+ },
+ "widgets_values": [
+ -1,
+ 1
+ ]
+ },
+ {
+ "id": 191,
+ "type": "GetImageRangeFromBatch",
+ "pos": [
+ 487.84234619140625,
+ -205.70770263671875
+ ],
+ "size": [
+ 340.3267517089844,
+ 102
+ ],
+ "flags": {},
+ "order": 7,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "shape": 7,
+ "type": "IMAGE",
+ "link": null
+ },
+ {
+ "name": "masks",
+ "shape": 7,
+ "type": "MASK",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": null
+ },
+ {
+ "name": "MASK",
+ "type": "MASK",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "876a6dd2929d88ec35c09c7cddc6c360f4b27013",
+ "Node name for S&R": "GetImageRangeFromBatch"
+ },
+ "widgets_values": [
+ -1,
+ 1
+ ]
+ },
+ {
+ "id": 161,
+ "type": "WanVideoDecode",
+ "pos": [
+ 1438.688232421875,
+ -1154.7078857421875
+ ],
+ "size": [
+ 270,
+ 198
+ ],
+ "flags": {},
+ "order": 31,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "vae",
+ "type": "WANVAE",
+ "link": 284
+ },
+ {
+ "name": "samples",
+ "type": "LATENT",
+ "link": 285
+ }
+ ],
+ "outputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "links": [
+ 338
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoDecode"
+ },
+ "widgets_values": [
+ false,
+ 272,
+ 272,
+ 144,
+ 128,
+ "default"
+ ]
+ },
+ {
+ "id": 193,
+ "type": "VHS_SplitImages",
+ "pos": [
+ 524.53759765625,
+ -21.65776252746582
+ ],
+ "size": [
+ 210,
+ 118
+ ],
+ "flags": {},
+ "order": 8,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE_A",
+ "type": "IMAGE",
+ "links": []
+ },
+ {
+ "name": "A_count",
+ "type": "INT",
+ "links": null
+ },
+ {
+ "name": "IMAGE_B",
+ "type": "IMAGE",
+ "links": null
+ },
+ {
+ "name": "B_count",
+ "type": "INT",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-videohelpersuite",
+ "ver": "8e4d79471bf1952154768e8435a9300077b534fa",
+ "Node name for S&R": "VHS_SplitImages"
+ },
+ "widgets_values": {
+ "split_index": -1
+ }
+ },
+ {
+ "id": 157,
+ "type": "WanVideoImageToVideoEncode",
+ "pos": [
+ 218.52210998535156,
+ -1111.871826171875
+ ],
+ "size": [
+ 308.2320251464844,
+ 390
+ ],
+ "flags": {},
+ "order": 26,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "vae",
+ "shape": 7,
+ "type": "WANVAE",
+ "link": 279
+ },
+ {
+ "name": "clip_embeds",
+ "shape": 7,
+ "type": "WANVIDIMAGE_CLIPEMBEDS",
+ "link": null
+ },
+ {
+ "name": "start_image",
+ "shape": 7,
+ "type": "IMAGE",
+ "link": 349
+ },
+ {
+ "name": "end_image",
+ "shape": 7,
+ "type": "IMAGE",
+ "link": null
+ },
+ {
+ "name": "control_embeds",
+ "shape": 7,
+ "type": "WANVIDIMAGE_EMBEDS",
+ "link": null
+ },
+ {
+ "name": "temporal_mask",
+ "shape": 7,
+ "type": "MASK",
+ "link": null
+ },
+ {
+ "name": "extra_latents",
+ "shape": 7,
+ "type": "LATENT",
+ "link": null
+ },
+ {
+ "name": "add_cond_latents",
+ "shape": 7,
+ "type": "ADD_COND_LATENTS",
+ "link": null
+ },
+ {
+ "name": "width",
+ "type": "INT",
+ "widget": {
+ "name": "width"
+ },
+ "link": 307
+ },
+ {
+ "name": "height",
+ "type": "INT",
+ "widget": {
+ "name": "height"
+ },
+ "link": 308
+ },
+ {
+ "name": "num_frames",
+ "type": "INT",
+ "widget": {
+ "name": "num_frames"
+ },
+ "link": 321
+ }
+ ],
+ "outputs": [
+ {
+ "name": "image_embeds",
+ "type": "WANVIDIMAGE_EMBEDS",
+ "links": [
+ 291
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoImageToVideoEncode"
+ },
+ "widgets_values": [
+ 608,
+ 1088,
+ 121,
+ 0.03,
+ 1,
+ 1,
+ true,
+ true,
+ false
+ ],
+ "color": "#322",
+ "bgcolor": "#533"
+ },
+ {
+ "id": 188,
+ "type": "Reroute",
+ "pos": [
+ -264.114013671875,
+ -1282.322509765625
+ ],
+ "size": [
+ 75,
+ 26
+ ],
+ "flags": {},
+ "order": 23,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "",
+ "type": "*",
+ "link": 343
+ }
+ ],
+ "outputs": [
+ {
+ "name": "",
+ "type": "IMAGE",
+ "links": [
+ 344,
+ 346,
+ 349
+ ]
+ }
+ ],
+ "properties": {
+ "showOutputText": false,
+ "horizontal": false
+ }
+ },
+ {
+ "id": 160,
+ "type": "ImageResizeKJv2",
+ "pos": [
+ -1132.5859375,
+ -1305.7125244140625
+ ],
+ "size": [
+ 270,
+ 336
+ ],
+ "flags": {},
+ "order": 20,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "image",
+ "type": "IMAGE",
+ "link": 280
+ },
+ {
+ "name": "mask",
+ "shape": 7,
+ "type": "MASK",
+ "link": null
+ },
+ {
+ "name": "width",
+ "type": "INT",
+ "widget": {
+ "name": "width"
+ },
+ "link": 316
+ },
+ {
+ "name": "height",
+ "type": "INT",
+ "widget": {
+ "name": "height"
+ },
+ "link": 317
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 343
+ ]
+ },
+ {
+ "name": "width",
+ "type": "INT",
+ "links": [
+ 307
+ ]
+ },
+ {
+ "name": "height",
+ "type": "INT",
+ "links": [
+ 308
+ ]
+ },
+ {
+ "name": "mask",
+ "type": "MASK",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "2f7300dc546ec2d36fa8b0feebe493d41026524c",
+ "Node name for S&R": "ImageResizeKJv2"
+ },
+ "widgets_values": [
+ 640,
+ 640,
+ "lanczos",
+ "pad",
+ "255,255,255",
+ "center",
+ 2,
+ "cpu",
+ "| Output: | 1 x 640 x 640 | 4.69MB |
"
+ ]
+ },
+ {
+ "id": 155,
+ "type": "WanVideoSetLoRAs",
+ "pos": [
+ 359.7154541015625,
+ -1536.7674560546875
+ ],
+ "size": [
+ 174.53378295898438,
+ 46
+ ],
+ "flags": {},
+ "order": 21,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "model",
+ "type": "WANVIDEOMODEL",
+ "link": 289
+ },
+ {
+ "name": "lora",
+ "shape": 7,
+ "type": "WANVIDLORA",
+ "link": 277
+ }
+ ],
+ "outputs": [
+ {
+ "name": "model",
+ "type": "WANVIDEOMODEL",
+ "links": [
+ 276
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoSetLoRAs"
+ },
+ "widgets_values": [],
+ "color": "#223",
+ "bgcolor": "#335"
+ },
+ {
+ "id": 156,
+ "type": "WanVideoLoraSelect",
+ "pos": [
+ 258.90850830078125,
+ -1785.316650390625
+ ],
+ "size": [
+ 648.00537109375,
+ 150
+ ],
+ "flags": {},
+ "order": 9,
+ "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": [
+ 277
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoLoraSelect"
+ },
+ "widgets_values": [
+ "lightx2v_I2V_not_clamped_rank_64_fp16_00001_.safetensors",
+ 1,
+ false,
+ false
+ ],
+ "color": "#223",
+ "bgcolor": "#335"
+ },
+ {
+ "id": 138,
+ "type": "VHS_LoadVideo",
+ "pos": [
+ -1801.016845703125,
+ -651.8135375976562
+ ],
+ "size": [
+ 408.2240295410156,
+ 996.921630859375
+ ],
+ "flags": {},
+ "order": 10,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "meta_batch",
+ "shape": 7,
+ "type": "VHS_BatchManager",
+ "link": null
+ },
+ {
+ "name": "vae",
+ "shape": 7,
+ "type": "VAE",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 341
+ ]
+ },
+ {
+ "name": "frame_count",
+ "type": "INT",
+ "links": [
+ 321
+ ]
+ },
+ {
+ "name": "audio",
+ "type": "AUDIO",
+ "links": null
+ },
+ {
+ "name": "video_info",
+ "type": "VHS_VIDEOINFO",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-videohelpersuite",
+ "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523",
+ "Node name for S&R": "VHS_LoadVideo"
+ },
+ "widgets_values": {
+ "video": "MTV_crafter_example_pose_00001.mp4",
+ "force_rate": 0,
+ "custom_width": 0,
+ "custom_height": 0,
+ "frame_load_cap": 49,
+ "skip_first_frames": 0,
+ "select_every_nth": 1,
+ "format": "AnimateDiff",
+ "choose video to upload": "image",
+ "videopreview": {
+ "hidden": false,
+ "paused": false,
+ "params": {
+ "filename": "MTV_crafter_example_pose_00001.mp4",
+ "type": "input",
+ "format": "video/mp4",
+ "force_rate": 0,
+ "custom_width": 0,
+ "custom_height": 0,
+ "frame_load_cap": 49,
+ "skip_first_frames": 0,
+ "select_every_nth": 1
+ }
+ }
+ }
+ },
+ {
+ "id": 144,
+ "type": "ImageResizeKJv2",
+ "pos": [
+ -1095.5374755859375,
+ -422.9293212890625
+ ],
+ "size": [
+ 270,
+ 336
+ ],
+ "flags": {},
+ "order": 17,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "image",
+ "type": "IMAGE",
+ "link": 341
+ },
+ {
+ "name": "mask",
+ "shape": 7,
+ "type": "MASK",
+ "link": null
+ },
+ {
+ "name": "width",
+ "type": "INT",
+ "widget": {
+ "name": "width"
+ },
+ "link": 318
+ },
+ {
+ "name": "height",
+ "type": "INT",
+ "widget": {
+ "name": "height"
+ },
+ "link": 319
+ }
+ ],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 294
+ ]
+ },
+ {
+ "name": "width",
+ "type": "INT",
+ "links": [
+ 316,
+ 322
+ ]
+ },
+ {
+ "name": "height",
+ "type": "INT",
+ "links": [
+ 317,
+ 323
+ ]
+ },
+ {
+ "name": "mask",
+ "type": "MASK",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "2f7300dc546ec2d36fa8b0feebe493d41026524c",
+ "Node name for S&R": "ImageResizeKJv2"
+ },
+ "widgets_values": [
+ 832,
+ 480,
+ "lanczos",
+ "pad",
+ "0, 0, 0",
+ "center",
+ 2,
+ "cpu",
+ "| Output: | 49 x 640 x 640 | 229.69MB |
"
+ ]
+ },
+ {
+ "id": 162,
+ "type": "VHS_VideoCombine",
+ "pos": [
+ 1848.469970703125,
+ -1164.3343505859375
+ ],
+ "size": [
+ 1004.5999755859375,
+ 1004.3999633789062
+ ],
+ "flags": {},
+ "order": 33,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "images",
+ "type": "IMAGE",
+ "link": 337
+ },
+ {
+ "name": "audio",
+ "shape": 7,
+ "type": "AUDIO",
+ "link": null
+ },
+ {
+ "name": "meta_batch",
+ "shape": 7,
+ "type": "VHS_BatchManager",
+ "link": null
+ },
+ {
+ "name": "vae",
+ "shape": 7,
+ "type": "VAE",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "Filenames",
+ "type": "VHS_FILENAMES",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfyui-videohelpersuite",
+ "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523",
+ "Node name for S&R": "VHS_VideoCombine"
+ },
+ "widgets_values": {
+ "frame_rate": 16,
+ "loop_count": 0,
+ "filename_prefix": "WanVideo_MTV_Crafter",
+ "format": "video/h264-mp4",
+ "pix_fmt": "yuv420p",
+ "crf": 19,
+ "save_metadata": true,
+ "trim_to_audio": false,
+ "pingpong": false,
+ "save_output": false,
+ "videopreview": {
+ "hidden": false,
+ "paused": false,
+ "params": {
+ "filename": "WanVideo_MTV_Crafter_00012.mp4",
+ "subfolder": "",
+ "type": "temp",
+ "format": "video/h264-mp4",
+ "frame_rate": 16,
+ "workflow": "WanVideo_MTV_Crafter_00012.png",
+ "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo_MTV_Crafter_00012.mp4"
+ }
+ }
+ }
+ },
+ {
+ "id": 154,
+ "type": "WanVideoSampler",
+ "pos": [
+ 1072.934814453125,
+ -1081.7371826171875
+ ],
+ "size": [
+ 327.80859375,
+ 1011.80859375
+ ],
+ "flags": {},
+ "order": 30,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "model",
+ "type": "WANVIDEOMODEL",
+ "link": 276
+ },
+ {
+ "name": "image_embeds",
+ "type": "WANVIDIMAGE_EMBEDS",
+ "link": 292
+ },
+ {
+ "name": "text_embeds",
+ "shape": 7,
+ "type": "WANVIDEOTEXTEMBEDS",
+ "link": 287
+ },
+ {
+ "name": "samples",
+ "shape": 7,
+ "type": "LATENT",
+ "link": null
+ },
+ {
+ "name": "feta_args",
+ "shape": 7,
+ "type": "FETAARGS",
+ "link": null
+ },
+ {
+ "name": "context_options",
+ "shape": 7,
+ "type": "WANVIDCONTEXT",
+ "link": null
+ },
+ {
+ "name": "cache_args",
+ "shape": 7,
+ "type": "CACHEARGS",
+ "link": null
+ },
+ {
+ "name": "flowedit_args",
+ "shape": 7,
+ "type": "FLOWEDITARGS",
+ "link": null
+ },
+ {
+ "name": "slg_args",
+ "shape": 7,
+ "type": "SLGARGS",
+ "link": null
+ },
+ {
+ "name": "loop_args",
+ "shape": 7,
+ "type": "LOOPARGS",
+ "link": null
+ },
+ {
+ "name": "experimental_args",
+ "shape": 7,
+ "type": "EXPERIMENTALARGS",
+ "link": null
+ },
+ {
+ "name": "sigmas",
+ "shape": 7,
+ "type": "SIGMAS",
+ "link": null
+ },
+ {
+ "name": "unianimate_poses",
+ "shape": 7,
+ "type": "UNIANIMATE_POSE",
+ "link": null
+ },
+ {
+ "name": "fantasytalking_embeds",
+ "shape": 7,
+ "type": "FANTASYTALKING_EMBEDS",
+ "link": null
+ },
+ {
+ "name": "uni3c_embeds",
+ "shape": 7,
+ "type": "UNI3C_EMBEDS",
+ "link": null
+ },
+ {
+ "name": "multitalk_embeds",
+ "shape": 7,
+ "type": "MULTITALK_EMBEDS",
+ "link": null
+ },
+ {
+ "name": "freeinit_args",
+ "shape": 7,
+ "type": "FREEINITARGS",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "samples",
+ "type": "LATENT",
+ "links": [
+ 285
+ ]
+ },
+ {
+ "name": "denoised_samples",
+ "type": "LATENT",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoSampler"
+ },
+ "widgets_values": [
+ 6,
+ 1,
+ 5,
+ 0,
+ "fixed",
+ true,
+ "dpm++_sde",
+ 0,
+ 1,
+ false,
+ "comfy",
+ 0,
+ -1,
+ false
+ ]
+ },
+ {
+ "id": 197,
+ "type": "WanVideoModelLoader",
+ "pos": [
+ -818.532958984375,
+ -1800.3616943359375
+ ],
+ "size": [
+ 602.063720703125,
+ 294
+ ],
+ "flags": {},
+ "order": 16,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "compile_args",
+ "shape": 7,
+ "type": "WANCOMPILEARGS",
+ "link": null
+ },
+ {
+ "name": "block_swap_args",
+ "shape": 7,
+ "type": "BLOCKSWAPARGS",
+ "link": null
+ },
+ {
+ "name": "lora",
+ "shape": 7,
+ "type": "WANVIDLORA",
+ "link": null
+ },
+ {
+ "name": "vram_management_args",
+ "shape": 7,
+ "type": "VRAM_MANAGEMENTARGS",
+ "link": null
+ },
+ {
+ "name": "extra_model",
+ "shape": 7,
+ "type": "VACEPATH",
+ "link": 353
+ },
+ {
+ "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
+ }
+ ],
+ "outputs": [
+ {
+ "name": "model",
+ "type": "WANVIDEOMODEL",
+ "links": [
+ 351
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "0b1fa14adec48f84adf565dc5616c3a9c7ef95c2",
+ "Node name for S&R": "WanVideoModelLoader"
+ },
+ "widgets_values": [
+ "WanVideo\\fp8_scaled_kj\\I2V\\Wan2_1-I2V-14B-MAGREF_fp8_e4m3fn_scaled_KJ.safetensors",
+ "fp16_fast",
+ "disabled",
+ "offload_device",
+ "sageattn"
+ ],
+ "color": "#223",
+ "bgcolor": "#335"
+ },
+ {
+ "id": 198,
+ "type": "WanVideoExtraModelSelect",
+ "pos": [
+ -1360.404541015625,
+ -1698.273193359375
+ ],
+ "size": [
+ 490.187744140625,
+ 58
+ ],
+ "flags": {},
+ "order": 11,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "extra_model",
+ "type": "VACEPATH",
+ "links": [
+ 353
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "0b1fa14adec48f84adf565dc5616c3a9c7ef95c2",
+ "Node name for S&R": "WanVideoExtraModelSelect"
+ },
+ "widgets_values": [
+ "WanVideo\\MTVCrafter\\Wan2_1_MTV-Crafter_motion_adapter_bf16.safetensors"
+ ]
+ },
+ {
+ "id": 159,
+ "type": "LoadImage",
+ "pos": [
+ -1522.158203125,
+ -1306.342529296875
+ ],
+ "size": [
+ 346.12261962890625,
+ 537.36181640625
+ ],
+ "flags": {},
+ "order": 12,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "IMAGE",
+ "type": "IMAGE",
+ "links": [
+ 280
+ ]
+ },
+ {
+ "name": "MASK",
+ "type": "MASK",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "comfy-core",
+ "ver": "0.3.50",
+ "Node name for S&R": "LoadImage"
+ },
+ "widgets_values": [
+ "example.png",
+ "image"
+ ]
+ },
+ {
+ "id": 163,
+ "type": "WanVideoTextEncodeCached",
+ "pos": [
+ 586.0681762695312,
+ -1102.898193359375
+ ],
+ "size": [
+ 400,
+ 302
+ ],
+ "flags": {},
+ "order": 13,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "extender_args",
+ "shape": 7,
+ "type": "WANVIDEOPROMPTEXTENDER_ARGS",
+ "link": null
+ }
+ ],
+ "outputs": [
+ {
+ "name": "text_embeds",
+ "type": "WANVIDEOTEXTEMBEDS",
+ "links": [
+ 287
+ ]
+ },
+ {
+ "name": "negative_text_embeds",
+ "type": "WANVIDEOTEXTEMBEDS",
+ "links": null
+ },
+ {
+ "name": "positive_prompt",
+ "type": "STRING",
+ "links": null
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoTextEncodeCached"
+ },
+ "widgets_values": [
+ "umt5_xxl_fp16.safetensors",
+ "bf16",
+ "animated character",
+ "bad quality video",
+ "disabled",
+ true,
+ "gpu"
+ ],
+ "color": "#432",
+ "bgcolor": "#653"
+ },
+ {
+ "id": 176,
+ "type": "INTConstant",
+ "pos": [
+ -1923.911376953125,
+ -980.3076171875
+ ],
+ "size": [
+ 210,
+ 58
+ ],
+ "flags": {},
+ "order": 14,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "value",
+ "type": "INT",
+ "links": [
+ 318
+ ]
+ }
+ ],
+ "title": "Width",
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "876a6dd2929d88ec35c09c7cddc6c360f4b27013",
+ "Node name for S&R": "INTConstant"
+ },
+ "widgets_values": [
+ 640
+ ],
+ "color": "#1b4669",
+ "bgcolor": "#29699c"
+ },
+ {
+ "id": 177,
+ "type": "INTConstant",
+ "pos": [
+ -1924.9029541015625,
+ -848.5099487304688
+ ],
+ "size": [
+ 210,
+ 58
+ ],
+ "flags": {},
+ "order": 15,
+ "mode": 0,
+ "inputs": [],
+ "outputs": [
+ {
+ "name": "value",
+ "type": "INT",
+ "links": [
+ 319
+ ]
+ }
+ ],
+ "title": "Height",
+ "properties": {
+ "cnr_id": "comfyui-kjnodes",
+ "ver": "876a6dd2929d88ec35c09c7cddc6c360f4b27013",
+ "Node name for S&R": "INTConstant"
+ },
+ "widgets_values": [
+ 640
+ ],
+ "color": "#1b4669",
+ "bgcolor": "#29699c"
+ },
+ {
+ "id": 167,
+ "type": "WanVideoAddMTVMotion",
+ "pos": [
+ 227.9759979248047,
+ -651.407958984375
+ ],
+ "size": [
+ 289.751953125,
+ 126
+ ],
+ "flags": {},
+ "order": 29,
+ "mode": 0,
+ "inputs": [
+ {
+ "name": "embeds",
+ "type": "WANVIDIMAGE_EMBEDS",
+ "link": 291
+ },
+ {
+ "name": "mtv_crafter_motion",
+ "type": "MTVCRAFTERMOTION",
+ "link": 293
+ }
+ ],
+ "outputs": [
+ {
+ "name": "image_embeds",
+ "type": "WANVIDIMAGE_EMBEDS",
+ "links": [
+ 292
+ ]
+ }
+ ],
+ "properties": {
+ "cnr_id": "ComfyUI-WanVideoWrapper",
+ "ver": "e4f422cf9952b13e8e2147c403ef7b5fed95920a",
+ "Node name for S&R": "WanVideoAddMTVMotion"
+ },
+ "widgets_values": [
+ 1.2,
+ 0,
+ 1
+ ],
+ "color": "#323",
+ "bgcolor": "#535"
+ }
+ ],
+ "links": [
+ [
+ 245,
+ 139,
+ 0,
+ 140,
+ 0,
+ "NLFMODEL"
+ ],
+ [
+ 259,
+ 146,
+ 0,
+ 147,
+ 0,
+ "VQVAE"
+ ],
+ [
+ 260,
+ 140,
+ 0,
+ 147,
+ 1,
+ "NLFPRED"
+ ],
+ [
+ 276,
+ 155,
+ 0,
+ 154,
+ 0,
+ "WANVIDEOMODEL"
+ ],
+ [
+ 277,
+ 156,
+ 0,
+ 155,
+ 1,
+ "WANVIDLORA"
+ ],
+ [
+ 279,
+ 158,
+ 0,
+ 157,
+ 0,
+ "WANVAE"
+ ],
+ [
+ 280,
+ 159,
+ 0,
+ 160,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 284,
+ 158,
+ 0,
+ 161,
+ 0,
+ "WANVAE"
+ ],
+ [
+ 285,
+ 154,
+ 0,
+ 161,
+ 1,
+ "LATENT"
+ ],
+ [
+ 287,
+ 163,
+ 0,
+ 154,
+ 2,
+ "WANVIDEOTEXTEMBEDS"
+ ],
+ [
+ 289,
+ 164,
+ 0,
+ 155,
+ 0,
+ "WANVIDEOMODEL"
+ ],
+ [
+ 290,
+ 165,
+ 0,
+ 164,
+ 1,
+ "BLOCKSWAPARGS"
+ ],
+ [
+ 291,
+ 157,
+ 0,
+ 167,
+ 0,
+ "WANVIDIMAGE_EMBEDS"
+ ],
+ [
+ 292,
+ 167,
+ 0,
+ 154,
+ 1,
+ "WANVIDIMAGE_EMBEDS"
+ ],
+ [
+ 293,
+ 147,
+ 0,
+ 167,
+ 1,
+ "MTVCRAFTERMOTION"
+ ],
+ [
+ 294,
+ 144,
+ 0,
+ 140,
+ 1,
+ "IMAGE"
+ ],
+ [
+ 296,
+ 147,
+ 1,
+ 169,
+ 0,
+ "NLFPRED"
+ ],
+ [
+ 297,
+ 169,
+ 0,
+ 170,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 301,
+ 172,
+ 0,
+ 171,
+ 0,
+ "CLIP_VISION"
+ ],
+ [
+ 307,
+ 160,
+ 1,
+ 157,
+ 8,
+ "INT"
+ ],
+ [
+ 308,
+ 160,
+ 2,
+ 157,
+ 9,
+ "INT"
+ ],
+ [
+ 316,
+ 144,
+ 1,
+ 160,
+ 2,
+ "INT"
+ ],
+ [
+ 317,
+ 144,
+ 2,
+ 160,
+ 3,
+ "INT"
+ ],
+ [
+ 318,
+ 176,
+ 0,
+ 144,
+ 2,
+ "INT"
+ ],
+ [
+ 319,
+ 177,
+ 0,
+ 144,
+ 3,
+ "INT"
+ ],
+ [
+ 321,
+ 138,
+ 1,
+ 157,
+ 10,
+ "INT"
+ ],
+ [
+ 322,
+ 144,
+ 1,
+ 169,
+ 1,
+ "INT"
+ ],
+ [
+ 323,
+ 144,
+ 2,
+ 169,
+ 2,
+ "INT"
+ ],
+ [
+ 336,
+ 169,
+ 0,
+ 174,
+ 1,
+ "IMAGE"
+ ],
+ [
+ 337,
+ 186,
+ 0,
+ 162,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 338,
+ 161,
+ 0,
+ 186,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 339,
+ 174,
+ 0,
+ 186,
+ 1,
+ "IMAGE"
+ ],
+ [
+ 341,
+ 138,
+ 0,
+ 144,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 343,
+ 160,
+ 0,
+ 188,
+ 0,
+ "*"
+ ],
+ [
+ 344,
+ 188,
+ 0,
+ 171,
+ 1,
+ "IMAGE"
+ ],
+ [
+ 346,
+ 188,
+ 0,
+ 174,
+ 0,
+ "IMAGE"
+ ],
+ [
+ 349,
+ 188,
+ 0,
+ 157,
+ 2,
+ "IMAGE"
+ ],
+ [
+ 351,
+ 197,
+ 0,
+ 164,
+ 0,
+ "WANVIDEOMODEL"
+ ],
+ [
+ 353,
+ 198,
+ 0,
+ 197,
+ 4,
+ "VACEPATH"
+ ]
+ ],
+ "groups": [],
+ "config": {},
+ "extra": {
+ "ds": {
+ "scale": 0.6115909044841946,
+ "offset": [
+ 830.4699444221586,
+ 1344.358206307376
+ ]
+ },
+ "frontendVersion": "1.26.3",
+ "node_versions": {
+ "ComfyUI-WanVideoWrapper": "f8f423eceeadf2edcb58fab73701333e83ca733e",
+ "comfy-core": "0.3.26",
+ "ComfyUI_essentials": "76e9d1e4399bd025ce8b12c290753d58f9f53e93",
+ "ComfyUI-KJNodes": "a5bd3c86c8ed6b83c55c2d0e7a59515b15a0137f",
+ "ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8"
+ },
+ "VHS_latentpreview": true,
+ "VHS_latentpreviewrate": 0,
+ "VHS_MetadataImage": true,
+ "VHS_KeepIntermediate": true
+ },
+ "version": 0.4
+}
\ No newline at end of file
diff --git a/fp8_optimization.py b/fp8_optimization.py
index dac0d3d..33d6f60 100644
--- a/fp8_optimization.py
+++ b/fp8_optimization.py
@@ -48,26 +48,6 @@ def apply_lora(weight, lora, step=None):
weight = weight.add(patch_diff, alpha=scale)
return weight
-
-def linear_with_lora_and_scale_forward(cls, input):
- # Handles both scaled and unscaled, with or without LoRA
- has_scale = hasattr(cls, "scale_weight")
- weight = cls.weight.to(input.dtype)
- bias = cls.bias.to(input.dtype) if cls.bias is not None else None
-
- if has_scale:
- scale_weight = cls.scale_weight.to(input.device)
- if weight.numel() < input.numel():
- weight = weight * scale_weight
- else:
- input = input * scale_weight
-
- lora = getattr(cls, "lora", None)
- if lora is not None:
- weight = apply_lora(weight, lora, cls.step).to(input.dtype)
-
- return torch.nn.functional.linear(input, weight, bias)
-
def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=None):
log.info("FP8 matmul enabled")
for name, submodule in module.named_modules():
@@ -81,57 +61,4 @@ def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=
original_forward = submodule.forward
setattr(submodule, "original_forward", original_forward)
setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input))
-
-def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=None, params_to_keep={}):
- log.info("Patching Linear layers...")
- for name, submodule in module.named_modules():
- if not any(keyword in name for keyword in params_to_keep):
- # Set scale_weight if present
- if scale_weight_keys is not None:
- scale_key = f"{name}.scale_weight"
- if scale_key in scale_weight_keys:
- setattr(submodule, "scale_weight", scale_weight_keys[scale_key])
- # Set LoRA if present
- if hasattr(submodule, "lora"):
- #print(f"removing old LoRA in {name}" )
- delattr(submodule, "lora")
- if patches is not None:
- patch_key1 = f"diffusion_model.{name}.weight"
- patch_key_compiled = f"diffusion_model.{name.replace('_orig_mod.', '')}.weight"
- patch = patches.get(patch_key1, []) or patches.get(patch_key_compiled, [])
- if len(patch) != 0:
- lora_diffs = []
- for p in patch:
- lora_obj = p[1]
- if "head" in name:
- continue # For now skip LoRA for head layers
- elif hasattr(lora_obj, "weights"):
- lora_diffs.append(lora_obj.weights)
- elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff":
- lora_diffs.append(lora_obj[1])
- else:
- continue
- lora_strengths = [p[0] for p in patch]
- lora = (lora_diffs, lora_strengths)
- setattr(submodule, "lora", lora)
- #print(f"Added LoRA to {name} with {len(lora_diffs)} diffs and strengths {lora_strengths}")
-
- # Set forward if Linear and has either scale or lora
- if isinstance(submodule, nn.Linear):
- has_scale = hasattr(submodule, "scale_weight")
- has_lora = hasattr(submodule, "lora")
- if not hasattr(submodule, "original_forward"):
- setattr(submodule, "original_forward", submodule.forward)
- if has_scale or has_lora:
- setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_and_scale_forward(m, input))
- setattr(submodule, "step", 0) # Initialize step for LoRA scheduling
-
-def remove_lora_from_module(module):
- unloaded = False
- for name, submodule in module.named_modules():
- if hasattr(submodule, "lora"):
- if not unloaded:
- log.info("Unloading all LoRAs")
- unloaded = True
- delattr(submodule, "lora")
diff --git a/gguf/gguf.py b/gguf/gguf.py
index 44948ad..41f37bd 100644
--- a/gguf/gguf.py
+++ b/gguf/gguf.py
@@ -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=""):
diff --git a/multitalk/multitalk.py b/multitalk/multitalk.py
index cdec103..220bb98 100644
--- a/multitalk/multitalk.py
+++ b/multitalk/multitalk.py
@@ -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
diff --git a/multitalk/nodes.py b/multitalk/nodes.py
index d63c480..9831705 100644
--- a/multitalk/nodes.py
+++ b/multitalk/nodes.py
@@ -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",
}
diff --git a/nodes.py b/nodes.py
index 40dbcc1..37e7ea4 100644
--- a/nodes.py
+++ b/nodes.py
@@ -16,9 +16,10 @@ from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fouri
add_noise_to_reference_video, optimized_scale, setup_radial_attention,
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance)
from .cache_methods.cache_methods import cache_report
+from .nodes_model_loading import load_weights
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"
"
+ 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",
}
diff --git a/nodes_model_loading.py b/nodes_model_loading.py
index 3316e31..a5f6567 100644
--- a/nodes_model_loading.py
+++ b/nodes_model_loading.py
@@ -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",
diff --git a/utils.py b/utils.py
index 85d8637..b00b811 100644
--- a/utils.py
+++ b/utils.py
@@ -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):
diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py
index 1030a30..b923419 100644
--- a/wanvideo/modules/attention.py
+++ b/wanvideo/modules/attention.py
@@ -189,6 +189,8 @@ def attention(
version=fa_version,
)
elif attention_mode == 'sdpa':
+ if not (q.dtype == k.dtype == v.dtype):
+ return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype)).transpose(1, 2).contiguous()
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous()
elif attention_mode == 'sageattn_3':
return sageattn_blackwell(
diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py
index 8e4a003..8bb23ac 100644
--- a/wanvideo/modules/model.py
+++ b/wanvideo/modules/model.py
@@ -28,6 +28,9 @@ from ...cache_methods.cache_methods import TeaCacheState, MagCacheState, EasyCac
from ...multitalk.multitalk import get_attn_map_with_target
from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
+from ...MTV.mtv import apply_rotary_emb
+
+
__all__ = ['WanModel']
from comfy import model_management as mm
@@ -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]
diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py
index dc63457..407e107 100644
--- a/wanvideo/schedulers/__init__.py
+++ b/wanvideo/schedulers/__init__.py
@@ -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
\ No newline at end of file
+ return sample_scheduler, timesteps
\ No newline at end of file