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"Sigmas Plot" + 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