Author SHA1 Message Date
kijai 2e7c00a03b initial syncammaster 2025-04-17 01:21:47 +03:00
213 changed files with 13440 additions and 599436 deletions
+1 -3
View File
@@ -9,6 +9,4 @@ logs/
.idea .idea
tools/ tools/
.vscode/ .vscode/
convert_* convert_*
*.pt
*.pth
-42
View File
@@ -1,42 +0,0 @@
# Copyright (c) 2024-2025 Bytedance Ltd. and/or its affiliates
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Dict, List, Optional, Tuple, Union
import numpy as np
import torch
def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], quant_multi: int = 8, **kwargs):
# tracks: shape [t, h, w, 3] => samples align with 24 fps, model trained with 16 fps.
# frame_size: tuple (W, H)
tracks = torch.from_numpy(tracks_np).float()
if tracks.shape[1] == 121:
tracks = torch.permute(tracks, (1, 0, 2, 3))
tracks, visibles = tracks[..., :2], tracks[..., 2:3]
short_edge = min(*frame_size)
tracks = tracks - torch.tensor([*frame_size]).type_as(tracks) / 2
tracks = tracks / short_edge * 2
visibles = visibles * 2 - 1
trange = torch.linspace(-1, 1, tracks.shape[0]).view(-1, 1, 1, 1).expand(*visibles.shape)
out_ = torch.cat([trange, tracks, visibles], dim=-1).view(121, -1, 4)
out_0 = out_[:1]
out_l = out_[1:] # 121 => 120 | 1
out_l = torch.repeat_interleave(out_l, 2, dim=0)[1::3] # 120 => 240 => 80
return torch.cat([out_0, out_l], dim=0)
-142
View File
@@ -1,142 +0,0 @@
# Copyright (c) 2024-2025 Bytedance Ltd. and/or its affiliates
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import List, Optional, Tuple, Union
import torch
# Refer to https://github.com/Angtian/VoGE/blob/main/VoGE/Utils.py
def ind_sel(target: torch.Tensor, ind: torch.Tensor, dim: int = 1):
"""
:param target: [... (can be k or 1), n > M, ...]
:param ind: [... (k), M]
:param dim: dim to apply index on
:return: sel_target [... (k), M, ...]
"""
assert (
len(ind.shape) > dim
), "Index must have the target dim, but get dim: %d, ind shape: %s" % (dim, str(ind.shape))
target = target.expand(
*tuple(
[ind.shape[k] if target.shape[k] == 1 else -1 for k in range(dim)]
+ [
-1,
]
* (len(target.shape) - dim)
)
)
ind_pad = ind
if len(target.shape) > dim + 1:
for _ in range(len(target.shape) - (dim + 1)):
ind_pad = ind_pad.unsqueeze(-1)
ind_pad = ind_pad.expand(*(-1,) * (dim + 1), *target.shape[(dim + 1) : :])
return torch.gather(target, dim=dim, index=ind_pad)
def merge_final(vert_attr: torch.Tensor, weight: torch.Tensor, vert_assign: torch.Tensor):
"""
:param vert_attr: [n, d] or [b, n, d] color or feature of each vertex
:param weight: [b(optional), w, h, M] weight of selected vertices
:param vert_assign: [b(optional), w, h, M] selective index
:return:
"""
target_dim = len(vert_assign.shape) - 1
if len(vert_attr.shape) == 2:
assert vert_attr.shape[0] > vert_assign.max()
# [n, d] ind: [b(optional), w, h, M]-> [b(optional), w, h, M, d]
# sel_attr = ind_sel(
# vert_attr[(None,) * target_dim], vert_assign.type(torch.long), dim=target_dim
# )
new_shape = [1] * target_dim + list(vert_attr.shape)
tensor = vert_attr.reshape(new_shape)
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
else:
assert vert_attr.shape[1] > vert_assign.max()
#sel_attr = ind_sel(
# vert_attr[:, *(None,) * (target_dim - 1)], vert_assign.type(torch.long), dim=target_dim
#)
new_shape = [vert_attr.shape[0]] + [1] * (target_dim - 1) + list(vert_attr.shape[1:])
tensor = vert_attr.reshape(new_shape)
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
# [b(optional), w, h, M]
final_attr = torch.sum(sel_attr * weight.unsqueeze(-1), dim=-2)
return final_attr
def patch_motion(
tracks: torch.FloatTensor, # (B, T, N, 4)
vid: torch.FloatTensor, # (C, T, H, W)
temperature: float = 220.0,
vae_divide: tuple = (4, 16),
topk: int = 2,
):
with torch.no_grad():
_, T, H, W = vid.shape
N = tracks.shape[2]
_, tracks, visible = torch.split(
tracks, [1, 2, 1], dim=-1
) # (B, T, N, 2) | (B, T, N, 1)
tracks_n = tracks / torch.tensor([W / min(H, W), H / min(H, W)], device=tracks.device)
tracks_n = tracks_n.clamp(-1, 1)
visible = visible.clamp(0, 1)
xx = torch.linspace(-W / min(H, W), W / min(H, W), W)
yy = torch.linspace(-H / min(H, W), H / min(H, W), H)
grid = torch.stack(torch.meshgrid(yy, xx, indexing="ij")[::-1], dim=-1).to(
tracks.device
)
tracks_pad = tracks[:, 1:]
visible_pad = visible[:, 1:]
visible_align = visible_pad.view(T - 1, 4, *visible_pad.shape[2:]).sum(1)
tracks_align = (tracks_pad * visible_pad).view(T - 1, 4, *tracks_pad.shape[2:]).sum(
1
) / (visible_align + 1e-5)
dist_ = (
(tracks_align[:, None, None] - grid[None, :, :, None]).pow(2).sum(-1)
) # T, H, W, N
weight = torch.exp(-dist_ * temperature) * visible_align.clamp(0, 1).view(
T - 1, 1, 1, N
)
vert_weight, vert_index = torch.topk(
weight, k=min(topk, weight.shape[-1]), dim=-1
)
grid_mode = "bilinear"
point_feature = torch.nn.functional.grid_sample(
vid[vae_divide[0]:].permute(1, 0, 2, 3)[:1],
tracks_n[:, :1].type(vid.dtype),
mode=grid_mode,
padding_mode="zeros",
align_corners=False,
)
point_feature = point_feature.squeeze(0).squeeze(1).permute(1, 0) # N, C=16
out_feature = merge_final(point_feature, vert_weight, vert_index).permute(3, 0, 1, 2) # T - 1, H, W, C => C, T - 1, H, W
out_weight = vert_weight.sum(-1) # T - 1, H, W
# out feature -> already soft weighted
mix_feature = out_feature + vid[vae_divide[0]:, 1:] * (1 - out_weight.clamp(0, 1))
out_feature_full = torch.cat([vid[vae_divide[0]:, :1], mix_feature], dim=1) # C, T, H, W
out_mask_full = torch.cat([torch.ones_like(out_weight[:1]), out_weight], dim=0) # T, H, W
return torch.cat([out_mask_full[None].expand(vae_divide[0], -1, -1, -1), out_feature_full], dim=0)
-329
View File
@@ -1,329 +0,0 @@
import json
from .motion import process_tracks
import numpy as np
from typing import List, Tuple
import torch
FIXED_LENGTH = 121
def pad_pts(tr):
"""Convert list of {x,y} to (FIXED_LENGTH,1,3) array, padding/truncating."""
pts = np.array([[p['x'], p['y'], 1] for p in tr], dtype=np.float32)
n = pts.shape[0]
if n < FIXED_LENGTH:
pad = np.zeros((FIXED_LENGTH - n, 3), dtype=np.float32)
pts = np.vstack((pts, pad))
else:
pts = pts[:FIXED_LENGTH]
return pts.reshape(FIXED_LENGTH, 1, 3)
def age_to_bgr(ratio: float) -> Tuple[int,int,int]:
"""
Map ratio∈[0,1] through: 0→blue, 1/3→green, 2/3→yellow, 1→red.
Returns (B,G,R) for OpenCV.
"""
if ratio <= 1/3:
# blue→green
t = ratio / (1/3)
b = int(255 * (1 - t))
g = int(255 * t)
r = 0
elif ratio <= 2/3:
# green→yellow
t = (ratio - 1/3) / (1/3)
b = 0
g = 255
r = int(255 * t)
else:
# yellow→red
t = (ratio - 2/3) / (1/3)
b = 0
g = int(255 * (1 - t))
r = 255
return (r, g, b)
def paint_point_track(
frames: np.ndarray,
point_tracks: np.ndarray,
visibles: np.ndarray,
min_radius: int = 1,
max_radius: int = 6,
max_retain: int = 50
) -> np.ndarray:
"""
Draws every past point of each track on each frame, with radius and color
interpolated by the point's age (old→small to new→large).
Args:
frames: [F, H, W, 3] uint8 RGB
point_tracks:[N, F, 2] float32 – (x,y) in pixel coords
visibles: [N, F] bool – visibility mask
min_radius: radius for the very first point (oldest)
max_radius: radius for the current point (newest)
Returns:
video: [F, H, W, 3] uint8 RGB
"""
import cv2
num_points, num_frames = point_tracks.shape[:2]
H, W = frames.shape[1:3]
video = frames.copy()
for t in range(num_frames):
# start from the original frame
frame = video[t].copy()
for i in range(num_points):
# draw every past step τ = 0..t
for τ in range(t + 1):
if not visibles[i, τ]:
continue
if t - τ > max_retain:
continue
# sub-pixel offset + clamp
x, y = point_tracks[i, τ] + 0.5
xi = int(np.clip(x, 0, W - 1))
yi = int(np.clip(y, 0, H - 1))
# age‐ratio in [0,1]
if num_frames > 1:
ratio = 1 - float(t - τ) / max_retain
else:
ratio = 1.0
# interpolated radius
radius = int(round(min_radius + (max_radius - min_radius) * ratio))
# OpenCV draws in BGR order:
color_rgb = age_to_bgr(ratio)
# filled circle
cv2.circle(frame, (xi, yi), radius, color_rgb, thickness=-1)
video[t] = frame
return video
def parse_json_tracks(tracks):
tracks_data = []
try:
# If tracks is a string, try to parse it as JSON
if isinstance(tracks, str):
parsed = json.loads(tracks.replace("'", '"'))
tracks_data.extend(parsed)
else:
# If tracks is a list of strings, parse each one
for track_str in tracks:
parsed = json.loads(track_str.replace("'", '"'))
tracks_data.append(parsed)
# Check if we have a single track (dict with x,y) or a list of tracks
if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]:
# Single track detected, wrap it in a list
tracks_data = [tracks_data]
elif tracks_data and isinstance(tracks_data[0], list) and tracks_data[0] and isinstance(tracks_data[0][0], dict) and 'x' in tracks_data[0][0]:
# Already a list of tracks, nothing to do
pass
else:
# Unexpected format
print(f"Warning: Unexpected track format: {type(tracks_data[0])}")
except json.JSONDecodeError as e:
print(f"Error parsing tracks JSON: {e}")
tracks_data = []
return tracks_data
class WanVideoATITracks:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("WANVIDEOMODEL", ),
"tracks": ("STRING",),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
"temperature": ("FLOAT", {"default": 220.0, "min": 0.0, "max": 1000.0, "step": 0.1}),
"topk": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply ATI"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply ATI"}),
},
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "patchmodel"
CATEGORY = "WanVideoWrapper"
def patchmodel(self, model, tracks, width, height, temperature, topk, start_percent, end_percent):
tracks_data = parse_json_tracks(tracks)
arrs = []
for track in tracks_data:
pts = pad_pts(track)
arrs.append(pts)
tracks_np = np.stack(arrs, axis=0)
processed_tracks = process_tracks(tracks_np, (width, height))
patcher = model.clone()
patcher.model_options["transformer_options"]["ati_tracks"] = processed_tracks.unsqueeze(0)
patcher.model_options["transformer_options"]["ati_temperature"] = temperature
patcher.model_options["transformer_options"]["ati_topk"] = topk
patcher.model_options["transformer_options"]["ati_start_percent"] = start_percent
patcher.model_options["transformer_options"]["ati_end_percent"] = end_percent
return (patcher,)
class WanVideoATITracksVisualize:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"images": ("IMAGE",),
"tracks": ("STRING",),
"min_radius": ("INT", {"default": 1, "min": 0, "max": 100, "step": 1, "tooltip": "radius for the very first point (oldest)"}),
"max_radius": ("INT", {"default": 6, "min": 0, "max": 100, "step": 1, "tooltip": "radius for the current point (newest)"}),
"max_retain": ("INT", {"default": 50, "min": 0, "max": 100, "step": 1, "tooltip": "Maximum number of points to retain"}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "patchmodel"
CATEGORY = "WanVideoWrapper"
def patchmodel(self, images, tracks, min_radius, max_radius, max_retain):
tracks_data = parse_json_tracks(tracks)
arrs = []
for track in tracks_data:
pts = pad_pts(track)
arrs.append(pts)
tracks_np = np.stack(arrs, axis=0)
track = np.repeat(tracks_np, 2, axis=1)[:, ::3]
points = track[:, :, 0, :2].astype(np.float32)
visibles = track[:, :, 0, 2].astype(np.float32)
if images.shape[0] < points.shape[1]:
repeat_count = (points.shape[1] + images.shape[0] - 1) // images.shape[0]
images = images.repeat(repeat_count, 1, 1, 1)
images = images[:points.shape[1]]
elif images.shape[0] > points.shape[1]:
images = images[:points.shape[1]]
video_viz = paint_point_track(images.cpu().numpy(), points, visibles, min_radius, max_radius, max_retain)
video_viz = torch.from_numpy(video_viz).float()
return (video_viz,)
from comfy import utils
import types
from .motion_patch import patch_motion
class WanConcatCondPatch:
def __init__(self, tracks, temperature, topk):
self.tracks = tracks
self.temperature = temperature
self.topk = topk
def __get__(self, obj, objtype=None):
# Create bound method with stored parameters
def wrapped_concat_cond(self_module, *args, **kwargs):
return modified_concat_cond(self_module, self.tracks, self.temperature, self.topk, *args, **kwargs)
return types.MethodType(wrapped_concat_cond, obj)
def modified_concat_cond(self, tracks, temperature, topk, **kwargs):
noise = kwargs.get("noise", None)
extra_channels = self.diffusion_model.patch_embedding.weight.shape[1] - noise.shape[1]
if extra_channels == 0:
return None
image = kwargs.get("concat_latent_image", None)
device = kwargs["device"]
if image is None:
shape_image = list(noise.shape)
shape_image[1] = extra_channels
image = torch.zeros(shape_image, dtype=noise.dtype, layout=noise.layout, device=noise.device)
else:
image = utils.common_upscale(image.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center")
for i in range(0, image.shape[1], 16):
image[:, i: i + 16] = self.process_latent_in(image[:, i: i + 16])
image = utils.resize_to_batch_size(image, noise.shape[0])
if not self.image_to_video or extra_channels == image.shape[1]:
return image
if image.shape[1] > (extra_channels - 4):
image = image[:, :(extra_channels - 4)]
mask = kwargs.get("concat_mask", kwargs.get("denoise_mask", None))
if mask is None:
mask = torch.zeros_like(noise)[:, :4]
else:
if mask.shape[1] != 4:
mask = torch.mean(mask, dim=1, keepdim=True)
mask = 1.0 - mask
mask = utils.common_upscale(mask.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center")
if mask.shape[-3] < noise.shape[-3]:
mask = torch.nn.functional.pad(mask, (0, 0, 0, 0, 0, noise.shape[-3] - mask.shape[-3]), mode='constant', value=0)
if mask.shape[1] == 1:
mask = mask.repeat(1, 4, 1, 1, 1)
mask = utils.resize_to_batch_size(mask, noise.shape[0])
image_cond = torch.cat((mask, image), dim=1)
image_cond_ati = patch_motion(tracks.to(image_cond.device, image_cond.dtype), image_cond[0],
temperature=temperature, topk=topk)
return image_cond_ati.unsqueeze(0)
class WanVideoATI_comfy:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": ("MODEL", ),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
"tracks": ("STRING",),
"temperature": ("FLOAT", {"default": 220.0, "min": 0.0, "max": 1000.0, "step": 0.1}),
"topk": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}),
},
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "patchcond"
CATEGORY = "WanVideoWrapper"
def patchcond(self, model, tracks, width, height, temperature, topk):
tracks_data = parse_json_tracks(tracks)
arrs = []
for track in tracks_data:
pts = pad_pts(track)
arrs.append(pts)
tracks_np = np.stack(arrs, axis=0)
processed_tracks = process_tracks(tracks_np, (width, height))
model_clone = model.clone()
model_clone.add_object_patch(
"concat_cond",
WanConcatCondPatch(
processed_tracks.unsqueeze(0), temperature, topk
).__get__(model.model, model.model.__class__)
)
return (model_clone,)
NODE_CLASS_MAPPINGS = {
"WanVideoATITracks": WanVideoATITracks,
"WanVideoATITracksVisualize": WanVideoATITracksVisualize,
"WanVideoATI_comfy": WanVideoATI_comfy,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoATITracks": "WanVideo ATI Tracks",
"WanVideoATITracksVisualize": "WanVideo ATI Tracks Visualize",
"WanVideoATI_comfy": "WanVideo ATI Comfy",
}
-157
View File
@@ -1,157 +0,0 @@
from einops import rearrange
import torch
import torch.nn as nn
import torch.nn.functional as F
CACHE_T = 2
class RMS_norm(nn.Module):
def __init__(self, dim, channel_first=True, images=True, bias=False):
super().__init__()
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
self.channel_first = channel_first
self.scale = dim**0.5
self.gamma = nn.Parameter(torch.ones(shape))
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.
def forward(self, x):
return F.normalize(
x, dim=(1 if self.channel_first else
-1)) * self.scale * self.gamma + self.bias
class CausalConv3d(nn.Conv3d):
"""
Causal 3d convolusion.
"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._padding = (self.padding[2], self.padding[2], self.padding[1],
self.padding[1], 2 * self.padding[0], 0)
self.padding = (0, 0, 0)
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding, mode='replicate')
return super().forward(x)
class PixelShuffle3d(nn.Module):
def __init__(self, ff, hh, ww):
super().__init__()
self.ff = ff
self.hh = hh
self.ww = ww
def forward(self, x):
# x: (B, C, F, H, W)
return rearrange(x,
'b c (f ff) (h hh) (w ww) -> b (c ff hh ww) f h w',
ff=self.ff, hh=self.hh, ww=self.ww)
class Buffer_LQ4x_Proj(nn.Module):
def __init__(self, in_dim, out_dim, layer_num=30):
super().__init__()
self.ff = 1
self.hh = 16
self.ww = 16
self.hidden_dim1 = 2048
self.hidden_dim2 = 3072
self.layer_num = layer_num
self.pixel_shuffle = PixelShuffle3d(self.ff, self.hh, self.ww)
self.conv1 = CausalConv3d(in_dim*self.ff*self.hh*self.ww, self.hidden_dim1, (4, 3, 3), stride=(2, 1, 1), padding=(1, 1, 1)) # f -> f/2 h -> h w -> w
self.norm1 = RMS_norm(self.hidden_dim1, images=False)
self.act1 = nn.SiLU()
self.conv2 = CausalConv3d(self.hidden_dim1, self.hidden_dim2, (4, 3, 3), stride=(2, 1, 1), padding=(1, 1, 1)) # f -> f/2 h -> h w -> w
self.norm2 = RMS_norm(self.hidden_dim2, images=False)
self.act2 = nn.SiLU()
self.linear_layers = nn.ModuleList([nn.Linear(self.hidden_dim2, out_dim) for _ in range(layer_num)])
self.clip_idx = 0
def forward(self, video):
self.clear_cache()
# x: (B, C, F, H, W)
t = video.shape[2]
iter_ = 1 + (t - 1) // 4
first_frame = video[:, :, :1, :, :].repeat(1, 1, 3, 1, 1)
video = torch.cat([first_frame, video], dim=2)
out_x = []
for i in range(iter_):
x = self.pixel_shuffle(video[:,:,i*4:(i+1)*4,:,:])
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv1'] = cache1_x
x = self.conv1(x, self.cache['conv1'])
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv2'] = cache2_x
if i == 0:
continue
x = self.conv2(x, self.cache['conv2'])
x = self.norm2(x)
x = self.act2(x)
out_x.append(x)
out_x = torch.cat(out_x, dim = 2)
out_x = rearrange(out_x, 'b c f h w -> b (f h w) c')
outputs = []
for i in range(self.layer_num):
outputs.append(self.linear_layers[i](out_x))
self.clear_cache()
return outputs
def clear_cache(self):
self.cache = {}
self.cache['conv1'] = None
self.cache['conv2'] = None
self.clip_idx = 0
def stream_forward(self, video_clip):
if self.clip_idx == 0:
# self.clear_cache()
first_frame = video_clip[:, :, :1, :, :].repeat(1, 1, 3, 1, 1)
video_clip = torch.cat([first_frame, video_clip], dim=2)
x = self.pixel_shuffle(video_clip)
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv1'] = cache1_x
x = self.conv1(x, self.cache['conv1'])
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv2'] = cache2_x
self.clip_idx += 1
return None
else:
x = self.pixel_shuffle(video_clip)
cache1_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv1'] = cache1_x
x = self.conv1(x, self.cache['conv1'])
x = self.norm1(x)
x = self.act1(x)
cache2_x = x[:, :, -CACHE_T:, :, :].clone()
self.cache['conv2'] = cache2_x
x = self.conv2(x, self.cache['conv2'])
x = self.norm2(x)
x = self.act2(x)
out_x = rearrange(x, 'b c f h w -> b (f h w) c')
outputs = []
for i in range(self.layer_num):
outputs.append(self.linear_layers[i](out_x))
self.clip_idx += 1
return outputs
-261
View File
@@ -1,261 +0,0 @@
"""
Tiny AutoEncoder for Hunyuan Video (Decoder-only, pruned)
- Encoder removed
- Transplant/widening helpers removed
- Deepening (IdentityConv2d+ReLU) is now built into the decoder structure itself
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm.auto import tqdm
from collections import namedtuple
from einops import rearrange
import torch.nn.init as init
DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
# ----------------------------
# Utility / building blocks
# ----------------------------
class IdentityConv2d(nn.Conv2d):
"""Same-shape Conv2d initialized to identity (Dirac)."""
def __init__(self, C, kernel_size=3, bias=False):
pad = kernel_size // 2
super().__init__(C, C, kernel_size, padding=pad, bias=bias)
with torch.no_grad():
init.dirac_(self.weight)
if self.bias is not None:
self.bias.zero_()
def conv(n_in, n_out, **kwargs):
return nn.Conv2d(n_in, n_out, 3, padding=1, **kwargs)
class Clamp(nn.Module):
def forward(self, x):
return torch.tanh(x / 3) * 3
class MemBlock(nn.Module):
def __init__(self, n_in, n_out):
super().__init__()
self.conv = nn.Sequential(
conv(n_in * 2, n_out), nn.ReLU(inplace=True),
conv(n_out, n_out), nn.ReLU(inplace=True),
conv(n_out, n_out)
)
self.skip = nn.Conv2d(n_in, n_out, 1, bias=False) if n_in != n_out else nn.Identity()
self.act = nn.ReLU(inplace=True)
def forward(self, x, past):
return self.act(self.conv(torch.cat([x, past], 1)) + self.skip(x))
class TPool(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f*stride, n_f, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
return self.conv(x.reshape(-1, self.stride * C, H, W))
class TGrow(nn.Module):
def __init__(self, n_f, stride):
super().__init__()
self.stride = stride
self.conv = nn.Conv2d(n_f, n_f*stride, 1, bias=False)
def forward(self, x):
_NT, C, H, W = x.shape
x = self.conv(x)
return x.reshape(-1, C, H, W)
class PixelShuffle3d(nn.Module):
def __init__(self, ff, hh, ww):
super().__init__()
self.ff = ff
self.hh = hh
self.ww = ww
def forward(self, x):
# x: (B, C, F, H, W)
B, C, F, H, W = x.shape
if F % self.ff != 0:
first_frame = x[:, :, 0:1, :, :].repeat(1, 1, self.ff - F % self.ff, 1, 1)
x = torch.cat([first_frame, x], dim=2)
return rearrange(
x,
'b c (f ff) (h hh) (w ww) -> b (c ff hh ww) f h w',
ff=self.ff, hh=self.hh, ww=self.ww
).transpose(1, 2)
# ----------------------------
# Generic NTCHW graph executor (kept; used by decoder)
# ----------------------------
def apply_model_with_memblocks(model, x, parallel, show_progress_bar, mem=None):
"""
Apply a sequential model with memblocks to the given input.
Args:
- model: nn.Sequential of blocks to apply
- x: input data, of dimensions NTCHW
- parallel: if True, parallelize over timesteps (fast but uses O(T) memory)
if False, each timestep will be processed sequentially (slow but uses O(1) memory)
- show_progress_bar: if True, enables tqdm progressbar display
Returns NTCHW tensor of output data.
"""
assert x.ndim == 5, f"TAEHV operates on NTCHW tensors, but got {x.ndim}-dim tensor"
N, T, C, H, W = x.shape
if parallel:
x = x.reshape(N*T, C, H, W)
for b in tqdm(model, disable=not show_progress_bar):
if isinstance(b, MemBlock):
NT, C, H, W = x.shape
T = NT // N
_x = x.reshape(N, T, C, H, W)
mem = F.pad(_x, (0,0,0,0,0,0,1,0), value=0)[:,:T].reshape(x.shape)
x = b(x, mem)
else:
x = b(x)
NT, C, H, W = x.shape
T = NT // N
x = x.view(N, T, C, H, W)
else:
out = []
work_queue = [TWorkItem(xt, 0) for t, xt in enumerate(x.reshape(N, T * C, H, W).chunk(T, dim=1))]
progress_bar = tqdm(range(T), disable=not show_progress_bar)
while work_queue:
xt, i = work_queue.pop(0)
if i == 0:
progress_bar.update(1)
if i == len(model):
out.append(xt)
else:
b = model[i]
if isinstance(b, MemBlock):
if mem[i] is None:
xt_new = b(xt, xt * 0)
mem[i] = xt
else:
xt_new = b(xt, mem[i])
mem[i].copy_(xt)
work_queue.insert(0, TWorkItem(xt_new, i+1))
elif isinstance(b, TPool):
if mem[i] is None:
mem[i] = []
mem[i].append(xt)
if len(mem[i]) > b.stride:
raise ValueError("TPool internal state invalid.")
elif len(mem[i]) == b.stride:
N_, C_, H_, W_ = xt.shape
xt = b(torch.cat(mem[i], 1).view(N_*b.stride, C_, H_, W_))
mem[i] = []
work_queue.insert(0, TWorkItem(xt, i+1))
elif isinstance(b, TGrow):
xt = b(xt)
NT, C_, H_, W_ = xt.shape
for xt_next in reversed(xt.view(N, b.stride*C_, H_, W_).chunk(b.stride, 1)):
work_queue.insert(0, TWorkItem(xt_next, i+1))
else:
xt = b(xt)
work_queue.insert(0, TWorkItem(xt, i+1))
progress_bar.close()
x = torch.stack(out, 1)
return x, mem
# ----------------------------
# Decoder-only TAEHV
# ----------------------------
class TAEHV(nn.Module):
image_channels = 3
def __init__(
self,
decoder_time_upscale=(True, True),
decoder_space_upscale=(True, True, True),
channels = [256, 128, 64, 64],
latent_channels = 16,
dtype=torch.float32
):
"""Initialize TAEHV (decoder-only) with built-in deepening after every ReLU.
Deepening config: how_many_each=1, k=3 (fixed as requested).
"""
super().__init__()
self.dtype = dtype
self.latent_channels = latent_channels
n_f = channels
self.frames_to_trim = 2**sum(decoder_time_upscale) - 1
# Build the decoder "skeleton"
base_decoder = nn.Sequential(
Clamp(), conv(self.latent_channels, n_f[0]), nn.ReLU(inplace=True),
MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]), MemBlock(n_f[0], n_f[0]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[0] else 1),
TGrow(n_f[0], 1),
conv(n_f[0], n_f[1], bias=False),
MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]), MemBlock(n_f[1], n_f[1]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[1] else 1),
TGrow(n_f[1], 2 if decoder_time_upscale[0] else 1),
conv(n_f[1], n_f[2], bias=False),
MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]), MemBlock(n_f[2], n_f[2]),
nn.Upsample(scale_factor=2 if decoder_space_upscale[2] else 1),
TGrow(n_f[2], 2 if decoder_time_upscale[1] else 1),
conv(n_f[2], n_f[3], bias=False),
nn.ReLU(inplace=True), conv(n_f[3], TAEHV.image_channels),
)
# Inline deepening: insert (IdentityConv2d(k=3) + ReLU) after every ReLU
self.decoder = self._apply_identity_deepen(base_decoder, how_many_each=1, k=3)
self.pixel_shuffle = PixelShuffle3d(4, 8, 8)
# Initialize decoder mem state
self.clean_mem()
@staticmethod
def _apply_identity_deepen(decoder: nn.Sequential, how_many_each=1, k=3) -> nn.Sequential:
"""Return a new Sequential where every nn.ReLU is followed by how_many_each*(IdentityConv2d(k)+ReLU)."""
new_layers = []
for b in decoder:
new_layers.append(b)
if isinstance(b, nn.ReLU):
# Deduce channel count from preceding layer
C = None
if len(new_layers) >= 2 and isinstance(new_layers[-2], nn.Conv2d):
C = new_layers[-2].out_channels
elif len(new_layers) >= 2 and isinstance(new_layers[-2], MemBlock):
C = new_layers[-2].conv[-1].out_channels
if C is not None:
for _ in range(how_many_each):
new_layers.append(IdentityConv2d(C, kernel_size=k, bias=False))
new_layers.append(nn.ReLU(inplace=True))
return nn.Sequential(*new_layers)
def decode_video(self, x, parallel=False, show_progress_bar=False, cond=None):
"""Decode a sequence of frames from latents.
x: NTCHW latent tensor; returns NTCHW RGB in ~[0, 1].
"""
trim_flag = self.mem[-8] is None # keeps original relative check
if cond is not None:
shuffled = self.pixel_shuffle(cond.to(x))
x = torch.cat([shuffled[:, :x.shape[1]], x], dim=2)
x, self.mem = apply_model_with_memblocks(self.decoder, x, parallel, show_progress_bar, mem=self.mem)
self.clean_mem()
if trim_flag:
return x[:, self.frames_to_trim:]
return x
def clean_mem(self):
self.mem = [None] * len(self.decoder)
def build_tcdecoder(new_channels = [512, 256, 128, 128], device="cuda", dtype=torch.bfloat16, new_latent_channels=None):
big = TAEHV(channels=new_channels, latent_channels=new_latent_channels, dtype=dtype).to(device).to(dtype)
return big
-71
View File
@@ -1,71 +0,0 @@
import folder_paths
import torch
from comfy.utils import load_torch_file
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddFlashVSRInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"images": ("IMAGE", {"tooltip": "Low-res video frames to enhance"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01, "tooltip": "Strength to apply the FlashVSR latent"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, images, strength):
updated = dict(embeds)
updated["flashvsr_LQ_images"] = images
updated["flashvsr_strength"] = strength
return (updated,)
class WanVideoFlashVSRDecoderLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_name": (folder_paths.get_filename_list("vae"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae'"}),
},
"optional": {
"precision": (["fp16", "fp32", "bf16"],
{"default": "bf16"}
),
}
}
RETURN_TYPES = ("WANVAE",)
RETURN_NAMES = ("vae", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'"
def loadmodel(self, model_name, precision):
from .TCDecoder import build_tcdecoder
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path("vae", model_name)
sd = load_torch_file(model_path, safe_load=True)
TCDecoder = build_tcdecoder(new_channels=[512, 256, 128, 128], new_latent_channels=16+768, dtype=dtype)
TCDecoder.load_state_dict(sd, strict=True)
TCDecoder.to(dtype)
return (TCDecoder,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddFlashVSRInput": WanVideoAddFlashVSRInput,
"WanVideoFlashVSRDecoderLoader": WanVideoFlashVSRDecoderLoader,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddFlashVSRInput": "WanVideo Add FlashVSR Input",
"WanVideoFlashVSRDecoderLoader": "WanVideo FlashVSR Decoder Loader",
}
-87
View File
@@ -1,87 +0,0 @@
import torch
from einops import rearrange
from torch import nn
from einops import rearrange
class WanRMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
r"""
Args:
x(Tensor): Shape [B, L, C]
"""
return self._norm(x.to(self.weight.dtype)) * self.weight
def _norm(self, x):
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)).to(x.dtype)
class DummyAdapterLayer(nn.Module):
def __init__(self, layer):
super().__init__()
self.layer = layer
def forward(self, *args, **kwargs):
return self.layer(*args, **kwargs)
class AudioProjModel(nn.Module):
def __init__(
self,
seq_len=5,
blocks=13, # add a new parameter blocks
channels=768, # add a new parameter channels
intermediate_dim=512,
output_dim=1536,
context_tokens=16,
):
super().__init__()
self.seq_len = seq_len
self.blocks = blocks
self.channels = channels
self.input_dim = seq_len * blocks * channels # update input_dim to be the product of blocks and channels.
self.intermediate_dim = intermediate_dim
self.context_tokens = context_tokens
self.output_dim = output_dim
# define multiple linear layers
self.audio_proj_glob_1 = DummyAdapterLayer(nn.Linear(self.input_dim, intermediate_dim))
self.audio_proj_glob_2 = DummyAdapterLayer(nn.Linear(intermediate_dim, intermediate_dim))
self.audio_proj_glob_3 = DummyAdapterLayer(nn.Linear(intermediate_dim, context_tokens * output_dim))
self.audio_proj_glob_norm = DummyAdapterLayer(nn.LayerNorm(output_dim))
self.initialize_weights()
def initialize_weights(self):
# Initialize transformer layers:
def _basic_init(module):
if isinstance(module, nn.Linear):
torch.nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.constant_(module.bias, 0)
self.apply(_basic_init)
def forward(self, audio_embeds):
video_length = audio_embeds.shape[1]
audio_embeds = rearrange(audio_embeds, "bz f w b c -> (bz f) w b c")
batch_size, window_size, blocks, channels = audio_embeds.shape
audio_embeds = audio_embeds.view(batch_size, window_size * blocks * channels)
audio_embeds = torch.relu(self.audio_proj_glob_1(audio_embeds))
audio_embeds = torch.relu(self.audio_proj_glob_2(audio_embeds))
context_tokens = self.audio_proj_glob_3(audio_embeds).reshape(batch_size, self.context_tokens, self.output_dim)
context_tokens = self.audio_proj_glob_norm(context_tokens.to(self.audio_proj_glob_norm.layer.weight.dtype)).to(audio_embeds.dtype)
context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length)
return context_tokens
-287
View File
@@ -1,287 +0,0 @@
import folder_paths
import torch
import torch.nn.functional as F
import os
import json
import torchaudio
from comfy.utils import load_torch_file, common_upscale
import comfy.model_management as mm
from accelerate import init_empty_weights
from ..utils import set_module_tensor_to_device, log
from ..nodes import WanVideoEncodeLatentBatch
script_directory = os.path.dirname(os.path.abspath(__file__))
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
def linear_interpolation_fps(features, input_fps, output_fps, output_len=None):
features = features.transpose(1, 2) # [1, C, T]
seq_len = features.shape[2] / float(input_fps)
if output_len is None:
output_len = int(seq_len * output_fps)
output_features = F.interpolate(features, size=output_len, align_corners=True, mode='linear')
return output_features.transpose(1, 2)
def get_audio_emb_window(audio_emb, frame_num, frame0_idx, audio_shift=2):
zero_audio_embed = torch.zeros((audio_emb.shape[1], audio_emb.shape[2]), dtype=audio_emb.dtype, device=audio_emb.device)
zero_audio_embed_3 = torch.zeros((3, audio_emb.shape[1], audio_emb.shape[2]), dtype=audio_emb.dtype, device=audio_emb.device)
iter_ = 1 + (frame_num - 1) // 4
audio_emb_wind = []
for lt_i in range(iter_):
if lt_i == 0:
st = frame0_idx + lt_i - 2
ed = frame0_idx + lt_i + 3
wind_feat = torch.stack([
audio_emb[i] if (0 <= i < audio_emb.shape[0]) else zero_audio_embed
for i in range(st, ed)
], dim=0)
wind_feat = torch.cat((zero_audio_embed_3, wind_feat), dim=0)
else:
st = frame0_idx + 1 + 4 * (lt_i - 1) - audio_shift
ed = frame0_idx + 1 + 4 * lt_i + audio_shift
wind_feat = torch.stack([
audio_emb[i] if (0 <= i < audio_emb.shape[0]) else zero_audio_embed
for i in range(st, ed)
], dim=0)
audio_emb_wind.append(wind_feat)
audio_emb_wind = torch.stack(audio_emb_wind, dim=0)
return audio_emb_wind, ed - audio_shift
class WhisperModelLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("audio_encoders"), {"tooltip": "These models are loaded from the 'ComfyUI/models/audio_encoders' folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"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"}),
},
}
RETURN_TYPES = ("WHISPERMODEL",)
RETURN_NAMES = ("whisper_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, load_device):
from transformers import WhisperConfig, WhisperModel, WhisperFeatureExtractor
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]
if load_device == "offload_device":
transformer_load_device = offload_device
else:
transformer_load_device = device
config_path = os.path.join(script_directory, "whisper_config.json")
whisper_config = WhisperConfig(**json.load(open(config_path)))
with init_empty_weights():
whisper = WhisperModel(whisper_config).eval()
whisper.decoder = None # we only need the encoder
feature_extractor_config = {
"chunk_length": 30,
"feature_extractor_type": "WhisperFeatureExtractor",
"feature_size": 128,
"hop_length": 160,
"n_fft": 400,
"n_samples": 480000,
"nb_max_frames": 3000,
"padding_side": "right",
"padding_value": 0.0,
"processor_class": "WhisperProcessor",
"return_attention_mask": False,
"sampling_rate": 16000
}
feature_extractor = WhisperFeatureExtractor(**feature_extractor_config)
model_path = folder_paths.get_full_path_or_raise("audio_encoders", model)
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
for name, param in whisper.named_parameters():
key = "model." + name
value=sd[key]
set_module_tensor_to_device(whisper, name, device=offload_device, dtype=base_dtype, value=value)
whisper_model = {
"feature_extractor": feature_extractor,
"model": whisper,
"dtype": base_dtype,
}
return (whisper_model,)
class HuMoEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"num_frames": ("INT", {"default": 81, "min": -1, "max": 10000, "step": 1, "tooltip": "The total frame count to generate."}),
"width": ("INT", {"default": 832, "min": 64, "max": 4096, "step": 16}),
"height": ("INT", {"default": 480, "min": 64, "max": 4096, "step": 16}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}),
"audio_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The percent of the video to start applying audio conditioning"}),
"audio_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "The percent of the video to stop applying audio conditioning"})
},
"optional" : {
"whisper_model": ("WHISPERMODEL",),
"vae": ("WANVAE", ),
"reference_images": ("IMAGE", {"tooltip": "reference images for the humo model"}),
"audio": ("AUDIO",),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds", )
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, num_frames, width, height, audio_scale, audio_cfg_scale, audio_start_percent, audio_end_percent, whisper_model=None, vae=None, reference_images=None, audio=None, tiled_vae=False):
if reference_images is not None and vae is None:
raise ValueError("VAE is required when reference images are provided")
if whisper_model is None and audio is not None:
raise ValueError("Whisper model is required when audio is provided")
model = whisper_model["model"]
feature_extractor = whisper_model["feature_extractor"]
dtype = whisper_model["dtype"]
sampling_rate = 16000
if audio is not None:
audio_input = audio["waveform"][0]
sample_rate = audio["sample_rate"]
if sample_rate != sampling_rate:
audio_input = torchaudio.functional.resample(audio_input, sample_rate, sampling_rate)
if audio_input.shape[1] == 2:
audio_input = audio_input.mean(dim=0, keepdim=False)
else:
audio_input = audio_input[0]
model.to(device)
audio_len = len(audio_input) // 640
# feature extraction
audio_features = []
window = 750*640
for i in range(0, len(audio_input), window):
audio_feature = feature_extractor(audio_input[i:i+window], sampling_rate=sampling_rate, return_tensors="pt").input_features
audio_features.append(audio_feature)
audio_features = torch.cat(audio_features, dim=-1).to(device, dtype)
# preprocess
window = 3000
audio_prompts = []
for i in range(0, audio_features.shape[-1], window):
audio_prompt = model.encoder(audio_features[:,:,i:i+window], output_hidden_states=True).hidden_states
audio_prompt = torch.stack(audio_prompt, dim=2)
audio_prompts.append(audio_prompt)
model.to(offload_device)
audio_prompts = torch.cat(audio_prompts, dim=1)
audio_prompts = audio_prompts[:,:audio_len*2]
feat0 = linear_interpolation_fps(audio_prompts[:, :, 0: 8].mean(dim=2), 50, 25)
feat1 = linear_interpolation_fps(audio_prompts[:, :, 8: 16].mean(dim=2), 50, 25)
feat2 = linear_interpolation_fps(audio_prompts[:, :, 16: 24].mean(dim=2), 50, 25)
feat3 = linear_interpolation_fps(audio_prompts[:, :, 24: 32].mean(dim=2), 50, 25)
feat4 = linear_interpolation_fps(audio_prompts[:, :, 32], 50, 25)
audio_emb = torch.stack([feat0, feat1, feat2, feat3, feat4], dim=2)[0] # [T, 5, 1280]
else:
audio_emb = torch.zeros(num_frames, 5, 1280, device=device)
audio_len = num_frames
pixel_frame_num = num_frames if num_frames != -1 else audio_len
pixel_frame_num = 4 * ((pixel_frame_num - 1) // 4) + 1
latent_frame_num = (pixel_frame_num - 1) // 4 + 1
log.info(f"HuMo set to generate {pixel_frame_num} frames")
#audio_emb, _ = get_audio_emb_window(audio_emb, pixel_frame_num, frame0_idx=0)
num_refs = 0
if reference_images is not None:
if reference_images.shape[1] != height or reference_images.shape[2] != width:
reference_images_in = common_upscale(reference_images.movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
else:
reference_images_in = reference_images
samples, = WanVideoEncodeLatentBatch.encode(self, vae, reference_images_in, tiled_vae, None, None, None, None)
samples = samples["samples"].transpose(0, 2).squeeze(0)
num_refs = samples.shape[1]
vae.to(device)
zero_frames = torch.zeros(1, 3, pixel_frame_num + 4*num_refs, height, width, device=device, dtype=vae.dtype)
zero_latents = vae.encode(zero_frames, device=device, tiled=tiled_vae)[0].to(offload_device)
vae.to(offload_device)
mm.soft_empty_cache()
target_shape = (16, latent_frame_num + num_refs, height // 8, width // 8)
mask = torch.ones(4, target_shape[1], target_shape[2], target_shape[3], device=offload_device, dtype=vae.dtype)
if reference_images is not None:
mask[:,:-num_refs] = 0
image_cond = torch.cat([zero_latents[:, :(target_shape[1]-num_refs)], samples], dim=1)
#zero_audio_pad = torch.zeros(num_refs, *audio_emb.shape[1:]).to(audio_emb.device)
#audio_emb = torch.cat([audio_emb, zero_audio_pad], dim=0)
else:
image_cond = zero_latents
mask = torch.zeros_like(mask)
image_cond = torch.cat([mask, image_cond], dim=0)
image_cond_neg = torch.cat([mask, zero_latents], dim=0)
embeds = {
"humo_audio_emb": audio_emb,
"humo_audio_emb_neg": torch.zeros_like(audio_emb, dtype=audio_emb.dtype, device=audio_emb.device),
"humo_image_cond": image_cond,
"humo_image_cond_neg": image_cond_neg,
"humo_reference_count": num_refs,
"target_shape": target_shape,
"num_frames": pixel_frame_num,
"humo_audio_scale": audio_scale,
"humo_audio_cfg_scale": audio_cfg_scale,
"humo_start_percent": audio_start_percent,
"humo_end_percent": audio_end_percent,
}
return (embeds, )
class WanVideoCombineEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds_1": ("WANVIDIMAGE_EMBEDS",),
"embeds_2": ("WANVIDIMAGE_EMBEDS",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def add(self, embeds_1, embeds_2):
# Combine the two sets of embeds
combined = {**embeds_1, **embeds_2}
return (combined,)
NODE_CLASS_MAPPINGS = {
"WhisperModelLoader": WhisperModelLoader,
"HuMoEmbeds": HuMoEmbeds,
"WanVideoCombineEmbeds": WanVideoCombineEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WhisperModelLoader": "Whisper Model Loader",
"HuMoEmbeds": "HuMo Embeds",
"WanVideoCombineEmbeds": "WanVideo Combine Embeds",
}
-50
View File
@@ -1,50 +0,0 @@
{
"_name_or_path": "openai/whisper-large-v3",
"activation_dropout": 0.0,
"activation_function": "gelu",
"apply_spec_augment": false,
"architectures": [
"WhisperForConditionalGeneration"
],
"attention_dropout": 0.0,
"begin_suppress_tokens": [
220,
50257
],
"bos_token_id": 50257,
"classifier_proj_size": 256,
"d_model": 1280,
"decoder_attention_heads": 20,
"decoder_ffn_dim": 5120,
"decoder_layerdrop": 0.0,
"decoder_layers": 32,
"decoder_start_token_id": 50258,
"dropout": 0.0,
"encoder_attention_heads": 20,
"encoder_ffn_dim": 5120,
"encoder_layerdrop": 0.0,
"encoder_layers": 32,
"eos_token_id": 50257,
"init_std": 0.02,
"is_encoder_decoder": true,
"mask_feature_length": 10,
"mask_feature_min_masks": 0,
"mask_feature_prob": 0.0,
"mask_time_length": 10,
"mask_time_min_masks": 2,
"mask_time_prob": 0.05,
"max_length": 448,
"max_source_positions": 1500,
"max_target_positions": 448,
"median_filter_width": 7,
"model_type": "whisper",
"num_hidden_layers": 32,
"num_mel_bins": 128,
"pad_token_id": 50256,
"scale_embedding": false,
"torch_dtype": "float16",
"transformers_version": "4.36.0.dev0",
"use_cache": true,
"use_weighted_layer_sum": false,
"vocab_size": 51866
}
-212
View File
@@ -1,212 +0,0 @@
import torch.nn as nn
import torch.nn.functional as F
import torch
import math
from einops import rearrange
from ..wanvideo.modules.model import WanRMSNorm, attention
from ..multitalk.multitalk import RotaryPositionalEmbedding1D, normalize_and_scale
class FeedForwardSwiGLU(nn.Module):
def __init__(
self,
dim: int,
hidden_dim: int,
multiple_of: int = 256,
):
super().__init__()
hidden_dim = int(2 * hidden_dim / 3)
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
self.dim = dim
self.hidden_dim = hidden_dim
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
"""
def __init__(self, t_embed_dim, frequency_embedding_size=256):
super().__init__()
self.t_embed_dim = t_embed_dim
self.frequency_embedding_size = frequency_embedding_size
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, t_embed_dim, bias=True),
nn.SiLU(),
nn.Linear(t_embed_dim, t_embed_dim, bias=True),
)
@staticmethod
def timestep_embedding(t, dim, max_period=10000):
"""
Create sinusoidal timestep embeddings.
:param t: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:param dim: the dimension of the output.
:param max_period: controls the minimum frequency of the embeddings.
:return: an (N, D) Tensor of positional embeddings.
"""
half = dim // 2
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half)
freqs = freqs.to(device=t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def forward(self, t, dtype):
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
if t_freq.dtype != dtype:
t_freq = t_freq.to(dtype)
t_emb = self.mlp(t_freq)
return t_emb
class SingleStreamAttention(nn.Module):
def __init__(
self,
dim: int,
encoder_hidden_states_dim: int,
num_heads: int,
qkv_bias: bool,
qk_norm: bool,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
eps: float = 1e-6,
class_range: int = 24,
class_interval: int = 4,
attention_mode: str = "sdpa",
) -> None:
super().__init__()
assert dim % num_heads == 0, "dim should be divisible by num_heads"
self.dim = dim
self.encoder_hidden_states_dim = encoder_hidden_states_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.scale = self.head_dim**-0.5
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
self.q_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.attn_drop = nn.Dropout(attn_drop)
self.proj = nn.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
self.k_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
# multitalk related params
self.class_interval = class_interval
self.class_range = class_range
self.rope_h1 = (0, self.class_interval)
self.rope_h2 = (self.class_range - self.class_interval, self.class_range)
self.rope_bak = int(self.class_range // 2)
self.rope_1d = RotaryPositionalEmbedding1D(self.head_dim)
def _process_cross_attn(self, x, cond, frames_num=None, x_ref_attn_map=None):
N_t = frames_num
out_dtype = x.dtype
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
# get q for hidden_state
B, N, C = x.shape
q = self.q_linear(x)
q_shape = (B, N, self.num_heads, self.head_dim)
q = q.view(q_shape).permute((0, 2, 1, 3)) # [B, H, N, D]
q = self.q_norm(q.to(self.q_norm.weight.dtype)).to(q.dtype)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
max_values = x_ref_attn_map.max(1).values[:, None, None]
min_values = x_ref_attn_map.min(1).values[:, None, None]
max_min_values = torch.cat([max_values, min_values], dim=2)
human1_max_value, human1_min_value = max_min_values[0, :, 0].max(), max_min_values[0, :, 1].min()
human2_max_value, human2_min_value = max_min_values[1, :, 0].max(), max_min_values[1, :, 1].min()
human1 = normalize_and_scale(x_ref_attn_map[0], (human1_min_value, human1_max_value), (self.rope_h1[0], self.rope_h1[1]))
human2 = normalize_and_scale(x_ref_attn_map[1], (human2_min_value, human2_max_value), (self.rope_h2[0], self.rope_h2[1]))
back = torch.full((x_ref_attn_map.size(1),), self.rope_bak, dtype=human1.dtype).to(human1.device)
max_indices = x_ref_attn_map.argmax(dim=0)
normalized_map = torch.stack([human1, human2, back], dim=1)
normalized_pos = normalized_map[range(x_ref_attn_map.size(1)), max_indices]
q = rearrange(q, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
q = self.rope_1d(q, normalized_pos)
q = rearrange(q, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
# get kv from encoder_hidden_states
_, N_a, _ = cond.shape
encoder_kv = self.kv_linear(cond)
encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim)
encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4))
encoder_k, encoder_v = encoder_kv.unbind(0)
encoder_k = self.k_norm(encoder_k.to(self.k_norm.weight.dtype)).to(encoder_k.dtype)
# multitalk with rope1d pe
if x_ref_attn_map is not None:
per_frame = torch.zeros(N_a, dtype=encoder_k.dtype).to(encoder_k.device)
per_frame[:per_frame.size(0)//2] = (self.rope_h1[0] + self.rope_h1[1]) / 2
per_frame[per_frame.size(0)//2:] = (self.rope_h2[0] + self.rope_h2[1]) / 2
encoder_pos = torch.concat([per_frame]*N_t, dim=0)
encoder_k = rearrange(encoder_k, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
encoder_k = self.rope_1d(encoder_k, encoder_pos)
encoder_k = rearrange(encoder_k, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
# Input tensors must be in format ``[B, M, H, K]``, where B is the batch size, M \
# the sequence length, H the number of heads, and K the embeding size per head
q = rearrange(q, "B H M K -> B M H K")
encoder_k = rearrange(encoder_k, "B H M K -> B M H K")
encoder_v = rearrange(encoder_v, "B H M K -> B M H K")
x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode)
x = rearrange(x, "B M H K -> B H M K")
# linear transform
x_output_shape = (B, N, C)
x = x.transpose(1, 2)
x = x.reshape(x_output_shape)
x = self.proj(x)
x = self.proj_drop(x)
# reshape x to origin shape
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
return x.type(out_dtype)
def forward(self, x, cond, num_latent_frames=None, num_cond_latents=None, x_ref_attn_map=None, human_num=None):
B, N, C = x.shape
if (num_cond_latents is None or num_cond_latents == 0):
# text to video
output = self._process_cross_attn(x, cond, num_latent_frames, x_ref_attn_map)
return None, output
elif num_cond_latents is not None and num_cond_latents > 0:
# image to video or video continuation
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
x_noise = x[:, num_cond_latents_thw:]
cond = rearrange(cond, "(B N_t) M C -> B N_t M C", B=B)
cond = cond[:, num_cond_latents:]
cond = rearrange(cond, "B N_t M C -> (B N_t) M C")
frames_num = num_latent_frames - num_cond_latents
if human_num is not None and human_num == 2:
# multitalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num, x_ref_attn_map)
else:
# singletalk mode
output_noise = self._process_cross_attn(x_noise, cond, frames_num)
output_cond = torch.zeros((B, num_cond_latents_thw, C), dtype=output_noise.dtype, device=output_noise.device)
return output_cond, output_noise
else:
raise NotImplementedError
File diff suppressed because it is too large Load Diff
-120
View File
@@ -1,120 +0,0 @@
import torch
from ..utils import log
import comfy.model_management as mm
from comfy_api.latest import io
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="WanVideoLongCatAvatarExtendEmbeds",
category="WanVideoWrapper",
inputs=[
io.Latent.Input("prev_latents", tooltip="Full previous latents to be used to continue generation, continuation frames are selected based on 'overlap' parameter"),
io.Custom("MULTITALK_EMBEDS").Input("audio_embeds", tooltip="Full length audio embeddings"),
io.Int.Input("num_frames", default=93, min=1, max=256, step=1, tooltip="Number of new frames to generate"),
io.Int.Input("overlap", default=13, min=0, max=16, step=1, tooltip="Number of overlapping frames from previous latents for video continuation, set to 0 for T2V"),
io.Int.Input("frames_processed", default=0, min=0, max=10000, step=1, tooltip="Number of frames already processed in the video, used to select audio features"),
io.Combo.Input("if_not_enough_audio", ["pad_with_start", "mirror_from_end"], default="pad_with_start", tooltip="What to do if there are not enough frames in pose_images for the window"),
io.Int.Input("ref_frame_index", default=10, min=0, max=1000, step=1, tooltip="Values between 0 - 24 ensures better consistency, while selecting other ranges (e.g., -10 or 30) helps reduce repeated actions"),
io.Int.Input("ref_mask_frame_range", default=3, min=0, max=20, step=1, tooltip="Larger range can further help mitigate repeated actions, but excessively large values may introduce artifacts"),
io.Latent.Input("ref_latent", optional=True, tooltip="Reference latent used for consistency, generally should be either the init image, or first latent from first generation"),
io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"),
],
outputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
io.Latent.Output(display_name="samples_slice", tooltip="Sliced latent samples for the new frames"),
],
)
@classmethod
def execute(cls, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed, ref_frame_index, ref_mask_frame_range, ref_latent=None, samples=None) -> io.NodeOutput:
new_audio_embed = audio_embeds.copy()
audio_features = torch.stack(new_audio_embed["audio_features"])
num_audio_features = audio_features.shape[1]
if audio_features.shape[1] < frames_processed + num_frames:
deficit = frames_processed + num_frames - audio_features.shape[1]
if if_not_enough_audio == "pad_with_start":
pad = audio_features[:, :1].repeat(1, deficit, 1, 1)
audio_features = torch.cat([audio_features, pad], dim=1)
elif if_not_enough_audio == "mirror_from_end":
to_add = audio_features[:, -deficit:, :].flip(dims=[1])
audio_features = torch.cat([audio_features, to_add], dim=1)
log.warning(f"Not enough audio features, padded with strategy '{if_not_enough_audio}' from {num_audio_features} to {audio_features.shape[1]} frames")
ref_target_masks = new_audio_embed.get("ref_target_masks", None)
if ref_target_masks is not None:
new_audio_embed["ref_target_masks"] = ref_target_masks[:, frames_processed:frames_processed+num_frames, :]
prev_samples = prev_latents["samples"].clone()
if overlap != 0:
latent_overlap = (overlap - 1) // 4 + 1
prev_samples = prev_samples[:, :, -latent_overlap:]
ref_sample = None
if ref_latent is not None:
ref_sample = ref_latent["samples"][0, :, :1].clone()
log.info(f"Previous latents shape: {prev_samples.shape}, using last {latent_overlap} latent frames for overlap.")
new_latent_frames = (num_frames - 1) // 4 + 1
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
audio_stride = 2
indices = torch.arange(2 * 2 + 1) - 2
if frames_processed == 0:
audio_start_idx = 0
else:
audio_start_idx = (frames_processed - overlap) * audio_stride
audio_end_idx = audio_start_idx + num_frames * audio_stride
log.info(f"Extracting audio embeddings from index {audio_start_idx} to {audio_end_idx}")
audio_embs = []
for human_idx in range(len(audio_features)):
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
center_indices = torch.clamp(center_indices, min=0, max=audio_features[human_idx].shape[0] - 1)
audio_emb = audio_features[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb)
audio_emb = torch.cat(audio_embs, dim=0)
new_audio_embed["audio_features"] = None
new_audio_embed["audio_emb_slice"] = audio_emb
longcat_avatar_options = {
"longcat_ref_latent": ref_sample,
"ref_frame_index": ref_frame_index,
"ref_mask_frame_range": ref_mask_frame_range,
}
embeds = {
"target_shape": target_shape,
"num_frames": num_frames,
"extra_latents": [{"samples": prev_samples, "index": 0}] if overlap != 0 else None,
"multitalk_embeds": new_audio_embed,
"longcat_avatar_options": longcat_avatar_options,
}
samples_slice = None
if samples is not None:
latent_start_index = (frames_processed - 1) // 4 + 1 if frames_processed > 0 else 0
latent_end_index = latent_start_index + new_latent_frames
samples_slice = samples.copy()
samples_slice["samples"] = samples["samples"][:, :, latent_start_index:latent_end_index].clone()
return io.NodeOutput(embeds, samples_slice)
NODE_CLASS_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
}
-199
View File
@@ -1,199 +0,0 @@
import torch
import torch.nn as nn
from einops import rearrange
from ..wanvideo.modules.attention import attention
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
return (x * (1 + scale) + shift)
def sinusoidal_embedding_1d(dim, position):
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x.to(position.dtype)
def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0):
# 3d rope precompute
f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta)
h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
return f_freqs_cis, h_freqs_cis, w_freqs_cis
def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0):
# 1d rope precompute
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
[: (dim // 2)].double() / dim))
freqs = torch.outer(torch.arange(end, device=freqs.device), freqs)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
return freqs_cis
def rope_apply(x, freqs, num_heads):
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
x_out = torch.view_as_complex(x.to(torch.float64).reshape(
x.shape[0], x.shape[1], x.shape[2], -1, 2))
x_out = torch.view_as_real(x_out * freqs).flatten(2)
return x_out.to(x.dtype)
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-5):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
def forward(self, x):
dtype = x.dtype
return self.norm(x.float()).to(dtype) * self.weight
class AttentionModule(nn.Module):
def __init__(self, num_heads, head_dim):
super().__init__()
self.num_heads = num_heads
self.head_dim = head_dim
def forward(self, q, k, v):
b, n, d = q.size(0), self.num_heads, self.head_dim
x = attention(
q.view(b, -1, n, d),
k.view(b, -1, n, d),
v.view(b, -1, n, d)
)
return x.flatten(2)
class SelfAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x, freqs):
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(x))
v = self.v(x)
q = rope_apply(q, freqs, self.num_heads)
k = rope_apply(k, freqs, self.num_heads)
x = self.attn(q, k, v)
return self.o(x)
class CrossAttention(nn.Module):
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
self.k_img = nn.Linear(dim, dim)
self.v_img = nn.Linear(dim, dim)
self.norm_k_img = RMSNorm(dim, eps=eps)
self.attn = AttentionModule(self.num_heads, self.head_dim)
def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None):
ctx = y
q = self.norm_q(self.q(x))
k = self.norm_k(self.k(ctx))
v = self.v(ctx)
x = self.attn(q, k, v)
if clip_fea is not None:
k_img = self.norm_k_img(self.k_img(clip_fea))
v_img = self.v_img(clip_fea)
y = self.attn(q, k_img, v_img)
x = x + y
return self.o(x)
class GateModule(nn.Module):
def __init__(self,):
super().__init__()
def forward(self, x, gate, residual):
return x + gate * residual
class DiTBlock(nn.Module):
def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.ffn_dim = ffn_dim
self.self_attn = SelfAttention(dim, num_heads, eps)
self.cross_attn = CrossAttention(dim, num_heads, eps)
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
self.norm3 = nn.LayerNorm(dim, eps=eps)
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(
approximate='tanh'), nn.Linear(ffn_dim, dim))
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
self.gate = GateModule()
def forward(self, x, context, t_mod, freqs, clip_fea=None):
has_seq = len(t_mod.shape) == 4
chunk_dim = 2 if has_seq else 1
# msa: multi-head self-attention mlp: multi-layer perceptron
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)
if has_seq:
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
)
input_x = modulate(self.norm1(x), shift_msa, scale_msa)
x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea)
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
x = self.gate(x, gate_mlp, self.ffn(input_x))
return x
class WanModelDualControl(torch.nn.Module):
def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12):
super().__init__()
self.control_layers = control_layers
self.control_blocks_dense = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_blocks_sparse = nn.ModuleList([
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
for _ in range(self.control_layers)
])
self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2)
self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2)
self.control_text_linear = torch.nn.Linear(dim, dim//2)
self.control_t_mod = torch.nn.Linear(dim, dim//2)
self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)])
head_dim = dim // num_heads
self.freqs = precompute_freqs_cis_3d(head_dim)
-88
View File
@@ -1,88 +0,0 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddDualControlEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
"first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}),
},
"optional": {
"dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}),
"sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}),
"prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None):
updated = dict(embeds)
updated.setdefault("dual_control", {})
if dense is None and sparse is None:
raise ValueError("At least one of dense or sparse inputs must be provided.")
num_frames = dense.shape[0] if dense is not None else sparse.shape[0]
height = dense.shape[1] if dense is not None else sparse.shape[1]
width = dense.shape[2] if dense is not None else sparse.shape[2]
msk = torch.ones(1, num_frames, height//8, width//8, device=device)
msk[:, 1:] = 0
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
msk = msk.transpose(1, 2)
dense_input_latent = sparse_input_latent = None
vae.to(device)
if dense is not None:
dense_images = 1 - dense[..., :3] # Invert colors for depth to match the usual range in comfy
dense_images = dense_images.permute(3, 0, 1, 2) * 2 - 1
dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False)
dense_first = (dense_images[:, :1]).to(device, vae.dtype)
vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False)
dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1)
dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1)
if sparse is not None:
sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1
sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False)
sparse_first = (sparse_images[:, :1]).to(device, vae.dtype)
vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False)
sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1)
sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1)
if prev_images is not None:
prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1
prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False)
updated["dual_control"]["prev_latent"] = prev_video_latent[0]
vae.to(offload_device)
updated["dual_control"]["dense_input_latent"] = dense_input_latent
updated["dual_control"]["sparse_input_latent"] = sparse_input_latent
updated["dual_control"]["strength"] = strength
updated["dual_control"]["start_percent"] = start_percent
updated["dual_control"]["end_percent"] = end_percent
updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds",
}
BIN
View File
Binary file not shown.
BIN
View File
Binary file not shown.
-213
View File
@@ -1,213 +0,0 @@
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, stick_width=1.0, point_radius=2, style="original"):
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)
if style == "original":
canvas = draw_3d_points(
canvas,
joints3d[0],
stickwidth=int(h / 350 * stick_width),
r=point_radius,
)
elif style == "scail":
canvas = draw_3d_points_scail(
canvas,
joints3d[0],
stickwidth=int(h / 350 * stick_width),
r=point_radius,
)
resized_canvas = cv2.resize(canvas, (w, h))
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
control_images.append(resized_canvas)
except Exception:
control_images.append(Image.fromarray(canvas))
control_pixel_values = np.array(control_images)
control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
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
def draw_3d_points_scail(canvas, points, stickwidth=2, r=2, draw_line=True):
connetions = [
[15,12],[12, 16],[16, 18],[18, 20],[20, 22], # 0-4: Left arm chain
[12,17],[17,19],[19,21], # 5-7: Right arm chain
[21,23], # 8: Right hand
[12,1],[1,4],[4,7], # 9-11: Neck to left leg (hip, thigh, shin)
[12,2],[2,5],[5,8], # 12-14: Neck to right leg (hip, thigh, shin)
]
# Warm colors for right side, cool colors for left side
connection_colors = [
[180, 180, 180], # 0: [15,12] - L. clavicle (Bright Cyan)
[0, 200, 255], # 1: [12,16] - L. shoulder (Bright Cyan)
[0, 120, 255], # 2: [16,18] - L. upper arm (Bright Blue)
[0, 60, 255], # 3: [18,20] - L. forearm (Deep Blue)
[60, 0, 255], # 4: [20,22] - L. hand (Blue-Purple)
[255, 0, 0], # 5: [12,17] - R. clavicle (Bright Red)
[255, 100, 0], # 6: [17,19] - R. upper arm (Bright Orange)
[255, 180, 0], # 7: [19,21] - R. forearm (Golden Orange)
[255, 255, 0], # 8: [21,23] - R. hand (Bright Yellow)
[30, 27, 160], # 9: [12,1] - Neck to L. hip (purple-blue)
[73, 27, 177], # 10: [1,4] - L. thigh (purple)
[145, 27, 194], # 11: [4,7] - L. shin (magenta)
[200, 255, 100], # 12: [12,2] - Neck to R. hip (yellow)
[54, 201, 52], # 13: [2,5] - R. thigh (green)
[30, 176, 85], # 14: [5,8] - R. shin (green)
]
# draw line
if draw_line:
# Collect all joints that are part of connections
joints_in_use = set()
for connection in connetions:
joints_in_use.add(connection[0])
joints_in_use.add(connection[1])
for i in range(len(connetions)):
point1_idx, point2_idx = connetions[i][0:2]
point1 = points[point1_idx]
point2 = points[point2_idx]
x1, y1 = int(point1[0]), int(point1[1])
x2, y2 = int(point2[0]), int(point2[1])
cv2.line(canvas, (x1, y1), (x2, y2), connection_colors[i], stickwidth)
# draw points for joints that have connections
joints_in_use = set()
for connection in connetions:
joints_in_use.add(connection[0])
joints_in_use.add(connection[1])
for joint_idx in joints_in_use:
if joint_idx >= len(points):
continue
x, y = points[joint_idx][0:2]
x, y = int(x), int(y)
# Use the color from the first connection involving this joint
joint_color = [180, 180, 180] # default grey
for i, connection in enumerate(connetions):
if connection[0] == joint_idx or connection[1] == joint_idx:
joint_color = connection_colors[i]
break
cv2.circle(canvas, (x, y), r, joint_color, thickness=-1)
return canvas
-1
View File
@@ -1 +0,0 @@
from .vqvae import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
-329
View File
@@ -1,329 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
class Encoder(nn.Module):
def __init__(
self,
in_channels=3,
mid_channels=[128, 512],
out_channels=3072,
downsample_time=[1, 1],
downsample_joint=[1, 1],
num_attention_heads=8,
attention_head_dim=64,
dim=3072,
):
super(Encoder, self).__init__()
self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
self.downsample1 = Downsample(mid_channels[0], mid_channels[0], downsample_time[0], downsample_joint[0])
self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
self.downsample2 = Downsample(mid_channels[1], mid_channels[1], downsample_time[1], downsample_joint[1])
self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = self.conv_in(x)
for resnet in self.resnet1:
x = resnet(x)
x = self.downsample1(x)
x = self.resnet2(x)
for resnet in self.resnet3:
x = resnet(x)
x = self.downsample2(x)
x = self.conv_out(x)
return x
class VectorQuantizer(nn.Module):
def __init__(self, nb_code, code_dim):
super().__init__()
self.nb_code = nb_code
self.code_dim = code_dim
self.mu = 0.99
self.reset_codebook()
self.reset_count = 0
self.usage = torch.zeros((self.nb_code, 1))
def reset_codebook(self):
self.init = False
self.code_sum = None
self.code_count = None
self.register_buffer('codebook', torch.zeros(self.nb_code, self.code_dim).cuda())
def _tile(self, x):
nb_code_x, code_dim = x.shape
if nb_code_x < self.nb_code:
n_repeats = (self.nb_code + nb_code_x - 1) // nb_code_x
std = 0.01 / np.sqrt(code_dim)
out = x.repeat(n_repeats, 1)
out = out + torch.randn_like(out) * std
else:
out = x
return out
def preprocess(self, x):
# [bs, c, f, j] -> [bs * f * j, c]
x = x.permute(0, 2, 3, 1).contiguous()
x = x.view(-1, x.shape[-1])
return x
def quantize(self, x):
# [bs * f * j, dim=3072]
# Calculate latent code x_l
k_w = self.codebook.t()
distance = torch.sum(x ** 2, dim=-1, keepdim=True) - 2 * torch.matmul(x, k_w) + torch.sum(k_w ** 2, dim=0, keepdim=True)
_, code_idx = torch.min(distance, dim=-1)
return code_idx
def dequantize(self, code_idx):
x = F.embedding(code_idx, self.codebook) # indexing: [bs * f * j, 32]
return x
def forward(self, x, return_vq=False):
bs, c, f, j = x.shape # SMPL data frames: [bs, 3072, f, j]
# Preprocess
x = self.preprocess(x)
# return x.view(bs, f*j, c).contiguous(), None
assert x.shape[-1] == self.code_dim
# quantize and dequantize through bottleneck
code_idx = self.quantize(x)
x_d = self.dequantize(code_idx)
# Loss
commit_loss = F.mse_loss(x, x_d.detach())
# Passthrough
x_d = x + (x_d - x).detach()
if return_vq:
return x_d.view(bs, f*j, c).contiguous(), commit_loss
# return (x_d, x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()), commit_loss, perplexity
# Postprocess
x_d = x_d.view(bs, f, j, c).permute(0, 3, 1, 2).contiguous()
return x_d, commit_loss
class Decoder(nn.Module):
def __init__(
self,
in_channels=3072,
mid_channels=[512, 128],
out_channels=3,
upsample_rate=None,
frame_upsample_rate=[1.0, 1.0],
joint_upsample_rate=[1.0, 1.0],
dim=128,
attention_head_dim=64,
num_attention_heads=8,
):
super(Decoder, self).__init__()
self.conv_in = nn.Conv2d(in_channels, mid_channels[0], kernel_size=3, stride=1, padding=1)
self.resnet1 = nn.ModuleList([ResBlock(mid_channels[0], mid_channels[0]) for _ in range(3)])
self.upsample1 = Upsample(mid_channels[0], mid_channels[0], frame_upsample_rate=frame_upsample_rate[0], joint_upsample_rate=joint_upsample_rate[0])
self.resnet2 = ResBlock(mid_channels[0], mid_channels[1])
self.resnet3 = nn.ModuleList([ResBlock(mid_channels[1], mid_channels[1]) for _ in range(3)])
self.upsample2 = Upsample(mid_channels[1], mid_channels[1], frame_upsample_rate=frame_upsample_rate[1], joint_upsample_rate=joint_upsample_rate[1])
self.conv_out = nn.Conv2d(mid_channels[-1], out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
x = self.conv_in(x)
for resnet in self.resnet1:
x = resnet(x)
x = self.upsample1(x)
x = self.resnet2(x)
for resnet in self.resnet3:
x = resnet(x)
x = self.upsample2(x)
x = self.conv_out(x)
return x
class Upsample(nn.Module):
def __init__(
self,
in_channels,
out_channels,
upsample_rate=None,
frame_upsample_rate=None,
joint_upsample_rate=None,
):
super(Upsample, self).__init__()
self.upsampler = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.upsample_rate = upsample_rate
self.frame_upsample_rate = frame_upsample_rate
self.joint_upsample_rate = joint_upsample_rate
self.upsample_rate = upsample_rate
def forward(self, inputs):
if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1:
# split first frame
x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:]
if self.upsample_rate is not None:
# import pdb; pdb.set_trace()
x_first = F.interpolate(x_first, scale_factor=self.upsample_rate)
x_rest = F.interpolate(x_rest, scale_factor=self.upsample_rate)
else:
# import pdb; pdb.set_trace()
# x_first = F.interpolate(x_first, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
x_rest = F.interpolate(x_rest, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
x_first = x_first[:, :, None, :]
inputs = torch.cat([x_first, x_rest], dim=2)
elif inputs.shape[2] > 1:
if self.upsample_rate is not None:
inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
else:
inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="bilinear", align_corners=True)
else:
inputs = inputs.squeeze(2)
if self.upsample_rate is not None:
inputs = F.interpolate(inputs, scale_factor=self.upsample_rate)
else:
inputs = F.interpolate(inputs, scale_factor=(self.frame_upsample_rate, self.joint_upsample_rate), mode="linear", align_corners=True)
inputs = inputs[:, :, None, :, :]
b, c, t, j = inputs.shape
inputs = inputs.permute(0, 2, 1, 3).reshape(b * t, c, j)
inputs = self.upsampler(inputs)
inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3)
return inputs
class Downsample(nn.Module):
def __init__(
self,
in_channels,
out_channels,
frame_downsample_rate,
joint_downsample_rate
):
super(Downsample, self).__init__()
self.frame_downsample_rate = frame_downsample_rate
self.joint_downsample_rate = joint_downsample_rate
self.joint_downsample = nn.Conv1d(in_channels, out_channels, kernel_size=3, stride=self.joint_downsample_rate, padding=1)
def forward(self, x):
# (batch_size, channels, frames, joints) -> (batch_size * joints, channels, frames)
if self.frame_downsample_rate > 1:
batch_size, channels, frames, joints = x.shape
x = x.permute(0, 3, 1, 2).reshape(batch_size * joints, channels, frames)
if x.shape[-1] % 2 == 1:
x_first, x_rest = x[..., 0], x[..., 1:]
if x_rest.shape[-1] > 0:
# (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2)
x_rest = F.avg_pool1d(x_rest, kernel_size=self.frame_downsample_rate, stride=self.frame_downsample_rate)
x = torch.cat([x_first[..., None], x_rest], dim=-1)
# (batch_size * joints, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, joints)
x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
else:
# (batch_size * joints, channels, frames) -> (batch_size * joints, channels, frames // 2)
x = F.avg_pool1d(x, kernel_size=2, stride=2)
# (batch_size * joints, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width)
x = x.reshape(batch_size, joints, channels, x.shape[-1]).permute(0, 2, 3, 1)
# Pad the tensor
# pad = (0, 1)
# x = F.pad(x, pad, mode="constant", value=0)
batch_size, channels, frames, joints = x.shape
# (batch_size, channels, frames, joints) -> (batch_size * frames, channels, joints)
x = x.permute(0, 2, 1, 3).reshape(batch_size * frames, channels, joints)
x = self.joint_downsample(x)
# (batch_size * frames, channels, joints) -> (batch_size, channels, frames, joints)
x = x.reshape(batch_size, frames, x.shape[1], x.shape[2]).permute(0, 2, 1, 3)
return x
class ResBlock(nn.Module):
def __init__(self,
in_channels,
out_channels,
group_num=32,
max_channels=512):
super(ResBlock, self).__init__()
skip = max(1, max_channels // out_channels - 1)
self.block = nn.Sequential(
nn.GroupNorm(group_num, in_channels, eps=1e-06, affine=True),
nn.SiLU(),
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=skip, dilation=skip),
nn.GroupNorm(group_num, out_channels, eps=1e-06, affine=True),
nn.SiLU(),
nn.Conv2d(out_channels, out_channels, kernel_size=1, stride=1, padding=0),
)
self.conv_short = nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, padding=0) if in_channels != out_channels else nn.Identity()
def forward(self, x):
hidden_states = self.block(x)
if hidden_states.shape != x.shape:
x = self.conv_short(x)
x = x + hidden_states
return x
class SMPL_VQVAE(nn.Module):
def __init__(self, encoder, decoder, vq):
super(SMPL_VQVAE, self).__init__()
self.encoder = encoder
self.decoder = decoder
self.vq = vq
def to(self, device):
self.encoder = self.encoder.to(device)
self.decoder = self.decoder.to(device)
self.vq = self.vq.to(device)
self.device = device
return self
def encdec_slice_frames(self, x, frame_batch_size, encdec, return_vq):
num_frames = x.shape[2]
remaining_frames = num_frames % frame_batch_size
x_output = []
for i in range(num_frames // frame_batch_size):
remaining_frames = num_frames % frame_batch_size
start_frame = frame_batch_size * i + (0 if i == 0 else remaining_frames)
end_frame = frame_batch_size * (i + 1) + remaining_frames
x_intermediate = x[:, :, start_frame:end_frame]
x_intermediate = encdec(x_intermediate)
x_output.append(x_intermediate)
if encdec == self.encoder and self.vq is not None:
x_output, loss = self.vq(torch.cat(x_output, dim=2), return_vq=return_vq)
return x_output, loss
else:
return torch.cat(x_output, dim=2), None, None
def forward(self, x, return_vq=False):
x = x.permute(0, 3, 1, 2)
x, loss = self.encdec_slice_frames(x, frame_batch_size=8, encdec=self.encoder, return_vq=return_vq)
if return_vq:
return x, loss
x, _, _ = self.encdec_slice_frames(x, frame_batch_size=2, encdec=self.decoder, return_vq=return_vq)
x = x.permute(0, 2, 3, 1)
return x, loss
-193
View File
@@ -1,193 +0,0 @@
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
-340
View File
@@ -1,340 +0,0 @@
import os
import torch
from ..utils import log
import numpy as np
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")
folder_paths.add_model_folder_path("nlf", os.path.join(folder_paths.models_dir, "nlf"))
from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
def check_jit_script_function():
if torch.jit.script.__name__ != "script":
# Get more details about what modified it
module = torch.jit.script.__module__
qualname = getattr(torch.jit.script, '__qualname__', 'unknown')
code_file = None
try:
code_file = torch.jit.script.__code__.co_filename
code_line = torch.jit.script.__code__.co_firstlineno
log.warning(f"torch.jit.script has been modified by another custom node.\n"
f" Function name: {torch.jit.script.__name__}\n"
f" Module: {module}\n"
f" Qualified name: {qualname}\n"
f" Defined in: {code_file}:{code_line}\n"
f"This may cause issues with the NLF model.")
except:
log.warning("--------------------------------")
log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, "
f"this has been modified by another custom node. This may cause issues with the NLF model.")
log.warning("--------------------------------")
model_list = [
"https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript",
"https://github.com/isarandi/nlf/releases/download/v0.2.2/nlf_l_multi_0.2.2.torchscript",
]
class DownloadAndLoadNLFModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"url": (model_list, {"default": "https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"}),
},
"optional": {
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
},
}
RETURN_TYPES = ("NLFMODEL",)
RETURN_NAMES = ("nlf_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, url, warmup=True):
if url not in model_list:
raise ValueError(f"URL {url} is not in the list of allowed models.")
check_jit_script_function()
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()
if warmup:
log.info("Warming up NLF model...")
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
try:
for _ in range(2):
_ = model.detect_smpl_batched(dummy_input)
finally:
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
log.info("NLF model warmed up")
model = model.to(offload_device)
return (model,)
class LoadNLFModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"nlf_model": (folder_paths.get_filename_list("nlf"), {"tooltip": "These models are loaded from the 'ComfyUI/models/nlf' -folder",}),
},
"optional": {
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
},
}
RETURN_TYPES = ("NLFMODEL",)
RETURN_NAMES = ("nlf_model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, nlf_model, warmup=True):
check_jit_script_function()
model = torch.jit.load(folder_paths.get_full_path_or_raise("nlf", nlf_model)).eval()
if warmup:
log.info("Warming up NLF model...")
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
try:
for _ in range(2):
_ = model.detect_smpl_batched(dummy_input)
finally:
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
log.info("NLF model warmed up")
model = model.to(offload_device)
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):
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"}),
},
"optional": {
"per_batch": ("INT", {"default": -1, "min": -1, "max": 10000, "step": 1, "tooltip": "How many images to process at once. -1 means all at once."}),
}
}
RETURN_TYPES = ("NLFPRED", "BBOX",)
RETURN_NAMES = ("pose_results", "bboxes")
FUNCTION = "predict"
CATEGORY = "WanVideoWrapper"
def predict(self, model, images, per_batch=-1):
check_jit_script_function()
model = model.to(device)
num_images = images.shape[0]
# Determine batch size
if per_batch == -1:
batch_size = num_images
else:
batch_size = per_batch
# Initialize result containers
all_boxes = []
all_joints3d_nonparam = []
# Process in batches
for i in range(0, num_images, batch_size):
end_idx = min(i + batch_size, num_images)
batch_images = images[i:end_idx]
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
try:
pred = model.detect_smpl_batched(batch_images.permute(0, 3, 1, 2).to(device))
finally:
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
# Collect boxes and joints from this batch
if 'boxes' in pred:
all_boxes.extend(pred['boxes'])
if 'joints3d_nonparam' in pred:
all_joints3d_nonparam.extend(pred['joints3d_nonparam'])
model = model.to(offload_device)
# Move collected results to offload device
all_boxes = [box.to(offload_device) for box in all_boxes]
all_joints3d_nonparam = [joints.to(offload_device) for joints in all_joints3d_nonparam]
# Maintain the original nested format: wrap in a list to match expected structure
pose_results = {
'joints3d_nonparam': [all_joints3d_nonparam],
}
# Convert bboxes to list format: [x_min, y_min, x_max, y_max] for each detection
# Each box tensor is shape (1, 5) with [x_min, y_min, x_max, y_max, confidence]
formatted_boxes = []
for box in all_boxes:
# Handle empty detections (no person detected in frame)
if box.numel() == 0 or box.shape[0] == 0:
formatted_boxes.append([0.0, 0.0, 0.0, 0.0])
else:
# Extract first 4 values (x_min, y_min, x_max, y_max), drop confidence
bbox_values = box[0, :4].cpu().tolist()
formatted_boxes.append(bbox_values)
return (pose_results, formatted_boxes)
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}),
},
"optional": {
"stick_width": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 1000.0, "step": 0.01, "tooltip": "Stick width multiplier"}),
"point_radius": ("INT", {"default": 5, "min": 1, "max": 10, "step": 1, "tooltip": "Point radius for drawing the pose"}),
"style": (["original", "scail"], {"default": "original", "tooltip": "style of the pose drawing"}),
}
}
RETURN_TYPES = ("IMAGE", )
RETURN_NAMES = ("image",)
FUNCTION = "predict"
CATEGORY = "WanVideoWrapper"
def predict(self, poses, width, height, stick_width=1.0, point_radius=2, style="original"):
from .draw_pose import get_control_conditions
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, stick_width=stick_width, point_radius=point_radius, style=style)
return (control_conditions,)
NODE_CLASS_MAPPINGS = {
"LoadNLFModel": LoadNLFModel,
"DownloadAndLoadNLFModel": DownloadAndLoadNLFModel,
"NLFPredict": NLFPredict,
"DrawNLFPoses": DrawNLFPoses,
"LoadVQVAE": LoadVQVAE,
"MTVCrafterEncodePoses": MTVCrafterEncodePoses
}
NODE_DISPLAY_NAME_MAPPINGS = {
"LoadNLFModel": "Load NLF Model",
"DownloadAndLoadNLFModel": "(Download)Load NLF Model",
"NLFPredict": "NLF Predict",
"DrawNLFPoses": "Draw NLF Poses",
"LoadVQVAE": "Load VQVAE",
"MTVCrafterEncodePoses": "MTV Crafter Encode Poses"
}
-48
View File
@@ -1,48 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class ChannelLastConv1d(nn.Conv1d):
def forward(self, x: torch.Tensor) -> torch.Tensor:
x = x.permute(0, 2, 1)
x = super().forward(x)
x = x.permute(0, 2, 1)
return x
class ConvMLP(nn.Module):
def __init__(
self,
dim: int,
hidden_dim: int,
multiple_of: int = 256,
kernel_size: int = 3,
padding: int = 1,
):
"""
Initialize the FeedForward module.
Args:
dim (int): Input dimension.
hidden_dim (int): Hidden dimension of the feedforward layer.
multiple_of (int): Value to ensure hidden dimension is a multiple of this value.
Attributes:
w1 (ColumnParallelLinear): Linear transformation for the first layer.
w2 (RowParallelLinear): Linear transformation for the second layer.
w3 (ColumnParallelLinear): Linear transformation for the third layer.
"""
super().__init__()
hidden_dim = int(2 * hidden_dim / 3)
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
self.w1 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
self.w2 = ChannelLastConv1d(hidden_dim, dim, bias=False, kernel_size=kernel_size, padding=padding)
self.w3 = ChannelLastConv1d(dim, hidden_dim, bias=False, kernel_size=kernel_size, padding=padding)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))
-21
View File
@@ -1,21 +0,0 @@
MIT License
Copyright (c) 2022 NVIDIA CORPORATION.
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
-1
View File
@@ -1 +0,0 @@
from .bigvgan import BigVGAN
-120
View File
@@ -1,120 +0,0 @@
# Implementation adapted from https://github.com/EdwardDixon/snake under the MIT license.
# LICENSE is in incl_licenses directory.
import torch
from torch import nn, sin, pow
from torch.nn import Parameter
class Snake(nn.Module):
'''
Implementation of a sine-based periodic activation function
Shape:
- Input: (B, C, T)
- Output: (B, C, T), same shape as the input
Parameters:
- alpha - trainable parameter
References:
- This activation function is from this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
https://arxiv.org/abs/2006.08195
Examples:
>>> a1 = snake(256)
>>> x = torch.randn(256)
>>> x = a1(x)
'''
def __init__(self, in_features, alpha=1.0, alpha_trainable=True, alpha_logscale=False):
'''
Initialization.
INPUT:
- in_features: shape of the input
- alpha: trainable parameter
alpha is initialized to 1 by default, higher values = higher-frequency.
alpha will be trained along with the rest of your model.
'''
super(Snake, self).__init__()
self.in_features = in_features
# initialize alpha
self.alpha_logscale = alpha_logscale
if self.alpha_logscale: # log scale alphas initialized to zeros
self.alpha = Parameter(torch.zeros(in_features) * alpha)
else: # linear scale alphas initialized to ones
self.alpha = Parameter(torch.ones(in_features) * alpha)
self.alpha.requires_grad = alpha_trainable
self.no_div_by_zero = 0.000000001
def forward(self, x):
'''
Forward pass of the function.
Applies the function to the input elementwise.
Snake ∶= x + 1/a * sin^2 (xa)
'''
alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # line up with x to [B, C, T]
if self.alpha_logscale:
alpha = torch.exp(alpha)
x = x + (1.0 / (alpha + self.no_div_by_zero)) * pow(sin(x * alpha), 2)
return x
class SnakeBeta(nn.Module):
'''
A modified Snake function which uses separate parameters for the magnitude of the periodic components
Shape:
- Input: (B, C, T)
- Output: (B, C, T), same shape as the input
Parameters:
- alpha - trainable parameter that controls frequency
- beta - trainable parameter that controls magnitude
References:
- This activation function is a modified version based on this paper by Liu Ziyin, Tilman Hartwig, Masahito Ueda:
https://arxiv.org/abs/2006.08195
Examples:
>>> a1 = snakebeta(256)
>>> x = torch.randn(256)
>>> x = a1(x)
'''
def __init__(self, in_features, alpha=1.0, alpha_trainable=True, alpha_logscale=False):
'''
Initialization.
INPUT:
- in_features: shape of the input
- alpha - trainable parameter that controls frequency
- beta - trainable parameter that controls magnitude
alpha is initialized to 1 by default, higher values = higher-frequency.
beta is initialized to 1 by default, higher values = higher-magnitude.
alpha will be trained along with the rest of your model.
'''
super(SnakeBeta, self).__init__()
self.in_features = in_features
# initialize alpha
self.alpha_logscale = alpha_logscale
if self.alpha_logscale: # log scale alphas initialized to zeros
self.alpha = Parameter(torch.zeros(in_features) * alpha)
self.beta = Parameter(torch.zeros(in_features) * alpha)
else: # linear scale alphas initialized to ones
self.alpha = Parameter(torch.ones(in_features) * alpha)
self.beta = Parameter(torch.ones(in_features) * alpha)
self.alpha.requires_grad = alpha_trainable
self.beta.requires_grad = alpha_trainable
self.no_div_by_zero = 0.000000001
def forward(self, x):
'''
Forward pass of the function.
Applies the function to the input elementwise.
SnakeBeta ∶= x + 1/b * sin^2 (xa)
'''
alpha = self.alpha.unsqueeze(0).unsqueeze(-1) # line up with x to [B, C, T]
beta = self.beta.unsqueeze(0).unsqueeze(-1)
if self.alpha_logscale:
alpha = torch.exp(alpha)
beta = torch.exp(beta)
x = x + (1.0 / (beta + self.no_div_by_zero)) * pow(sin(x * alpha), 2)
return x
-6
View File
@@ -1,6 +0,0 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
from .filter import *
from .resample import *
from .act import *
-28
View File
@@ -1,28 +0,0 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
import torch.nn as nn
from .resample import UpSample1d, DownSample1d
class Activation1d(nn.Module):
def __init__(self,
activation,
up_ratio: int = 2,
down_ratio: int = 2,
up_kernel_size: int = 12,
down_kernel_size: int = 12):
super().__init__()
self.up_ratio = up_ratio
self.down_ratio = down_ratio
self.act = activation
self.upsample = UpSample1d(up_ratio, up_kernel_size)
self.downsample = DownSample1d(down_ratio, down_kernel_size)
# x: [B,C,T]
def forward(self, x):
x = self.upsample(x)
x = self.act(x)
x = self.downsample(x)
return x
-95
View File
@@ -1,95 +0,0 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
if 'sinc' in dir(torch):
sinc = torch.sinc
else:
# This code is adopted from adefossez's julius.core.sinc under the MIT License
# https://adefossez.github.io/julius/julius/core.html
# LICENSE is in incl_licenses directory.
def sinc(x: torch.Tensor):
"""
Implementation of sinc, i.e. sin(pi * x) / (pi * x)
__Warning__: Different to julius.sinc, the input is multiplied by `pi`!
"""
return torch.where(x == 0,
torch.tensor(1., device=x.device, dtype=x.dtype),
torch.sin(math.pi * x) / math.pi / x)
# This code is adopted from adefossez's julius.lowpass.LowPassFilters under the MIT License
# https://adefossez.github.io/julius/julius/lowpass.html
# LICENSE is in incl_licenses directory.
def kaiser_sinc_filter1d(cutoff, half_width, kernel_size): # return filter [1,1,kernel_size]
even = (kernel_size % 2 == 0)
half_size = kernel_size // 2
#For kaiser window
delta_f = 4 * half_width
A = 2.285 * (half_size - 1) * math.pi * delta_f + 7.95
if A > 50.:
beta = 0.1102 * (A - 8.7)
elif A >= 21.:
beta = 0.5842 * (A - 21)**0.4 + 0.07886 * (A - 21.)
else:
beta = 0.
window = torch.kaiser_window(kernel_size, beta=beta, periodic=False)
# ratio = 0.5/cutoff -> 2 * cutoff = 1 / ratio
if even:
time = (torch.arange(-half_size, half_size) + 0.5)
else:
time = torch.arange(kernel_size) - half_size
if cutoff == 0:
filter_ = torch.zeros_like(time)
else:
filter_ = 2 * cutoff * window * sinc(2 * cutoff * time)
# Normalize filter to have sum = 1, otherwise we will have a small leakage
# of the constant component in the input signal.
filter_ /= filter_.sum()
filter = filter_.view(1, 1, kernel_size)
return filter
class LowPassFilter1d(nn.Module):
def __init__(self,
cutoff=0.5,
half_width=0.6,
stride: int = 1,
padding: bool = True,
padding_mode: str = 'replicate',
kernel_size: int = 12):
# kernel_size should be even number for stylegan3 setup,
# in this implementation, odd number is also possible.
super().__init__()
if cutoff < -0.:
raise ValueError("Minimum cutoff must be larger than zero.")
if cutoff > 0.5:
raise ValueError("A cutoff above 0.5 does not make sense.")
self.kernel_size = kernel_size
self.even = (kernel_size % 2 == 0)
self.pad_left = kernel_size // 2 - int(self.even)
self.pad_right = kernel_size // 2
self.stride = stride
self.padding = padding
self.padding_mode = padding_mode
filter = kaiser_sinc_filter1d(cutoff, half_width, kernel_size)
self.register_buffer("filter", filter)
#input [B, C, T]
def forward(self, x):
_, C, _ = x.shape
if self.padding:
x = F.pad(x, (self.pad_left, self.pad_right),
mode=self.padding_mode)
out = F.conv1d(x, self.filter.expand(C, -1, -1),
stride=self.stride, groups=C)
return out
-49
View File
@@ -1,49 +0,0 @@
# Adapted from https://github.com/junjun3518/alias-free-torch under the Apache License 2.0
# LICENSE is in incl_licenses directory.
import torch.nn as nn
from torch.nn import functional as F
from .filter import LowPassFilter1d
from .filter import kaiser_sinc_filter1d
class UpSample1d(nn.Module):
def __init__(self, ratio=2, kernel_size=None):
super().__init__()
self.ratio = ratio
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.stride = ratio
self.pad = self.kernel_size // ratio - 1
self.pad_left = self.pad * self.stride + (self.kernel_size - self.stride) // 2
self.pad_right = self.pad * self.stride + (self.kernel_size - self.stride + 1) // 2
filter = kaiser_sinc_filter1d(cutoff=0.5 / ratio,
half_width=0.6 / ratio,
kernel_size=self.kernel_size)
self.register_buffer("filter", filter)
# x: [B, C, T]
def forward(self, x):
_, C, _ = x.shape
x = F.pad(x, (self.pad, self.pad), mode='replicate')
x = self.ratio * F.conv_transpose1d(
x, self.filter.expand(C, -1, -1), stride=self.stride, groups=C)
x = x[..., self.pad_left:-self.pad_right]
return x
class DownSample1d(nn.Module):
def __init__(self, ratio=2, kernel_size=None):
super().__init__()
self.ratio = ratio
self.kernel_size = int(6 * ratio // 2) * 2 if kernel_size is None else kernel_size
self.lowpass = LowPassFilter1d(cutoff=0.5 / ratio,
half_width=0.6 / ratio,
stride=ratio,
kernel_size=self.kernel_size)
def forward(self, x):
xx = self.lowpass(x)
return xx
-62
View File
@@ -1,62 +0,0 @@
import torch
import torch.nn as nn
from types import SimpleNamespace
from .models import BigVGANVocoder
from comfy.utils import load_torch_file
# BigVGAN vocoder configuration
_bigvgan_vocoder_config = {
'resblock': '1',
'num_gpus': 0,
'batch_size': 64,
'num_mels': 80,
'learning_rate': 0.0001,
'adam_b1': 0.8,
'adam_b2': 0.99,
'lr_decay': 0.999,
'seed': 1234,
'upsample_rates': [4, 4, 2, 2, 2, 2],
'upsample_kernel_sizes': [8, 8, 4, 4, 4, 4],
'upsample_initial_channel': 1536,
'resblock_kernel_sizes': [3, 7, 11],
'resblock_dilation_sizes': [
[1, 3, 5],
[1, 3, 5],
[1, 3, 5]
],
'activation': 'snakebeta',
'snake_logscale': True,
'resolutions': [
[1024, 120, 600],
[2048, 240, 1200],
[512, 50, 240]
],
'mpd_reshapes': [2, 3, 5, 7, 11],
'use_spectral_norm': False,
'discriminator_channel_mult': 1,
}
class BigVGAN(nn.Module):
def __init__(self, ckpt_path):
super().__init__()
# Convert dictionary to namespace object for attribute access
vocoder_cfg = SimpleNamespace(**_bigvgan_vocoder_config)
self.vocoder = BigVGANVocoder(vocoder_cfg).eval()
vocoder_ckpt = load_torch_file(ckpt_path)
self.vocoder.load_state_dict(vocoder_ckpt)
self.weight_norm_removed = False
self.remove_weight_norm()
@torch.inference_mode()
def forward(self, x):
assert self.weight_norm_removed, 'call remove_weight_norm() before inference'
return self.vocoder(x)
def remove_weight_norm(self):
self.vocoder.remove_weight_norm()
self.weight_norm_removed = True
return self
-21
View File
@@ -1,21 +0,0 @@
MIT License
Copyright (c) 2020 Jungil Kong
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
-21
View File
@@ -1,21 +0,0 @@
MIT License
Copyright (c) 2020 Edward Dixon
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
-201
View File
@@ -1,201 +0,0 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. For the purposes
of this License, Derivative Works shall not include works that remain
separable from, or merely link (or bind by name) to the interfaces of,
the Work and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
-29
View File
@@ -1,29 +0,0 @@
BSD 3-Clause License
Copyright (c) 2019, Seungwon Park 박승원
All rights reserved.
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met:
1. Redistributions of source code must retain the above copyright notice, this
list of conditions and the following disclaimer.
2. Redistributions in binary form must reproduce the above copyright notice,
this list of conditions and the following disclaimer in the documentation
and/or other materials provided with the distribution.
3. Neither the name of the copyright holder nor the names of its
contributors may be used to endorse or promote products derived from
this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
-16
View File
@@ -1,16 +0,0 @@
Copyright 2020 Alexandre Défossez
Permission is hereby granted, free of charge, to any person obtaining a copy of this software and
associated documentation files (the "Software"), to deal in the Software without restriction,
including without limitation the rights to use, copy, modify, merge, publish, distribute,
sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all copies or
substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT
NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM,
DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
-255
View File
@@ -1,255 +0,0 @@
# Copyright (c) 2022 NVIDIA CORPORATION.
# Licensed under the MIT license.
# Adapted from https://github.com/jik876/hifi-gan under the MIT license.
# LICENSE is in incl_licenses directory.
import torch
import torch.nn as nn
from torch.nn import Conv1d, ConvTranspose1d
from torch.nn.utils.parametrizations import weight_norm
from torch.nn.utils.parametrize import remove_parametrizations
from . import activations
from .alias_free_torch import *
from .utils import get_padding, init_weights
LRELU_SLOPE = 0.1
class AMPBlock1(torch.nn.Module):
def __init__(self, h, channels, kernel_size=3, dilation=(1, 3, 5), activation=None):
super(AMPBlock1, self).__init__()
self.h = h
self.convs1 = nn.ModuleList([
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1]))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[2],
padding=get_padding(kernel_size, dilation[2])))
])
self.convs1.apply(init_weights)
self.convs2 = nn.ModuleList([
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1)))
])
self.convs2.apply(init_weights)
self.num_layers = len(self.convs1) + len(self.convs2) # total number of conv layers
if activation == 'snake': # periodic nonlinearity with snake function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.Snake(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
elif activation == 'snakebeta': # periodic nonlinearity with snakebeta function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.SnakeBeta(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
else:
raise NotImplementedError(
"activation incorrectly specified. check the config file and look for 'activation'."
)
def forward(self, x):
acts1, acts2 = self.activations[::2], self.activations[1::2]
for c1, c2, a1, a2 in zip(self.convs1, self.convs2, acts1, acts2):
xt = a1(x)
xt = c1(xt)
xt = a2(xt)
xt = c2(xt)
x = xt + x
return x
def remove_weight_norm(self):
for l in self.convs1:
remove_parametrizations(l, 'weight')
for l in self.convs2:
remove_parametrizations(l, 'weight')
class AMPBlock2(torch.nn.Module):
def __init__(self, h, channels, kernel_size=3, dilation=(1, 3), activation=None):
super(AMPBlock2, self).__init__()
self.h = h
self.convs = nn.ModuleList([
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]))),
weight_norm(
Conv1d(channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1])))
])
self.convs.apply(init_weights)
self.num_layers = len(self.convs) # total number of conv layers
if activation == 'snake': # periodic nonlinearity with snake function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.Snake(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
elif activation == 'snakebeta': # periodic nonlinearity with snakebeta function and anti-aliasing
self.activations = nn.ModuleList([
Activation1d(
activation=activations.SnakeBeta(channels, alpha_logscale=h.snake_logscale))
for _ in range(self.num_layers)
])
else:
raise NotImplementedError(
"activation incorrectly specified. check the config file and look for 'activation'."
)
def forward(self, x):
for c, a in zip(self.convs, self.activations):
xt = a(x)
xt = c(xt)
x = xt + x
return x
def remove_weight_norm(self):
for l in self.convs:
remove_parametrizations(l, 'weight')
class BigVGANVocoder(torch.nn.Module):
# this is our main BigVGAN model. Applies anti-aliased periodic activation for resblocks.
def __init__(self, h):
super().__init__()
self.h = h
self.num_kernels = len(h.resblock_kernel_sizes)
self.num_upsamples = len(h.upsample_rates)
# pre conv
self.conv_pre = weight_norm(Conv1d(h.num_mels, h.upsample_initial_channel, 7, 1, padding=3))
# define which AMPBlock to use. BigVGAN uses AMPBlock1 as default
resblock = AMPBlock1 if h.resblock == '1' else AMPBlock2
# transposed conv-based upsamplers. does not apply anti-aliasing
self.ups = nn.ModuleList()
for i, (u, k) in enumerate(zip(h.upsample_rates, h.upsample_kernel_sizes)):
self.ups.append(
nn.ModuleList([
weight_norm(
ConvTranspose1d(h.upsample_initial_channel // (2**i),
h.upsample_initial_channel // (2**(i + 1)),
k,
u,
padding=(k - u) // 2))
]))
# residual blocks using anti-aliased multi-periodicity composition modules (AMP)
self.resblocks = nn.ModuleList()
for i in range(len(self.ups)):
ch = h.upsample_initial_channel // (2**(i + 1))
for j, (k, d) in enumerate(zip(h.resblock_kernel_sizes, h.resblock_dilation_sizes)):
self.resblocks.append(resblock(h, ch, k, d, activation=h.activation))
# post conv
if h.activation == "snake": # periodic nonlinearity with snake function and anti-aliasing
activation_post = activations.Snake(ch, alpha_logscale=h.snake_logscale)
self.activation_post = Activation1d(activation=activation_post)
elif h.activation == "snakebeta": # periodic nonlinearity with snakebeta function and anti-aliasing
activation_post = activations.SnakeBeta(ch, alpha_logscale=h.snake_logscale)
self.activation_post = Activation1d(activation=activation_post)
else:
raise NotImplementedError(
"activation incorrectly specified. check the config file and look for 'activation'."
)
self.conv_post = weight_norm(Conv1d(ch, 1, 7, 1, padding=3))
# weight initialization
for i in range(len(self.ups)):
self.ups[i].apply(init_weights)
self.conv_post.apply(init_weights)
def forward(self, x):
# pre conv
x = self.conv_pre(x)
for i in range(self.num_upsamples):
# upsampling
for i_up in range(len(self.ups[i])):
x = self.ups[i][i_up](x)
# AMP blocks
xs = None
for j in range(self.num_kernels):
if xs is None:
xs = self.resblocks[i * self.num_kernels + j](x)
else:
xs += self.resblocks[i * self.num_kernels + j](x)
x = xs / self.num_kernels
# post conv
x = self.activation_post(x)
x = self.conv_post(x)
x = torch.tanh(x)
return x
def remove_weight_norm(self):
print('Removing weight norm...')
for l in self.ups:
for l_i in l:
remove_parametrizations(l_i, 'weight')
for l in self.resblocks:
l.remove_weight_norm()
remove_parametrizations(self.conv_pre, 'weight')
remove_parametrizations(self.conv_post, 'weight')
-20
View File
@@ -1,20 +0,0 @@
# Adapted from https://github.com/jik876/hifi-gan under the MIT license.
# LICENSE is in incl_licenses directory.
from torch.nn.utils.parametrizations import weight_norm
def init_weights(m, mean=0.0, std=0.01):
classname = m.__class__.__name__
if classname.find("Conv") != -1:
m.weight.data.normal_(mean, std)
def apply_weight_norm(m):
classname = m.__class__.__name__
if classname.find("Conv") != -1:
weight_norm(m)
def get_padding(kernel_size, dilation=1):
return int((kernel_size * dilation - dilation) / 2)
-212
View File
@@ -1,212 +0,0 @@
# Reference: # https://github.com/bytedance/Make-An-Audio-2
from typing import Literal
import torch
import torch.nn as nn
import numpy as np
# following is from librosa
def hz_to_mel(frequencies, *, htk = False):
frequencies = np.asanyarray(frequencies)
if htk:
mels: np.ndarray = 2595.0 * np.log10(1.0 + frequencies / 700.0)
return mels
# Fill in the linear part
f_min = 0.0
f_sp = 200.0 / 3
mels = (frequencies - f_min) / f_sp
# Fill in the log-scale part
min_log_hz = 1000.0 # beginning of log region (Hz)
min_log_mel = (min_log_hz - f_min) / f_sp # same (Mels)
logstep = np.log(6.4) / 27.0 # step size for log region
if frequencies.ndim:
# If we have array data, vectorize
log_t = frequencies >= min_log_hz
mels[log_t] = min_log_mel + np.log(frequencies[log_t] / min_log_hz) / logstep
elif frequencies >= min_log_hz:
# If we have scalar data, heck directly
mels = min_log_mel + np.log(frequencies / min_log_hz) / logstep
return mels
def mel_to_hz(mels, *, htk = False):
mels = np.asanyarray(mels)
if htk:
return 700.0 * (10.0 ** (mels / 2595.0) - 1.0)
# Fill in the linear scale
f_min = 0.0
f_sp = 200.0 / 3
freqs = f_min + f_sp * mels
# And now the nonlinear scale
min_log_hz = 1000.0 # beginning of log region (Hz)
min_log_mel = (min_log_hz - f_min) / f_sp # same (Mels)
logstep = np.log(6.4) / 27.0 # step size for log region
if mels.ndim:
# If we have vector data, vectorize
log_t = mels >= min_log_mel
freqs[log_t] = min_log_hz * np.exp(logstep * (mels[log_t] - min_log_mel))
elif mels >= min_log_mel:
# If we have scalar data, check directly
freqs = min_log_hz * np.exp(logstep * (mels - min_log_mel))
return freqs
def mel_frequencies(n_mels = 128, *, fmin = 0.0, fmax = 11025.0, htk = False):
min_mel = hz_to_mel(fmin, htk=htk)
max_mel = hz_to_mel(fmax, htk=htk)
mels = np.linspace(min_mel, max_mel, n_mels)
hz: np.ndarray = mel_to_hz(mels, htk=htk)
return hz
def librosa_mel_fn(
*,
sr: float,
n_fft: int,
n_mels: int = 128,
fmin: float = 0.0,
fmax = None,
htk = False,
norm = "slaney",
dtype = np.float32,
) -> np.ndarray:
if fmax is None:
fmax = float(sr) / 2
# Initialize the weights
n_mels = int(n_mels)
weights = np.zeros((n_mels, int(1 + n_fft // 2)), dtype=dtype)
# Center freqs of each FFT bin
fftfreqs = np.fft.rfftfreq(n=n_fft, d=1.0 / sr)
# 'Center freqs' of mel bands - uniformly spaced between limits
mel_f = mel_frequencies(n_mels + 2, fmin=fmin, fmax=fmax, htk=htk)
fdiff = np.diff(mel_f)
ramps = np.subtract.outer(mel_f, fftfreqs)
for i in range(n_mels):
# lower and upper slopes for all bins
lower = -ramps[i] / fdiff[i]
upper = ramps[i + 2] / fdiff[i + 1]
# .. then intersect them with each other and zero
weights[i] = np.maximum(0, np.minimum(lower, upper))
# Slaney-style mel is scaled to be approx constant energy per channel
enorm = 2.0 / (mel_f[2 : n_mels + 2] - mel_f[:n_mels])
weights *= enorm[:, np.newaxis]
return weights
def dynamic_range_compression_torch(x, C=1, clip_val=1e-5, *, norm_fn):
return norm_fn(torch.clamp(x, min=clip_val) * C)
def spectral_normalize_torch(magnitudes, norm_fn):
output = dynamic_range_compression_torch(magnitudes, norm_fn=norm_fn)
return output
class MelConverter(nn.Module):
def __init__(
self,
*,
sampling_rate: float,
n_fft: int,
num_mels: int,
hop_size: int,
win_size: int,
fmin: float,
fmax: float,
norm_fn,
):
super().__init__()
self.sampling_rate = sampling_rate
self.n_fft = n_fft
self.num_mels = num_mels
self.hop_size = hop_size
self.win_size = win_size
self.fmin = fmin
self.fmax = fmax
self.norm_fn = norm_fn
mel = librosa_mel_fn(sr=self.sampling_rate,
n_fft=self.n_fft,
n_mels=self.num_mels,
fmin=self.fmin,
fmax=self.fmax)
mel_basis = torch.from_numpy(mel).float()
hann_window = torch.hann_window(self.win_size)
self.register_buffer('mel_basis', mel_basis)
self.register_buffer('hann_window', hann_window)
@property
def device(self):
return self.mel_basis.device
def forward(self, waveform: torch.Tensor, center: bool = False) -> torch.Tensor:
waveform = waveform.clamp(min=-1., max=1.).to(self.device)
waveform = torch.nn.functional.pad(
waveform.unsqueeze(1),
[int((self.n_fft - self.hop_size) / 2),
int((self.n_fft - self.hop_size) / 2)],
mode='reflect')
waveform = waveform.squeeze(1)
spec = torch.stft(waveform,
self.n_fft,
hop_length=self.hop_size,
win_length=self.win_size,
window=self.hann_window,
center=center,
pad_mode='reflect',
normalized=False,
onesided=True,
return_complex=True)
spec = torch.view_as_real(spec)
spec = torch.sqrt(spec.pow(2).sum(-1) + (1e-9)).float()
spec = torch.matmul(self.mel_basis, spec)
spec = spectral_normalize_torch(spec, self.norm_fn)
return spec
def get_mel_converter(mode: Literal['16k', '44k']) -> MelConverter:
if mode == '16k':
return MelConverter(sampling_rate=16_000,
n_fft=1024,
num_mels=80,
hop_size=256,
win_size=1024,
fmin=0,
fmax=8_000,
norm_fn=torch.log10)
elif mode == '44k':
return MelConverter(sampling_rate=44_100,
n_fft=2048,
num_mels=128,
hop_size=512,
win_size=2048,
fmin=0,
fmax=44100 / 2,
norm_fn=torch.log)
else:
raise ValueError(f'Unknown mode: {mode}')
-267
View File
@@ -1,267 +0,0 @@
import torch
import torch.nn as nn
import folder_paths
import os
from .mel_converter import get_mel_converter
from .vae.autoencoder import AutoEncoderModule
from .vae.distributions import DiagonalGaussianDistribution
import torchaudio
from ..utils import log
from comfy import model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class FeaturesUtils(nn.Module):
def __init__(
self,
*,
tod_vae_ckpt: str,
bigvgan_vocoder_ckpt = None,
mode=['16k', '44k'],
need_vae_encoder: bool = True,
):
super().__init__()
self.mel_converter = get_mel_converter(mode)
self.tod = AutoEncoderModule(vae_ckpt_path=tod_vae_ckpt,
vocoder_ckpt_path=bigvgan_vocoder_ckpt,
mode=mode,
need_vae_encoder=need_vae_encoder)
def encode_audio(self, x) -> DiagonalGaussianDistribution:
assert self.tod is not None, 'VAE is not loaded'
# x: (B * L)
mel = self.mel_converter(x)
dist = self.tod.encode(mel)
return dist
def vocode(self, mel: torch.Tensor) -> torch.Tensor:
assert self.tod is not None, 'VAE is not loaded'
return self.tod.vocode(mel)
def decode(self, z: torch.Tensor) -> torch.Tensor:
assert self.tod is not None, 'VAE is not loaded'
return self.tod.decode(z)
@property
def device(self):
return next(self.parameters()).device
@property
def dtype(self):
return next(self.parameters()).dtype
def wrapped_decode(self, z):
with torch.amp.autocast('cuda', dtype=self.dtype):
mel_decoded = self.decode(z)
audio = self.vocode(mel_decoded)
return audio
def wrapped_encode(self, audio):
with torch.amp.autocast('cuda', dtype=self.dtype):
dist = self.encode_audio(audio)
return dist.mean
if not "mmaudio" in folder_paths.folder_names_and_paths:
folder_paths.add_model_folder_path("mmaudio", os.path.join(folder_paths.models_dir, "mmaudio"))
class OviMMAudioVAELoader:
"""Loads MMAudio VAE for audio encoding/decoding in Ovi"""
@classmethod
def INPUT_TYPES(s):
s.vae_files = folder_paths.get_filename_list("vae")
s.mmaudio_files = folder_paths.get_filename_list("mmaudio")
s.all_files = s.vae_files + s.mmaudio_files
return {
"required": {
"vae": (s.all_files, {"tooltip": "MMAudio VAE 16k (v1-16.pth) model from models/vae or models/mmaudio"}),
"vocoder": (s.all_files, {"tooltip": "BigVGAN vocoder (best_netG.pt) from models/vae or models/mmaudio"}),
"precision": (["bf16", "fp16", "fp32"], {"default": "bf16"}),
}
}
RETURN_TYPES = ("MMAUDIOVAE",)
RETURN_NAMES = ("mmaudio_vae",)
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper/Ovi"
DESCRIPTION = "Loads MMAudio VAE for Ovi audio generation"
def loadmodel(self, vae, vocoder, precision):
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
vae_path = folder_paths.get_full_path("vae", vae) if vae in self.vae_files else folder_paths.get_full_path("mmaudio", vae)
vocoder_path = folder_paths.get_full_path("vae", vocoder) if vocoder in self.vae_files else folder_paths.get_full_path("mmaudio", vocoder)
vae = FeaturesUtils(
tod_vae_ckpt=vae_path,
bigvgan_vocoder_ckpt=vocoder_path,
mode='16k',
need_vae_encoder=True
)
vae.to(device=offload_device, dtype=dtype)
vae.eval()
return (vae,)
class WanVideoDecodeOviAudio:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"mmaudio_vae": ("MMAUDIOVAE",),
"samples": ("LATENT",),
}
}
RETURN_TYPES = ("AUDIO",)
RETURN_NAMES = ("audio",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, mmaudio_vae, samples):
mm.soft_empty_cache()
audio_latents = samples.get("latent_ovi_audio", None)
if audio_latents is None:
raise ValueError("No Ovi audio latents found in input samples")
mmaudio_vae.to(device)
waveform = mmaudio_vae.wrapped_decode(audio_latents.to(device=device, dtype=mmaudio_vae.dtype))
audio = {"waveform": waveform.cpu().float(), "sample_rate": 16000}
mmaudio_vae.to(offload_device)
mm.soft_empty_cache()
return (audio,)
class WanVideoEncodeOviAudio:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"mmaudio_vae": ("MMAUDIOVAE",),
"audio": ("AUDIO",),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, mmaudio_vae, audio):
mmaudio_vae.to(device)
waveform = audio.get("waveform", None)
sample_rate = audio.get("sample_rate", None)
if sample_rate != 16000:
waveform = torchaudio.functional.resample(waveform, sample_rate, 16000)
waveform = waveform.to(device=device, dtype=mmaudio_vae.dtype)[0][0].unsqueeze(0)
samples = mmaudio_vae.wrapped_encode(waveform)
mmaudio_vae.to(offload_device)
mm.soft_empty_cache()
return ({"latent_ovi_audio": samples},)
class WanVideoAddOviAudioToLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"original_samples": ("LATENT",),
"audio_samples": ("LATENT",),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, original_samples, audio_samples):
samples = original_samples.copy()
samples.update(audio_samples)
return (samples,)
class WanVideoEmptyMMAudioLatents:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"length": ("INT", {"default": 157, "min": 1, "max": 10000, "step": 1, "tooltip": "Length of the audio latent sequence"}),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "decode"
CATEGORY = "WanVideoWrapper/Ovi"
def decode(self, length):
audio_latents = torch.zeros((length, 20), device=torch.device("cpu"), dtype=torch.float32) # 1, l c -> l, c
return ({"latent_ovi_audio": audio_latents},)
class WanVideoOviCFG:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"original_text_embeds": ("WANVIDEOTEXTEMBEDS",),
"ovi_audio_cfg": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step": 0.01}),
},
"optional": {
"ovi_negative_text_embeds": ("WANVIDEOTEXTEMBEDS",),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper/Ovi"
DESCRIPTION = "Adds Ovi negative text embeddings and audio CFG scale to the text embeddings dictionary"
def process(self, original_text_embeds, ovi_audio_cfg, ovi_negative_text_embeds=None):
negative_text_embeds = None
if ovi_negative_text_embeds is not None:
negative_text_embeds = ovi_negative_text_embeds.get("prompt_embeds", None)
if negative_text_embeds is None:
negative_text_embeds = original_text_embeds["prompt_embeds"]
log.info("WanVideoOviCFG: Ovi negative text embeddings not provided, using original prompt embeddings as negative embeddings")
else:
log.info("WanVideoOviCFG: Using provided Ovi audio negative text embeddings")
log.info("WanVideoOviCFG: negative text embedding shape: {}".format(negative_text_embeds[0].shape))
prompt_embeds_dict_copy = original_text_embeds.copy()
prompt_embeds_dict_copy.update({
"ovi_negative_prompt_embeds": negative_text_embeds,
"ovi_audio_cfg": ovi_audio_cfg,
})
return (prompt_embeds_dict_copy,)
NODE_CLASS_MAPPINGS = {
"OviMMAudioVAELoader": OviMMAudioVAELoader,
"WanVideoDecodeOviAudio": WanVideoDecodeOviAudio,
"WanVideoEncodeOviAudio": WanVideoEncodeOviAudio,
"WanVideoOviCFG": WanVideoOviCFG,
"WanVideoAddOviAudioToLatents": WanVideoAddOviAudioToLatents,
"WanVideoEmptyMMAudioLatents": WanVideoEmptyMMAudioLatents,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"OviMMAudioVAELoader": "Ovi MMAudio VAE Loader",
"WanVideoDecodeOviAudio": "WanVideo Decode Ovi Audio",
"WanVideoEncodeOviAudio": "WanVideo Encode Ovi Audio",
"WanVideoOviCFG": "WanVideo Ovi CFG",
"WanVideoAddOviAudioToLatents": "WanVideo Add MMAudio To Latents",
"WanVideoEmptyMMAudioLatents": "WanVideo Empty MMAudio Latents",
}
-54
View File
@@ -1,54 +0,0 @@
from typing import Literal, Optional
import torch
import torch.nn as nn
from .vae import VAE, get_my_vae
from .distributions import DiagonalGaussianDistribution
from ..bigvgan import BigVGAN
from comfy.utils import load_torch_file
class AutoEncoderModule(nn.Module):
def __init__(self,
*,
vae_ckpt_path,
vocoder_ckpt_path: Optional[str] = None,
mode: Literal['16k', '44k'],
need_vae_encoder: bool = True):
super().__init__()
self.vae: VAE = get_my_vae(mode).eval()
#vae_state_dict = torch.load(vae_ckpt_path, weights_only=True, map_location='cpu')'
vae_state_dict = load_torch_file(vae_ckpt_path)
self.vae.load_state_dict(vae_state_dict)
self.vae.remove_weight_norm()
if mode == '16k':
assert vocoder_ckpt_path is not None
self.vocoder = BigVGAN(vocoder_ckpt_path).eval()
elif mode == '44k':
raise NotImplementedError("44k mode requires BigVGANv2 which is not currently supported in this environment.")
self.vocoder = BigVGANv2.from_pretrained('nvidia/bigvgan_v2_44khz_128band_512x',
use_cuda_kernel=False)
self.vocoder.remove_weight_norm()
else:
raise ValueError(f'Unknown mode: {mode}')
for param in self.parameters():
param.requires_grad = False
if not need_vae_encoder:
del self.vae.encoder
@torch.inference_mode()
def encode(self, x: torch.Tensor) -> DiagonalGaussianDistribution:
return self.vae.encode(x)
@torch.inference_mode()
def decode(self, z: torch.Tensor) -> torch.Tensor:
return self.vae.decode(z)
@torch.inference_mode()
def vocode(self, spec: torch.Tensor) -> torch.Tensor:
return self.vocoder(spec)
-45
View File
@@ -1,45 +0,0 @@
from typing import Optional
import torch
import numpy as np
class DiagonalGaussianDistribution:
def __init__(self, parameters, deterministic=False):
self.parameters = parameters
self.mean, self.logvar = torch.chunk(parameters, 2, dim=1)
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
self.deterministic = deterministic
self.std = torch.exp(0.5 * self.logvar)
self.var = torch.exp(self.logvar)
if self.deterministic:
self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device)
def sample(self, rng: Optional[torch.Generator] = None):
# x = self.mean + self.std * torch.randn(self.mean.shape).to(device=self.parameters.device)
r = torch.empty_like(self.mean).normal_(generator=rng)
x = self.mean + self.std * r
return x
def kl(self, other=None):
if self.deterministic:
return torch.Tensor([0.])
else:
if other is None:
return 0.5 * torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar
else:
return 0.5 * (torch.pow(self.mean - other.mean, 2) / other.var +
self.var / other.var - 1.0 - self.logvar + other.logvar)
def nll(self, sample, dims=[1, 2, 3]):
if self.deterministic:
return torch.Tensor([0.])
logtwopi = np.log(2.0 * np.pi)
return 0.5 * torch.sum(logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
dim=dims)
def mode(self):
return self.mean
-168
View File
@@ -1,168 +0,0 @@
# Copyright (c) 2024, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# This work is licensed under a Creative Commons
# Attribution-NonCommercial-ShareAlike 4.0 International License.
# You should have received a copy of the license along with this
# work. If not, see http://creativecommons.org/licenses/by-nc-sa/4.0/
"""Improved diffusion model architecture proposed in the paper
"Analyzing and Improving the Training Dynamics of Diffusion Models"."""
import numpy as np
import torch
#----------------------------------------------------------------------------
# Variant of constant() that inherits dtype and device from the given
# reference tensor by default.
_constant_cache = dict()
def constant(value, shape=None, dtype=None, device=None, memory_format=None):
value = np.asarray(value)
if shape is not None:
shape = tuple(shape)
if dtype is None:
dtype = torch.get_default_dtype()
if device is None:
device = torch.device('cpu')
if memory_format is None:
memory_format = torch.contiguous_format
key = (value.shape, value.dtype, value.tobytes(), shape, dtype, device, memory_format)
tensor = _constant_cache.get(key, None)
if tensor is None:
tensor = torch.as_tensor(value.copy(), dtype=dtype, device=device)
if shape is not None:
tensor, _ = torch.broadcast_tensors(tensor, torch.empty(shape))
tensor = tensor.contiguous(memory_format=memory_format)
_constant_cache[key] = tensor
return tensor
def const_like(ref, value, shape=None, dtype=None, device=None, memory_format=None):
if dtype is None:
dtype = ref.dtype
if device is None:
device = ref.device
return constant(value, shape=shape, dtype=dtype, device=device, memory_format=memory_format)
#----------------------------------------------------------------------------
# Normalize given tensor to unit magnitude with respect to the given
# dimensions. Default = all dimensions except the first.
def normalize(x, dim=None, eps=1e-4):
if dim is None:
dim = list(range(1, x.ndim))
norm = torch.linalg.vector_norm(x, dim=dim, keepdim=True, dtype=torch.float32)
norm = torch.add(eps, norm, alpha=np.sqrt(norm.numel() / x.numel()))
return x / norm.to(x.dtype)
class Normalize(torch.nn.Module):
def __init__(self, dim=None, eps=1e-4):
super().__init__()
self.dim = dim
self.eps = eps
def forward(self, x):
return normalize(x, dim=self.dim, eps=self.eps)
#----------------------------------------------------------------------------
# Upsample or downsample the given tensor with the given filter,
# or keep it as is.
def resample(x, f=[1, 1], mode='keep'):
if mode == 'keep':
return x
f = np.float32(f)
assert f.ndim == 1 and len(f) % 2 == 0
pad = (len(f) - 1) // 2
f = f / f.sum()
f = np.outer(f, f)[np.newaxis, np.newaxis, :, :]
f = const_like(x, f)
c = x.shape[1]
if mode == 'down':
return torch.nn.functional.conv2d(x,
f.tile([c, 1, 1, 1]),
groups=c,
stride=2,
padding=(pad, ))
assert mode == 'up'
return torch.nn.functional.conv_transpose2d(x, (f * 4).tile([c, 1, 1, 1]),
groups=c,
stride=2,
padding=(pad, ))
#----------------------------------------------------------------------------
# Magnitude-preserving SiLU (Equation 81).
def mp_silu(x):
return torch.nn.functional.silu(x) / 0.596
class MPSiLU(torch.nn.Module):
def forward(self, x):
return mp_silu(x)
#----------------------------------------------------------------------------
# Magnitude-preserving sum (Equation 88).
def mp_sum(a, b, t=0.5):
return a.lerp(b, t) / np.sqrt((1 - t)**2 + t**2)
#----------------------------------------------------------------------------
# Magnitude-preserving concatenation (Equation 103).
def mp_cat(a, b, dim=1, t=0.5):
Na = a.shape[dim]
Nb = b.shape[dim]
C = np.sqrt((Na + Nb) / ((1 - t)**2 + t**2))
wa = C / np.sqrt(Na) * (1 - t)
wb = C / np.sqrt(Nb) * t
return torch.cat([wa * a, wb * b], dim=dim)
#----------------------------------------------------------------------------
# Magnitude-preserving convolution or fully-connected layer (Equation 47)
# with force weight normalization (Equation 66).
class MPConv1D(torch.nn.Module):
def __init__(self, in_channels, out_channels, kernel_size):
super().__init__()
self.out_channels = out_channels
self.weight = torch.nn.Parameter(torch.randn(out_channels, in_channels, kernel_size))
self.weight_norm_removed = False
def forward(self, x, gain=1):
assert self.weight_norm_removed, 'call remove_weight_norm() before inference'
w = self.weight * gain
if w.ndim == 2:
return x @ w.t()
assert w.ndim == 3
return torch.nn.functional.conv1d(x, w, padding=(w.shape[-1] // 2, ))
def remove_weight_norm(self):
w = self.weight.to(torch.float32)
w = normalize(w) # traditional weight normalization
w = w / np.sqrt(w[0].numel())
w = w.to(self.weight.dtype)
self.weight.data.copy_(w)
self.weight_norm_removed = True
return self
-376
View File
@@ -1,376 +0,0 @@
import logging
from typing import Optional
import torch
import torch.nn as nn
from .edm2_utils import MPConv1D
from .vae_modules import (AttnBlock1D, Downsample1D, ResnetBlock1D,
Upsample1D, nonlinearity)
from .distributions import DiagonalGaussianDistribution
log = logging.getLogger()
DATA_MEAN_80D = [
-1.6058, -1.3676, -1.2520, -1.2453, -1.2078, -1.2224, -1.2419, -1.2439, -1.2922, -1.2927,
-1.3170, -1.3543, -1.3401, -1.3836, -1.3907, -1.3912, -1.4313, -1.4152, -1.4527, -1.4728,
-1.4568, -1.5101, -1.5051, -1.5172, -1.5623, -1.5373, -1.5746, -1.5687, -1.6032, -1.6131,
-1.6081, -1.6331, -1.6489, -1.6489, -1.6700, -1.6738, -1.6953, -1.6969, -1.7048, -1.7280,
-1.7361, -1.7495, -1.7658, -1.7814, -1.7889, -1.8064, -1.8221, -1.8377, -1.8417, -1.8643,
-1.8857, -1.8929, -1.9173, -1.9379, -1.9531, -1.9673, -1.9824, -2.0042, -2.0215, -2.0436,
-2.0766, -2.1064, -2.1418, -2.1855, -2.2319, -2.2767, -2.3161, -2.3572, -2.3954, -2.4282,
-2.4659, -2.5072, -2.5552, -2.6074, -2.6584, -2.7107, -2.7634, -2.8266, -2.8981, -2.9673
]
DATA_STD_80D = [
1.0291, 1.0411, 1.0043, 0.9820, 0.9677, 0.9543, 0.9450, 0.9392, 0.9343, 0.9297, 0.9276, 0.9263,
0.9242, 0.9254, 0.9232, 0.9281, 0.9263, 0.9315, 0.9274, 0.9247, 0.9277, 0.9199, 0.9188, 0.9194,
0.9160, 0.9161, 0.9146, 0.9161, 0.9100, 0.9095, 0.9145, 0.9076, 0.9066, 0.9095, 0.9032, 0.9043,
0.9038, 0.9011, 0.9019, 0.9010, 0.8984, 0.8983, 0.8986, 0.8961, 0.8962, 0.8978, 0.8962, 0.8973,
0.8993, 0.8976, 0.8995, 0.9016, 0.8982, 0.8972, 0.8974, 0.8949, 0.8940, 0.8947, 0.8936, 0.8939,
0.8951, 0.8956, 0.9017, 0.9167, 0.9436, 0.9690, 1.0003, 1.0225, 1.0381, 1.0491, 1.0545, 1.0604,
1.0761, 1.0929, 1.1089, 1.1196, 1.1176, 1.1156, 1.1117, 1.1070
]
DATA_MEAN_128D = [
-3.3462, -2.6723, -2.4893, -2.3143, -2.2664, -2.3317, -2.1802, -2.4006, -2.2357, -2.4597,
-2.3717, -2.4690, -2.5142, -2.4919, -2.6610, -2.5047, -2.7483, -2.5926, -2.7462, -2.7033,
-2.7386, -2.8112, -2.7502, -2.9594, -2.7473, -3.0035, -2.8891, -2.9922, -2.9856, -3.0157,
-3.1191, -2.9893, -3.1718, -3.0745, -3.1879, -3.2310, -3.1424, -3.2296, -3.2791, -3.2782,
-3.2756, -3.3134, -3.3509, -3.3750, -3.3951, -3.3698, -3.4505, -3.4509, -3.5089, -3.4647,
-3.5536, -3.5788, -3.5867, -3.6036, -3.6400, -3.6747, -3.7072, -3.7279, -3.7283, -3.7795,
-3.8259, -3.8447, -3.8663, -3.9182, -3.9605, -3.9861, -4.0105, -4.0373, -4.0762, -4.1121,
-4.1488, -4.1874, -4.2461, -4.3170, -4.3639, -4.4452, -4.5282, -4.6297, -4.7019, -4.7960,
-4.8700, -4.9507, -5.0303, -5.0866, -5.1634, -5.2342, -5.3242, -5.4053, -5.4927, -5.5712,
-5.6464, -5.7052, -5.7619, -5.8410, -5.9188, -6.0103, -6.0955, -6.1673, -6.2362, -6.3120,
-6.3926, -6.4797, -6.5565, -6.6511, -6.8130, -6.9961, -7.1275, -7.2457, -7.3576, -7.4663,
-7.6136, -7.7469, -7.8815, -8.0132, -8.1515, -8.3071, -8.4722, -8.7418, -9.3975, -9.6628,
-9.7671, -9.8863, -9.9992, -10.0860, -10.1709, -10.5418, -11.2795, -11.3861
]
DATA_STD_128D = [
2.3804, 2.4368, 2.3772, 2.3145, 2.2803, 2.2510, 2.2316, 2.2083, 2.1996, 2.1835, 2.1769, 2.1659,
2.1631, 2.1618, 2.1540, 2.1606, 2.1571, 2.1567, 2.1612, 2.1579, 2.1679, 2.1683, 2.1634, 2.1557,
2.1668, 2.1518, 2.1415, 2.1449, 2.1406, 2.1350, 2.1313, 2.1415, 2.1281, 2.1352, 2.1219, 2.1182,
2.1327, 2.1195, 2.1137, 2.1080, 2.1179, 2.1036, 2.1087, 2.1036, 2.1015, 2.1068, 2.0975, 2.0991,
2.0902, 2.1015, 2.0857, 2.0920, 2.0893, 2.0897, 2.0910, 2.0881, 2.0925, 2.0873, 2.0960, 2.0900,
2.0957, 2.0958, 2.0978, 2.0936, 2.0886, 2.0905, 2.0845, 2.0855, 2.0796, 2.0840, 2.0813, 2.0817,
2.0838, 2.0840, 2.0917, 2.1061, 2.1431, 2.1976, 2.2482, 2.3055, 2.3700, 2.4088, 2.4372, 2.4609,
2.4731, 2.4847, 2.5072, 2.5451, 2.5772, 2.6147, 2.6529, 2.6596, 2.6645, 2.6726, 2.6803, 2.6812,
2.6899, 2.6916, 2.6931, 2.6998, 2.7062, 2.7262, 2.7222, 2.7158, 2.7041, 2.7485, 2.7491, 2.7451,
2.7485, 2.7233, 2.7297, 2.7233, 2.7145, 2.6958, 2.6788, 2.6439, 2.6007, 2.4786, 2.2469, 2.1877,
2.1392, 2.0717, 2.0107, 1.9676, 1.9140, 1.7102, 0.9101, 0.7164
]
class VAE(nn.Module):
def __init__(
self,
*,
data_dim: int,
embed_dim: int,
hidden_dim: int,
):
super().__init__()
if data_dim == 80:
data_mean = torch.tensor(DATA_MEAN_80D, dtype=torch.float32)
data_std = torch.tensor(DATA_STD_80D, dtype=torch.float32)
elif data_dim == 128:
data_mean = torch.tensor(DATA_MEAN_128D, dtype=torch.float32)
data_std = torch.tensor(DATA_STD_128D, dtype=torch.float32)
else:
raise ValueError(f"Unsupported data_dim={data_dim}, expected 80 or 128")
# match old shape: (1, channels, 1)
data_mean = data_mean.view(1, -1, 1)
data_std = data_std.view(1, -1, 1)
# register as buffers so they move with .to(device) / .cuda()
self.register_buffer("data_mean", data_mean)
self.register_buffer("data_std", data_std)
self.encoder = Encoder1D(
dim=hidden_dim,
ch_mult=(1, 2, 4),
num_res_blocks=2,
attn_layers=[3],
down_layers=[0],
in_dim=data_dim,
embed_dim=embed_dim,
)
self.decoder = Decoder1D(
dim=hidden_dim,
ch_mult=(1, 2, 4),
num_res_blocks=2,
attn_layers=[3],
down_layers=[0],
in_dim=data_dim,
out_dim=data_dim,
embed_dim=embed_dim,
)
self.embed_dim = embed_dim
# self.quant_conv = nn.Conv1d(2 * embed_dim, 2 * embed_dim, 1)
# self.post_quant_conv = nn.Conv1d(embed_dim, embed_dim, 1)
self.initialize_weights()
def initialize_weights(self):
pass
def encode(self, x: torch.Tensor, normalize: bool = True) -> DiagonalGaussianDistribution:
if normalize:
x = self.normalize(x)
moments = self.encoder(x)
posterior = DiagonalGaussianDistribution(moments)
return posterior
def decode(self, z: torch.Tensor, unnormalize: bool = True) -> torch.Tensor:
dec = self.decoder(z)
if unnormalize:
dec = self.unnormalize(dec)
return dec
def normalize(self, x: torch.Tensor) -> torch.Tensor:
return (x - self.data_mean) / self.data_std
def unnormalize(self, x: torch.Tensor) -> torch.Tensor:
return x * self.data_std + self.data_mean
def forward(
self,
x: torch.Tensor,
sample_posterior: bool = True,
rng: Optional[torch.Generator] = None,
normalize: bool = True,
unnormalize: bool = True,
) -> tuple[torch.Tensor, DiagonalGaussianDistribution]:
posterior = self.encode(x, normalize=normalize)
if sample_posterior:
z = posterior.sample(rng)
else:
z = posterior.mode()
dec = self.decode(z, unnormalize=unnormalize)
return dec, posterior
def load_weights(self, src_dict) -> None:
self.load_state_dict(src_dict, strict=True)
@property
def device(self) -> torch.device:
return next(self.parameters()).device
def get_last_layer(self):
return self.decoder.conv_out.weight
def remove_weight_norm(self):
for name, m in self.named_modules():
if isinstance(m, MPConv1D):
m.remove_weight_norm()
log.debug(f"Removed weight norm from {name}")
return self
class Encoder1D(nn.Module):
def __init__(self,
*,
dim: int,
ch_mult: tuple[int] = (1, 2, 4, 8),
num_res_blocks: int,
attn_layers: list[int] = [],
down_layers: list[int] = [],
resamp_with_conv: bool = True,
in_dim: int,
embed_dim: int,
double_z: bool = True,
kernel_size: int = 3,
clip_act: float = 256.0):
super().__init__()
self.dim = dim
self.num_layers = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.in_channels = in_dim
self.clip_act = clip_act
self.down_layers = down_layers
self.attn_layers = attn_layers
self.conv_in = MPConv1D(in_dim, self.dim, kernel_size=kernel_size)
in_ch_mult = (1, ) + tuple(ch_mult)
self.in_ch_mult = in_ch_mult
# downsampling
self.down = nn.ModuleList()
for i_level in range(self.num_layers):
block = nn.ModuleList()
attn = nn.ModuleList()
block_in = dim * in_ch_mult[i_level]
block_out = dim * ch_mult[i_level]
for i_block in range(self.num_res_blocks):
block.append(
ResnetBlock1D(in_dim=block_in,
out_dim=block_out,
kernel_size=kernel_size,
use_norm=True))
block_in = block_out
if i_level in attn_layers:
attn.append(AttnBlock1D(block_in))
down = nn.Module()
down.block = block
down.attn = attn
if i_level in down_layers:
down.downsample = Downsample1D(block_in, resamp_with_conv)
self.down.append(down)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock1D(in_dim=block_in,
out_dim=block_in,
kernel_size=kernel_size,
use_norm=True)
self.mid.attn_1 = AttnBlock1D(block_in)
self.mid.block_2 = ResnetBlock1D(in_dim=block_in,
out_dim=block_in,
kernel_size=kernel_size,
use_norm=True)
# end
self.conv_out = MPConv1D(block_in,
2 * embed_dim if double_z else embed_dim,
kernel_size=kernel_size)
self.learnable_gain = nn.Parameter(torch.zeros([]))
def forward(self, x):
# downsampling
hs = [self.conv_in(x)]
for i_level in range(self.num_layers):
for i_block in range(self.num_res_blocks):
h = self.down[i_level].block[i_block](hs[-1])
if len(self.down[i_level].attn) > 0:
h = self.down[i_level].attn[i_block](h)
h = h.clamp(-self.clip_act, self.clip_act)
hs.append(h)
if i_level in self.down_layers:
hs.append(self.down[i_level].downsample(hs[-1]))
# middle
h = hs[-1]
h = self.mid.block_1(h)
h = self.mid.attn_1(h)
h = self.mid.block_2(h)
h = h.clamp(-self.clip_act, self.clip_act)
# end
h = nonlinearity(h)
h = self.conv_out(h, gain=(self.learnable_gain + 1))
return h
class Decoder1D(nn.Module):
def __init__(self,
*,
dim: int,
out_dim: int,
ch_mult: tuple[int] = (1, 2, 4, 8),
num_res_blocks: int,
attn_layers: list[int] = [],
down_layers: list[int] = [],
kernel_size: int = 3,
resamp_with_conv: bool = True,
in_dim: int,
embed_dim: int,
clip_act: float = 256.0):
super().__init__()
self.ch = dim
self.num_layers = len(ch_mult)
self.num_res_blocks = num_res_blocks
self.in_channels = in_dim
self.clip_act = clip_act
self.down_layers = [i + 1 for i in down_layers] # each downlayer add one
# compute in_ch_mult, block_in and curr_res at lowest res
block_in = dim * ch_mult[self.num_layers - 1]
# z to block_in
self.conv_in = MPConv1D(embed_dim, block_in, kernel_size=kernel_size)
# middle
self.mid = nn.Module()
self.mid.block_1 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
self.mid.attn_1 = AttnBlock1D(block_in)
self.mid.block_2 = ResnetBlock1D(in_dim=block_in, out_dim=block_in, use_norm=True)
# upsampling
self.up = nn.ModuleList()
for i_level in reversed(range(self.num_layers)):
block = nn.ModuleList()
attn = nn.ModuleList()
block_out = dim * ch_mult[i_level]
for i_block in range(self.num_res_blocks + 1):
block.append(ResnetBlock1D(in_dim=block_in, out_dim=block_out, use_norm=True))
block_in = block_out
if i_level in attn_layers:
attn.append(AttnBlock1D(block_in))
up = nn.Module()
up.block = block
up.attn = attn
if i_level in self.down_layers:
up.upsample = Upsample1D(block_in, resamp_with_conv)
self.up.insert(0, up) # prepend to get consistent order
# end
self.conv_out = MPConv1D(block_in, out_dim, kernel_size=kernel_size)
self.learnable_gain = nn.Parameter(torch.zeros([]))
def forward(self, z):
# z to block_in
h = self.conv_in(z)
# middle
h = self.mid.block_1(h)
h = self.mid.attn_1(h)
h = self.mid.block_2(h)
h = h.clamp(-self.clip_act, self.clip_act)
# upsampling
for i_level in reversed(range(self.num_layers)):
for i_block in range(self.num_res_blocks + 1):
h = self.up[i_level].block[i_block](h)
if len(self.up[i_level].attn) > 0:
h = self.up[i_level].attn[i_block](h)
h = h.clamp(-self.clip_act, self.clip_act)
if i_level in self.down_layers:
h = self.up[i_level].upsample(h)
h = nonlinearity(h)
h = self.conv_out(h, gain=(self.learnable_gain + 1))
return h
def VAE_16k(**kwargs) -> VAE:
return VAE(data_dim=80, embed_dim=20, hidden_dim=384, **kwargs)
def VAE_44k(**kwargs) -> VAE:
return VAE(data_dim=128, embed_dim=40, hidden_dim=512, **kwargs)
def get_my_vae(name: str, **kwargs) -> VAE:
if name == '16k':
return VAE_16k(**kwargs)
if name == '44k':
return VAE_44k(**kwargs)
raise ValueError(f'Unknown model: {name}')
if __name__ == '__main__':
network = get_my_vae('standard')
# print the number of parameters in terms of millions
num_params = sum(p.numel() for p in network.parameters()) / 1e6
print(f'Number of parameters: {num_params:.2f}M')
-117
View File
@@ -1,117 +0,0 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from .edm2_utils import (MPConv1D, mp_silu, mp_sum, normalize)
def nonlinearity(x):
# swish
return mp_silu(x)
class ResnetBlock1D(nn.Module):
def __init__(self, *, in_dim, out_dim=None, conv_shortcut=False, kernel_size=3, use_norm=True):
super().__init__()
self.in_dim = in_dim
out_dim = in_dim if out_dim is None else out_dim
self.out_dim = out_dim
self.use_conv_shortcut = conv_shortcut
self.use_norm = use_norm
self.conv1 = MPConv1D(in_dim, out_dim, kernel_size=kernel_size)
self.conv2 = MPConv1D(out_dim, out_dim, kernel_size=kernel_size)
if self.in_dim != self.out_dim:
if self.use_conv_shortcut:
self.conv_shortcut = MPConv1D(in_dim, out_dim, kernel_size=kernel_size)
else:
self.nin_shortcut = MPConv1D(in_dim, out_dim, kernel_size=1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# pixel norm
if self.use_norm:
x = normalize(x, dim=1)
h = x
h = nonlinearity(h)
h = self.conv1(h)
h = nonlinearity(h)
h = self.conv2(h)
if self.in_dim != self.out_dim:
if self.use_conv_shortcut:
x = self.conv_shortcut(x)
else:
x = self.nin_shortcut(x)
return mp_sum(x, h, t=0.3)
class AttnBlock1D(nn.Module):
def __init__(self, in_channels, num_heads=1):
super().__init__()
self.in_channels = in_channels
self.num_heads = num_heads
self.qkv = MPConv1D(in_channels, in_channels * 3, kernel_size=1)
self.proj_out = MPConv1D(in_channels, in_channels, kernel_size=1)
def forward(self, x):
h = x
y = self.qkv(h)
y = y.reshape(y.shape[0], self.num_heads, -1, 3, y.shape[-1])
q, k, v = normalize(y, dim=2).unbind(3)
q = rearrange(q, 'b h c l -> b h l c')
k = rearrange(k, 'b h c l -> b h l c')
v = rearrange(v, 'b h c l -> b h l c')
h = F.scaled_dot_product_attention(q, k, v)
h = rearrange(h, 'b h l c -> b (h c) l')
h = self.proj_out(h)
return mp_sum(x, h, t=0.3)
class Upsample1D(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = MPConv1D(in_channels, in_channels, kernel_size=3)
def forward(self, x):
x = F.interpolate(x, scale_factor=2.0, mode='nearest-exact') # support 3D tensor(B,C,T)
if self.with_conv:
x = self.conv(x)
return x
class Downsample1D(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
# no asymmetric padding in torch conv, must do it ourselves
self.conv1 = MPConv1D(in_channels, in_channels, kernel_size=1)
self.conv2 = MPConv1D(in_channels, in_channels, kernel_size=1)
def forward(self, x):
if self.with_conv:
x = self.conv1(x)
x = F.avg_pool1d(x, kernel_size=2, stride=2)
if self.with_conv:
x = self.conv2(x)
return x
-96
View File
@@ -1,96 +0,0 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddSCAILReferenceEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"ref_image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, clip_embeds=None):
updated = dict(embeds)
vae.to(device)
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
ref_latent = vae.encode([ref_image_in], device, tiled=False)[0]
log.info(f"SCAIL ref_latent shape: {ref_latent.shape}")
ref_mask = torch.ones_like(ref_latent[:4])
ref_latent = torch.cat([ref_latent, ref_mask], dim=0)
vae.to(offload_device)
updated.setdefault("scail_embeds", {})
updated["scail_embeds"]["ref_latent_pos"] = ref_latent * strength
updated["scail_embeds"]["ref_latent_neg"] = torch.zeros_like(ref_latent)
updated["scail_embeds"]["ref_start_percent"] = start_percent
updated["scail_embeds"]["ref_end_percent"] = end_percent
updated["clip_context"] = clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None
return (updated,)
class WanVideoAddSCAILPoseEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
},
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, pose_images, strength, start_percent=0.0, end_percent=1.0):
updated = dict(embeds)
vae.to(device)
pose_images_in = (pose_images[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
pose_latent = vae.encode([pose_images_in], device, tiled=False)[0]
pose_mask = torch.ones_like(pose_latent[:4])
pose_latent = torch.cat([pose_latent, pose_mask], dim=0)
log.info(f"SCAIL pose_latent shape: {pose_latent.shape}")
vae.to(offload_device)
updated.setdefault("scail_embeds", {})
updated["scail_embeds"]["pose_latent"] = pose_latent
updated["scail_embeds"]["pose_strength"] = strength
updated["scail_embeds"]["pose_start_percent"] = start_percent
updated["scail_embeds"]["pose_end_percent"] = end_percent
return (updated,)
NODE_CLASS_MAPPINGS = {
"WanVideoAddSCAILPoseEmbeds": WanVideoAddSCAILPoseEmbeds,
"WanVideoAddSCAILReferenceEmbeds": WanVideoAddSCAILReferenceEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddSCAILReferenceEmbeds": "WanVideo Add SCAIL Reference Embeds",
"WanVideoAddSCAILPoseEmbeds": "WanVideo Add SCAIL Pose Embeds",
}
Binary file not shown.
Binary file not shown.
-207
View File
@@ -1,207 +0,0 @@
import json
import torch
import torchvision.transforms.functional as TF
from ..utils import log
from .trajectory import create_pos_feature_map, draw_tracks_on_video, replace_feature
import os
from comfy import model_management as mm
device = mm.get_torch_device()
script_directory = os.path.dirname(os.path.abspath(__file__))
VAE_STRIDE = (4, 8, 8) # t, h, w
class WanVideoWanDrawWanMoveTracks:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"images": ("IMAGE",),
"tracks": ("TRACKS",),
},
"optional": {
"line_resolution": ("INT", {"default": 24, "min": 4, "max": 64, "step": 1, "tooltip": "Number of points to use for each line segment"}),
"circle_size": ("INT", {"default": 10, "min": 1, "max": 20, "step": 1, "tooltip": "Size of the circle to draw for each track point"}),
"opacity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Opacity of the circle to draw for each track point"}),
"line_width": ("INT", {"default": 14, "min": 1, "max": 50, "step": 1, "tooltip": "Width of the line to draw for each track"}),
}
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("image",)
FUNCTION = "execute"
CATEGORY = "WanVideoWrapper"
def execute(self, images, tracks, line_resolution=24, circle_size=10, opacity=0.5, line_width=14):
if tracks is None or "track_path" not in tracks:
log.warning("WanVideoWanDrawWanMoveTracks: No tracks provided.")
return (images.float().cpu(), )
track = tracks["track_path"].unsqueeze(0)
track_visibility = tracks["track_visibility"].unsqueeze(0)
images_in = images * 255.0
if images_in.shape[0] != track.shape[1]:
repeat_count = track.shape[1] // images.shape[0]
images_in = images_in.repeat(repeat_count, 1, 1, 1)
track_video = draw_tracks_on_video(images_in, track, track_visibility, track_frame=line_resolution, circle_size=circle_size, opacity=opacity, line_width=line_width)
track_video = torch.stack([TF.to_tensor(frame) for frame in track_video], dim=0).movedim(1, -1)
return (track_video.float().cpu(), )
class WanVideoAddWanMoveTracks:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
},
"optional": {
"track_mask": ("MASK",),
"track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}),
"tracks": ("TRACKS", {"tooltip": "Alternatively use Comfy Tracks dictionary"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "TRACKS")
RETURN_NAMES = ("image_embeds", "tracks")
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, image_embeds, track_coords=None, tracks=None, strength=1.0, track_mask=None):
updated = dict(image_embeds)
track_visibility = None
target_shape = image_embeds.get("target_shape")
if target_shape is not None:
height = target_shape[2] * VAE_STRIDE[1]
width = target_shape[3] * VAE_STRIDE[2]
else:
height = image_embeds["lat_h"] * VAE_STRIDE[1]
width = image_embeds["lat_w"] * VAE_STRIDE[2]
num_frames = image_embeds["num_frames"]
if track_coords is not None:
tracks_data = parse_json_tracks(track_coords)
track_list = [
[[track[frame]['x'], track[frame]['y']] for track in tracks_data]
for frame in range(len(tracks_data[0]))
]
track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2)
elif tracks is not None and "track_path" in tracks:
track = tracks["track_path"]
if track_mask is None:
track_visibility = tracks.get("track_visibility", None)
track = track[:num_frames]
num_tracks = track.shape[-2]
if track_visibility is None:
if track_mask is None:
track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device)
else:
track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1)
feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=num_tracks, device=device)
updated.setdefault("wanmove_embeds", {})
updated["wanmove_embeds"]["track_pos"] = track_pos
updated["wanmove_embeds"]["strength"] = strength
tracks_dict = {
"track_path": track,
"track_visibility": track_visibility,
}
return (updated, tracks_dict,)
def parse_json_tracks(tracks):
tracks_data = []
try:
# If tracks is a string, try to parse it as JSON
if isinstance(tracks, str):
parsed = json.loads(tracks.replace("'", '"'))
tracks_data.extend(parsed)
else:
# If tracks is a list of strings, parse each one
for track_str in tracks:
parsed = json.loads(track_str.replace("'", '"'))
tracks_data.append(parsed)
# Check if we have a single track (dict with x,y) or a list of tracks
if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]:
# Single track detected, wrap it in a list
tracks_data = [tracks_data]
elif tracks_data and isinstance(tracks_data[0], list) and tracks_data[0] and isinstance(tracks_data[0][0], dict) and 'x' in tracks_data[0][0]:
# Already a list of tracks, nothing to do
pass
else:
# Unexpected format
log.warning(f"Warning: Unexpected track format: {type(tracks_data[0])}")
except json.JSONDecodeError as e:
log.warning(f"Error parsing tracks JSON: {e}")
tracks_data = []
return tracks_data
import node_helpers
class WanMove_native:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"positive": ("CONDITIONING",),
"track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}),
},
"optional": {
"track_mask": ("MASK",),
}
}
RETURN_TYPES = ("CONDITIONING", "TRACKS")
RETURN_NAMES = ("positive", "tracks")
FUNCTION = "patchcond"
CATEGORY = "WanVideoWrapper"
DEPRECATED = True
def patchcond(self, positive, track_coords, track_mask=None):
concat_latent_image = positive[0][1]["concat_latent_image"]
B, C, T, H, W = concat_latent_image.shape
num_frames = (T-1) * 4 + 1
width = W * 8
height = H * 8
tracks_data = parse_json_tracks(track_coords)
track_list = [
[[track[frame]['x'], track[frame]['y']] for track in tracks_data]
for frame in range(len(tracks_data[0]))
]
track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2)
track = track[:num_frames]
num_tracks = track.shape[-2]
if track_mask is None:
track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device)
else:
track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1)
feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=num_tracks, device=device)
wanmove_cond = replace_feature(concat_latent_image, track_pos.unsqueeze(0))
positive = node_helpers.conditioning_set_values(positive, {"concat_latent_image": wanmove_cond})
tracks_dict = {
"track_path": track,
"track_visibility": track_visibility,
}
return (positive, tracks_dict)
NODE_CLASS_MAPPINGS = {
"WanVideoAddWanMoveTracks": WanVideoAddWanMoveTracks,
"WanVideoWanDrawWanMoveTracks": WanVideoWanDrawWanMoveTracks,
"WanMove_native": WanMove_native,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddWanMoveTracks": "WanVideo Add WanMove Tracks",
"WanVideoWanDrawWanMoveTracks": "WanVideo Draw WanMove Tracks",
"WanMove_native": "WanMove Native",
}
-340
View File
@@ -1,340 +0,0 @@
# https://github.com/ali-vilab/Wan-Move/blob/main/wan/modules/trajectory.py
import numpy as np
import torch
from PIL import Image, ImageDraw
SKIP_ZERO = False
def get_pos_emb(
pos_k: torch.Tensor,
pos_emb_dim: int,
theta_func: callable = lambda i, d: torch.pow(10000, torch.mul(2, torch.div(i.to(torch.float32), d))),
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
dtype: torch.dtype = torch.float32,
) -> torch.Tensor:
"""
Generate batch position embeddings.
Args:
pos_k (torch.Tensor): A 1D tensor containing positions for which to generate embeddings.
pos_emb_dim (int): The dimension of position embeddings.
theta_func (callable): Function to compute thetas based on position and embedding dimensions.
device (torch.device): Device to store the position embeddings.
dtype (torch.dtype): Desired data type for computations.
Returns:
torch.Tensor: The position embeddings with shape (batch_size, pos_emb_dim).
"""
assert pos_emb_dim % 2 == 0, "The dimension of position embeddings must be even."
pos_k = pos_k.to(device, dtype)
if SKIP_ZERO:
pos_k = pos_k + 1
batch_size = pos_k.size(0)
denominator = torch.arange(0, pos_emb_dim // 2, device=device, dtype=dtype)
# Expand denominator to match the shape needed for broadcasting
denominator_expanded = denominator.view(1, -1).expand(batch_size, -1)
thetas = theta_func(denominator_expanded, pos_emb_dim)
# Ensure pos_k is in the correct shape for broadcasting
pos_k_expanded = pos_k.view(-1, 1).to(dtype)
sin_thetas = torch.sin(torch.div(pos_k_expanded, thetas))
cos_thetas = torch.cos(torch.div(pos_k_expanded, thetas))
# Concatenate sine and cosine embeddings along the last dimension
pos_emb = torch.cat([sin_thetas, cos_thetas], dim=-1)
return pos_emb
def create_pos_feature_map(
pred_tracks: torch.Tensor, # [T, N, 2]
pred_visibility: torch.Tensor, # [T, N]
downsample_ratios: list[int],
height: int,
width: int,
pos_emb_dim: int,
track_num: int = -1,
t_down_strategy: str = "sample",
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
dtype: torch.dtype = torch.float32,
):
"""
Create a feature map from the predicted tracks.
Args:
- pred_tracks: torch.Tensor, the predicted tracks, [T, N, 2]
- pred_visibility: torch.Tensor, the predicted visibility, [T, N]
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
- height: int, the height of the feature map
- width: int, the width of the feature map
- pos_emb_dim: int, the dimension of the position embeddings
- track_num: int, the number of tracks to use
- t_down_strategy: str, the strategy for downsampling time dimension
- device: torch.device, the device
- dtype: torch.dtype, the data type
Returns:
- feature_map: torch.Tensor, the feature map, [T', H', W', pos_emb_dim]
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
"""
assert t_down_strategy in ["sample", "average"], "Invalid strategy for downsampling time dimension."
t, n, _ = pred_tracks.shape
t_down, h_down, w_down = downsample_ratios
feature_map = torch.zeros((t-1) // t_down + 1, height // h_down, width // w_down, pos_emb_dim, device=device, dtype=dtype)
track_pos = - torch.ones(n, (t-1) // t_down + 1, 2, dtype=torch.long)
if track_num == -1:
track_num = n
tracks_idx = torch.randperm(n)[:track_num]
tracks = pred_tracks[:, tracks_idx]
visibility = pred_visibility[:, tracks_idx]
#tracks_embs = get_pos_emb(torch.randperm(n)[:track_num], pos_emb_dim, device=device, dtype=dtype)
for t_idx in range(0, t, t_down):
if t_down_strategy == "sample" or t_idx == 0:
cur_tracks = tracks[t_idx] # [N, 2]
cur_visibility = visibility[t_idx] # [N]
else:
cur_tracks = tracks[t_idx:t_idx+t_down].mean(dim=0)
cur_visibility = torch.any(visibility[t_idx:t_idx+t_down], dim=0)
for i in range(track_num):
if not cur_visibility[i] or cur_tracks[i][0] < 0 or cur_tracks[i][1] < 0 or cur_tracks[i][0] >= width or cur_tracks[i][1] >= height:
continue
x, y = cur_tracks[i]
x, y = int(x // w_down), int(y // h_down)
#feature_map[t_idx // t_down, y, x] += tracks_embs[i]
track_pos[i, t_idx // t_down, 0], track_pos[i, t_idx // t_down, 1] = y, x
return feature_map, track_pos
def replace_feature(
vae_feature: torch.Tensor, # [B, C', T', H', W']
track_pos: torch.Tensor, # [B, N, T', 2]
strength: float = 1.0,
) -> torch.Tensor:
b, _, t, h, w = vae_feature.shape
assert b == track_pos.shape[0], "Batch size mismatch."
n = track_pos.shape[1]
# Shuffle the trajectory order
track_pos = track_pos[:, torch.randperm(n), :, :]
# Extract coordinates at time steps ≥ 1 and generate a valid mask
current_pos = track_pos[:, :, 1:, :] # [B, N, T-1, 2]
mask = (current_pos[..., 0] >= 0) & (current_pos[..., 1] >= 0) # [B, N, T-1]
# Get all valid indices
valid_indices = mask.nonzero(as_tuple=False) # [num_valid, 3]
num_valid = valid_indices.shape[0]
if num_valid == 0:
return vae_feature
# Decompose valid indices into each dimension
batch_idx = valid_indices[:, 0]
track_idx = valid_indices[:, 1]
t_rel = valid_indices[:, 2]
t_target = t_rel + 1 # Convert to original time step indices
# Extract target position coordinates
h_target = current_pos[batch_idx, track_idx, t_rel, 0].long() # Ensure integer indices
w_target = current_pos[batch_idx, track_idx, t_rel, 1].long()
# Extract source position coordinates (t=0)
h_source = track_pos[batch_idx, track_idx, 0, 0].long()
w_source = track_pos[batch_idx, track_idx, 0, 1].long()
# Get source features and assign to target positions
src_features = vae_feature[batch_idx, :, 0, h_source, w_source]
dst_features = vae_feature[batch_idx, :, t_target, h_target, w_target]
vae_feature[batch_idx, :, t_target, h_target, w_target] = dst_features + (src_features - dst_features) * strength
return vae_feature
def get_video_track_video(
model,
video_tensor: torch.Tensor, # [T, C, H, W]
downsample_ratios: list[int],
pos_emb_dim: int,
grid_size: int = 32,
track_num: int = -1,
t_down_strategy: str = "sample",
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
dtype: torch.dtype = torch.float32,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Get the track video from the video tensor.
Args:
- model: torch.nn.Module, the model for tracking, CoTracker
- video_tensor: torch.Tensor, the video tensor, [T, C, H, W]
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
- height: int, the height of the feature map
- width: int, the width of the feature map
- pos_emb_dim: int, the dimension of the position embeddings
- grid_size: int, the size of the grid
- track_num: int, the number of tracks to use
- t_down_strategy: str, the strategy for downsampling time dimension
- device: torch.device, the device
- dtype: torch.dtype, the data type
Returns:
- track_video: torch.Tensor, the track video, [pos_emb_dim, T', H', W']
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
- pred_tracks: the predicted point trajectories
- pred_visibility: visibility of the predicted point trajectories
"""
t, c, height, width = video_tensor.shape
with (
torch.autocast(device_type=device.type, dtype=dtype),
torch.no_grad(),
):
pred_tracks, pred_visibility = model(
video_tensor.unsqueeze(0),
grid_size=grid_size,
backward_tracking=False,
)
track_video, track_pos = create_pos_feature_map(
pred_tracks[0], pred_visibility[0], downsample_ratios, height, width, pos_emb_dim, track_num, t_down_strategy, device, dtype
)
return track_video.permute(3, 0, 1, 2), track_pos, pred_tracks, pred_visibility
# ---------------------------
# Visualize functions
# --------------------------
def add_weighted(rgb, track):
rgb = np.array(rgb) # [H, W, C] "RGB"
track = np.array(track) # [H, W, C] "RGBA"
# Compute weights from the alpha channel
alpha = track[:, :, 3] / 255.0
# Expand alpha to 3 channels to match RGB
alpha = np.stack([alpha] * 3, axis=-1)
# Blend the two images
blend_img = track[:, :, :3] * alpha + rgb * (1 - alpha)
return Image.fromarray(blend_img.astype(np.uint8))
def draw_tracks_on_video(video, tracks, visibility=None, track_frame=24, circle_size=12, opacity=0.5, line_width=16):
color_map = [(102, 153, 255), (0, 255, 255), (255, 255, 0), (255, 102, 204), (0, 255, 0)]
video = video.byte().cpu().numpy() # (81, 480, 832, 3)
tracks = tracks[0].long().detach().cpu().numpy()
if visibility is not None:
visibility = visibility[0].detach().cpu().numpy()
num_frames, height, width = video.shape[:3]
num_tracks = tracks.shape[1]
alpha_opacity = int(255 * opacity)
output_frames = []
for t in range(num_frames):
frame_rgb = video[t].astype(np.float32)
# Create a single RGBA overlay for all tracks in this frame
overlay = Image.new("RGBA", (width, height), (0, 0, 0, 0))
draw_overlay = ImageDraw.Draw(overlay)
polyline_data = []
# Draw all circles on a single overlay
for n in range(num_tracks):
if visibility is not None and visibility[t, n] == 0:
continue
track_coord = tracks[t, n]
color = color_map[n % len(color_map)]
circle_color = color + (alpha_opacity,)
draw_overlay.ellipse(
(
track_coord[0] - circle_size,
track_coord[1] - circle_size,
track_coord[0] + circle_size,
track_coord[1] + circle_size
),
fill=circle_color
)
# Store polyline data for batch processing
tracks_coord = tracks[max(t - track_frame, 0):t + 1, n]
if len(tracks_coord) > 1:
polyline_data.append((tracks_coord, color))
# Blend circles overlay once
overlay_np = np.array(overlay)
alpha = overlay_np[:, :, 3:4] / 255.0
frame_rgb = overlay_np[:, :, :3] * alpha + frame_rgb * (1 - alpha)
# Draw all polylines on a single overlay
if polyline_data:
polyline_overlay = Image.new("RGBA", (width, height), (0, 0, 0, 0))
for tracks_coord, color in polyline_data:
_draw_gradient_polyline_on_overlay(polyline_overlay, line_width, tracks_coord, color, opacity)
# Blend polylines overlay once
polyline_np = np.array(polyline_overlay)
alpha = polyline_np[:, :, 3:4] / 255.0
frame_rgb = polyline_np[:, :, :3] * alpha + frame_rgb * (1 - alpha)
output_frames.append(Image.fromarray(frame_rgb.astype(np.uint8)))
return output_frames
def _draw_gradient_polyline_on_overlay(overlay, line_width, points, start_color, opacity=1.0):
"""
Draw a gradient polyline directly onto an existing RGBA overlay image.
This is an optimized version that doesn't create new images.
"""
draw = ImageDraw.Draw(overlay, 'RGBA')
points = points[::-1]
# Compute total length
total_length = 0
segment_lengths = []
for i in range(len(points) - 1):
dx = points[i + 1][0] - points[i][0]
dy = points[i + 1][1] - points[i][1]
length = (dx * dx + dy * dy) ** 0.5
segment_lengths.append(length)
total_length += length
if total_length == 0:
return
accumulated_length = 0
# Draw the gradient polyline
for idx, (start_point, end_point) in enumerate(zip(points[:-1], points[1:])):
segment_length = segment_lengths[idx]
steps = max(int(segment_length), 1)
for i in range(steps):
current_length = accumulated_length + (i / steps) * segment_length
ratio = current_length / total_length
alpha = int(255 * (1 - ratio) * opacity)
color = (*start_color, alpha)
x = int(start_point[0] + (end_point[0] - start_point[0]) * i / steps)
y = int(start_point[1] + (end_point[1] - start_point[1]) * i / steps)
dynamic_line_width = max(int(line_width * (1 - ratio)), 1)
draw.line([(x, y), (x + 1, y)], fill=color, width=dynamic_line_width)
accumulated_length += segment_length
+5 -73
View File
@@ -1,75 +1,7 @@
try: from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
from .utils import check_duplicate_nodes, log, color_text from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
duplicate_dirs = check_duplicate_nodes()
if duplicate_dirs:
warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n"
for dir_path in duplicate_dirs:
warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
except:
pass
from .utils import log
NODE_CLASS_MAPPINGS = {} NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
# Required modules (will raise on import failure)
REQUIRED_MODULES = [
(".nodes", "Main"),
(".nodes_sampler", "Sampler"),
(".nodes_model_loading", "ModelLoading"),
(".nodes_utility", "Utility"),
(".cache_methods.nodes_cache", "Cache"),
]
# Optional modules (will warn on import failure)
OPTIONAL_MODULES = [
(".nodes_deprecated", "Deprecated"),
(".s2v.nodes", "S2V"),
(".FlashVSR.flashvsr_nodes", "FlashVSR"),
(".mocha.nodes", "Mocha"),
(".fun_camera.nodes", "FunCamera"),
(".uni3c.nodes", "Uni3C"),
(".controlnet.nodes", "ControlNet"),
(".ATI.nodes", "ATI"),
(".multitalk.nodes", "MultiTalk"),
(".recammaster.nodes", "RecamMaster"),
(".skyreels.nodes", "SkyReels"),
(".fantasytalking.nodes", "FantasyTalking"),
(".qwen.qwen", "Qwen"),
(".fantasyportrait.nodes", "FantasyPortrait"),
(".unianimate.nodes", "UniAnimate"),
(".MTV.nodes", "MTV"),
(".HuMo.nodes", "HuMo"),
(".lynx.nodes", "Lynx"),
(".Ovi.nodes_ovi", "Ovi"),
(".steadydancer.nodes", "SteadyDancer"),
(".onetoall.nodes", "OneToAll"),
(".WanMove.nodes", "WanMove"),
(".SCAIL.nodes", "SCAIL"),
(".LongCat.nodes", "LongCat"),
(".LongVie2.nodes", "LongVie2"),
]
def register_nodes(module_path: str, name: str, optional: bool) -> None:
"""Import and register nodes from a module."""
try:
import importlib
module = importlib.import_module(module_path, package=__package__)
NODE_CLASS_MAPPINGS.update(getattr(module, "NODE_CLASS_MAPPINGS", {}))
NODE_DISPLAY_NAME_MAPPINGS.update(getattr(module, "NODE_DISPLAY_NAME_MAPPINGS", {}))
except Exception as e:
if optional:
log.warning(f"WanVideoWrapper WARNING: {name} nodes not available: {e}")
else:
raise
# Register all node modules
for module_path, name in REQUIRED_MODULES:
register_nodes(module_path, name, optional=False)
for module_path, name in OPTIONAL_MODULES:
register_nodes(module_path, name, optional=True)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
-159
View File
@@ -1,159 +0,0 @@
from ..utils import log
import torch
def set_transformer_cache_method(transformer, timesteps, cache_args=None):
transformer.cache_device = cache_args["cache_device"]
if cache_args["cache_type"] == "TeaCache":
log.info(f"TeaCache: Using cache device: {transformer.cache_device}")
transformer.teacache_state.clear_all()
transformer.enable_teacache = True
transformer.rel_l1_thresh = cache_args["rel_l1_thresh"]
transformer.teacache_start_step = cache_args["start_step"]
transformer.teacache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"]
transformer.teacache_use_coefficients = cache_args["use_coefficients"]
transformer.teacache_mode = cache_args["mode"]
elif cache_args["cache_type"] == "MagCache":
log.info(f"MagCache: Using cache device: {transformer.cache_device}")
transformer.magcache_state.clear_all()
transformer.enable_magcache = True
transformer.magcache_start_step = cache_args["start_step"]
transformer.magcache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"]
transformer.magcache_thresh = cache_args["magcache_thresh"]
transformer.magcache_K = cache_args["magcache_K"]
elif cache_args["cache_type"] == "EasyCache":
log.info(f"EasyCache: Using cache device: {transformer.cache_device}")
transformer.easycache_state.clear_all()
transformer.enable_easycache = True
transformer.easycache_start_step = cache_args["start_step"]
transformer.easycache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"]
transformer.easycache_thresh = cache_args["easycache_thresh"]
return transformer
class TeaCacheState:
def __init__(self, cache_device='cpu'):
self.cache_device = cache_device
self.states = {}
self._next_pred_id = 0
def new_prediction(self, cache_device='cpu'):
"""Create new prediction state and return its ID"""
self.cache_device = cache_device
pred_id = self._next_pred_id
self._next_pred_id += 1
self.states[pred_id] = {
'previous_residual': None,
'accumulated_rel_l1_distance': 0,
'previous_modulated_input': None,
'skipped_steps': [],
}
return pred_id
def update(self, pred_id, **kwargs):
"""Update state for specific prediction"""
if pred_id not in self.states:
return None
for key, value in kwargs.items():
self.states[pred_id][key] = value
def get(self, pred_id):
return self.states.get(pred_id, {})
def clear_all(self):
self.states = {}
self._next_pred_id = 0
class MagCacheState:
def __init__(self, cache_device='cpu'):
self.cache_device = cache_device
self.states = {}
self._next_pred_id = 0
def new_prediction(self, cache_device='cpu'):
"""Create new prediction state and return its ID"""
self.cache_device = cache_device
pred_id = self._next_pred_id
self._next_pred_id += 1
self.states[pred_id] = {
'residual_cache': None,
'accumulated_ratio': 1.0,
'accumulated_steps': 0,
'accumulated_err': 0,
'skipped_steps': [],
}
return pred_id
def update(self, pred_id, **kwargs):
"""Update state for specific prediction"""
if pred_id not in self.states:
return None
for key, value in kwargs.items():
self.states[pred_id][key] = value
def get(self, pred_id):
return self.states.get(pred_id, {})
def clear_all(self):
self.states = {}
self._next_pred_id = 0
class EasyCacheState:
def __init__(self, cache_device='cpu'):
self.cache_device = cache_device
self.states = {}
self._next_pred_id = 0
def new_prediction(self, cache_device='cpu'):
"""Create a new prediction state and return its ID."""
self.cache_device = cache_device
pred_id = self._next_pred_id
self._next_pred_id += 1
self.states[pred_id] = {
'previous_raw_input': None,
'previous_raw_output': None,
'cache': None,
'accumulated_error': 0.0,
'skipped_steps': [],
'cache_ovi': None,
}
return pred_id
def update(self, pred_id, **kwargs):
"""Update state for a specific prediction."""
if pred_id not in self.states:
return None
for key, value in kwargs.items():
self.states[pred_id][key] = value
def get(self, pred_id):
return self.states.get(pred_id, {})
def clear_all(self):
self.states = {}
self._next_pred_id = 0
def relative_l1_distance(last_tensor, current_tensor):
l1_distance = torch.abs(last_tensor.to(current_tensor.device) - current_tensor).mean()
norm = torch.abs(last_tensor).mean()
relative_l1_distance = l1_distance / norm
return relative_l1_distance.to(torch.float32).to(current_tensor.device)
def cache_report(transformer, cache_args):
cache_type = cache_args["cache_type"]
states = (
transformer.teacache_state.states if cache_type == "TeaCache" else
transformer.magcache_state.states if cache_type == "MagCache" else
transformer.easycache_state.states if cache_type == "EasyCache" else
None
)
state_names = {
0: "conditional",
1: "unconditional"
}
for pred_id, state in states.items():
name = state_names.get(pred_id, f"prediction_{pred_id}")
if 'skipped_steps' in state:
log.info(f"{cache_type} skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}")
transformer.teacache_state.clear_all()
transformer.magcache_state.clear_all()
transformer.easycache_state.clear_all()
del states
-128
View File
@@ -1,128 +0,0 @@
from comfy import model_management as mm
class WanVideoTeaCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"rel_l1_thresh": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.001,
"tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts. Good value range for 1.3B: 0.05 - 0.08, for other models 0.15-0.30"}),
"start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Start percentage of the steps to apply TeaCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "End steps to apply TeaCache"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
"use_coefficients": ("BOOLEAN", {"default": True, "tooltip": "Use calculated coefficients for more accuracy. When enabled therel_l1_thresh should be about 10 times higher than without"}),
},
"optional": {
"mode": (["e", "e0"], {"default": "e", "tooltip": "Choice between using e (time embeds, default) or e0 (modulated time embeds)"}),
},
}
RETURN_TYPES = ("CACHEARGS",)
RETURN_NAMES = ("cache_args",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = """
Patch WanVideo model to use TeaCache. Speeds up inference by caching the output and
applying it instead of doing the step. Best results are achieved by choosing the
appropriate coefficients for the model. Early steps should never be skipped, with too
aggressive values this can happen and the motion suffers. Starting later can help with that too.
When NOT using coefficients, the threshold value should be
about 10 times smaller than the value used with coefficients.
Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1
"""
def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"):
if cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
cache_args = {
"cache_type": "TeaCache",
"rel_l1_thresh": rel_l1_thresh,
"start_step": start_step,
"end_step": end_step,
"cache_device": cache_device,
"use_coefficients": use_coefficients,
"mode": mode,
}
return (cache_args,)
class WanVideoMagCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"magcache_thresh": ("FLOAT", {"default": 0.02, "min": 0.0, "max": 0.3, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}),
"magcache_K": ("INT", {"default": 4, "min": 0, "max": 6, "step": 1, "tooltip": "The maxium skip steps of MagCache."}),
"start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying MagCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying MagCache"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
},
}
RETURN_TYPES = ("CACHEARGS",)
RETURN_NAMES = ("cache_args",)
FUNCTION = "setargs"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
DESCRIPTION = "MagCache for WanVideoWrapper, source https://github.com/Zehong-Ma/MagCache"
def setargs(self, magcache_thresh, magcache_K, start_step, end_step, cache_device):
if cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
cache_args = {
"cache_type": "MagCache",
"magcache_thresh": magcache_thresh,
"magcache_K": magcache_K,
"start_step": start_step,
"end_step": end_step,
"cache_device": cache_device,
}
return (cache_args,)
class WanVideoEasyCache:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"easycache_thresh": ("FLOAT", {"default": 0.015, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}),
"start_step": ("INT", {"default": 10, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying EasyCache"}),
"end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying EasyCache"}),
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
},
}
RETURN_TYPES = ("CACHEARGS",)
RETURN_NAMES = ("cache_args",)
FUNCTION = "setargs"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
DESCRIPTION = "EasyCache for WanVideoWrapper, source https://github.com/H-EmbodVis/EasyCache"
def setargs(self, easycache_thresh, start_step, end_step, cache_device):
if cache_device == "main_device":
cache_device = mm.get_torch_device()
else:
cache_device = mm.unet_offload_device()
cache_args = {
"cache_type": "EasyCache",
"easycache_thresh": easycache_thresh,
"start_step": start_step,
"end_step": end_step,
"cache_device": cache_device,
}
return (cache_args,)
NODE_CLASS_MAPPINGS = {
"WanVideoTeaCache": WanVideoTeaCache,
"WanVideoMagCache": WanVideoMagCache,
"WanVideoEasyCache": WanVideoEasyCache,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoTeaCache": "WanVideo TeaCache",
"WanVideoMagCache": "WanVideo MagCache",
"WanVideoEasyCache": "WanVideo EasyCache"
}
+1 -75
View File
@@ -1,7 +1,6 @@
import numpy as np import numpy as np
from typing import Callable, Optional, List from typing import Callable, Optional, List
import torch
from ..utils import log
def ordered_halving(val): def ordered_halving(val):
bin_str = f"{val:064b}" bin_str = f"{val:064b}"
@@ -183,76 +182,3 @@ def get_total_steps(
) )
for i in range(len(timesteps)) for i in range(len(timesteps))
) )
def create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=False, window_type="linear"):
window_mask = torch.ones_like(noise_pred_context)
if window_type == "pyramid":
# Create pyramid weights that peak in the middle
length = noise_pred_context.shape[1]
if length % 2 == 0:
max_weight = length // 2
weight_sequence = list(range(1, max_weight + 1, 1)) + list(range(max_weight, 0, -1))
else:
max_weight = (length + 1) // 2
weight_sequence = list(range(1, max_weight, 1)) + [max_weight] + list(range(max_weight - 1, 0, -1))
# Normalize weights to range from 0 to 1
max_val = max(weight_sequence)
weight_sequence = [w / max_val for w in weight_sequence]
# Apply the weights to create the mask
weights_tensor = torch.tensor(weight_sequence, device=noise_pred_context.device)
weights_tensor = weights_tensor.view(1, -1, 1, 1)
window_mask = weights_tensor.expand_as(window_mask).clone()
# Adjust for position in sequence if needed
if not looped:
if min(c) == 0: # First chunk
left_ramp = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device).view(1, -1, 1, 1)
# Clone to avoid in-place memory conflict
left_section = window_mask[:, :context_overlap].clone()
window_mask[:, :context_overlap] = torch.maximum(left_section, left_ramp)
if max(c) == latent_video_length - 1: # Last chunk
right_ramp = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device).view(1, -1, 1, 1)
# Clone to avoid in-place memory conflict
right_section = window_mask[:, -context_overlap:].clone()
window_mask[:, -context_overlap:] = torch.maximum(right_section, right_ramp)
else: # Original "linear" window masking
# Apply left-side blending for all except first chunk (or always in loop mode)
if min(c) > 0 or (looped and max(c) == latent_video_length - 1):
ramp_up = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device)
ramp_up = ramp_up.view(1, -1, 1, 1)
window_mask[:, :context_overlap] = ramp_up
# Apply right-side blending for all except last chunk (or always in loop mode)
if max(c) < latent_video_length - 1 or (looped and min(c) == 0):
ramp_down = torch.linspace(1, 0, context_overlap, device=noise_pred_context.device)
ramp_down = ramp_down.view(1, -1, 1, 1)
window_mask[:, -context_overlap:] = ramp_down
return window_mask
class WindowTracker:
def __init__(self, verbose=False):
self.window_map = {} # Maps frame sequence to persistent ID
self.next_id = 0
self.cache_states = {} # Maps persistent ID to teacache state
self.verbose = verbose
def get_window_id(self, frames):
key = tuple(sorted(frames)) # Order-independent frame sequence
if key not in self.window_map:
self.window_map[key] = self.next_id
if self.verbose:
log.info(f"New window pattern {key} -> ID {self.next_id}")
self.next_id += 1
return self.window_map[key]
def get_teacache(self, window_id, base_state):
if window_id not in self.cache_states:
if self.verbose:
log.info(f"Initializing persistent teacache for window {window_id}")
self.cache_states[window_id] = base_state.copy()
return self.cache_states[window_id]
-173
View File
@@ -1,173 +0,0 @@
import torch
from ..utils import log
import comfy.model_management as mm
from comfy.utils import load_torch_file
from tqdm import tqdm
import gc
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import folder_paths
class WanVideoControlnetLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("controlnet"), {"tooltip": "These models are loaded from the 'ComfyUI/models/controlnet' -folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "bf16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn'], {"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"}),
},
}
RETURN_TYPES = ("WANVIDEOCONTROLNET",)
RETURN_NAMES = ("controlnet", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Loads ControlNet model from 'https://huggingface.co/collections/TheDenk/wan21-controlnets-68302b430411dafc0d74d2fc'"
def loadmodel(self, model, base_precision, load_device, quantization):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
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]
model_path = folder_paths.get_full_path_or_raise("controlnet", model)
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
num_layers = 8 if "blocks.7.scale_shift_table" in sd else 6
out_proj_dim = sd["controlnet_blocks.0.bias"].shape[0]
downscale_coef = 16 if out_proj_dim == 3072 else 8
vae_channels = 48 if out_proj_dim == 3072 else 16
if not "control_encoder.0.0.weight" in sd:
raise ValueError("Invalid ControlNet model")
controlnet_cfg = {
"added_kv_proj_dim": None,
"attention_head_dim": 128,
"cross_attn_norm": None,
"downscale_coef": downscale_coef,
"eps": 1e-06,
"ffn_dim": 8960,
"freq_dim": 256,
"image_dim": None,
"in_channels": 3,
"num_attention_heads": 12,
"num_layers": num_layers,
"out_proj_dim": out_proj_dim,
"patch_size": [
1,
2,
2
],
"qk_norm": "rms_norm_across_heads",
"rope_max_seq_len": 1024,
"text_dim": 4096,
"vae_channels": vae_channels
}
print(f"Loading WanControlnet with config: {controlnet_cfg}")
from .wan_controlnet import WanControlnet
with init_empty_weights():
controlnet = WanControlnet(**controlnet_cfg)
controlnet.eval()
if quantization == "disabled":
for k, v in sd.items():
if isinstance(v, torch.Tensor):
if v.dtype == torch.float8_e4m3fn:
quantization = "fp8_e4m3fn"
break
elif v.dtype == torch.float8_e5m2:
quantization = "fp8_e5m2"
break
if "fp8_e4m3fn" in quantization:
dtype = torch.float8_e4m3fn
elif quantization == "fp8_e5m2":
dtype = torch.float8_e5m2
else:
dtype = base_dtype
params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter"}
log.info("Using accelerate to load and assign controlnet model weights to device...")
param_count = sum(1 for _ in controlnet.named_parameters())
for name, param in tqdm(controlnet.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
if "controlnet_patch_embedding" in name:
dtype_to_use = torch.float32
set_module_tensor_to_device(controlnet, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
del sd
if load_device == "offload_device" and controlnet.device != offload_device:
log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}")
controlnet.to(offload_device)
gc.collect()
mm.soft_empty_cache()
return (controlnet,)
class WanVideoControlnetApply:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("WANVIDEOMODEL", ),
"controlnet": ("WANVIDEOCONTROLNET", ),
"control_images": ("IMAGE", ),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.0001, "tooltip": "controlnet strength"}),
"control_stride": ("INT", {"default": 3, "min": 1, "max": 8, "step": 1, "tooltip": "controlnet stride"}),
"control_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply controlnet"}),
"control_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply controlnet"}),
}
}
RETURN_TYPES = ("WANVIDEOMODEL",)
RETURN_NAMES = ("model", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, controlnet, control_images, strength, control_stride, control_start_percent, control_end_percent):
patcher = model.clone()
if 'transformer_options' not in patcher.model_options:
patcher.model_options['transformer_options'] = {}
control_input = control_images.permute(3, 0, 1, 2).unsqueeze(0).contiguous()
control_input = control_input * 2.0 - 1.0
controlnet = {
"controlnet": controlnet,
"control_latents": control_input,
"controlnet_strength": strength,
"control_stride": control_stride,
"controlnet_start": control_start_percent,
"controlnet_end": control_end_percent
}
patcher.model_options["transformer_options"]["controlnet"] = controlnet
return (patcher,)
NODE_CLASS_MAPPINGS = {
"WanVideoControlnetLoader": WanVideoControlnetLoader,
"WanVideoControlnet": WanVideoControlnetApply,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoControlnetLoader": "WanVideo Controlnet Loader",
"WanVideoControlnet": "WanVideo Controlnet Apply",
}
-236
View File
@@ -1,236 +0,0 @@
# source https://github.com/TheDenk/wan2.1-dilated-controlnet/blob/main/wan_controlnet.py
from typing import Any, Dict, Optional, Tuple, Union
import torch
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.transformers.transformer_wan import (
WanTimeTextImageEmbedding,
WanRotaryPosEmbed,
WanTransformerBlock
)
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r"""
A Controlnet Transformer model for video-like data used in the Wan model.
Args:
patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`):
3D patch dimensions for video embedding (t_patch, h_patch, w_patch).
num_attention_heads (`int`, defaults to `40`):
Fixed length for text embeddings.
attention_head_dim (`int`, defaults to `128`):
The number of channels in each head.
vae_channels (`int`, defaults to `16`):
The number of channels in the vae input.
in_channels (`int`, defaults to `16`):
The number of channels in the controlnet input.
text_dim (`int`, defaults to `512`):
Input dimension for text embeddings.
freq_dim (`int`, defaults to `256`):
Dimension for sinusoidal time embeddings.
ffn_dim (`int`, defaults to `13824`):
Intermediate dimension in feed-forward network.
num_layers (`int`, defaults to `40`):
The number of layers of transformer blocks to use.
window_size (`Tuple[int]`, defaults to `(-1, -1)`):
Window size for local attention (-1 indicates global attention).
cross_attn_norm (`bool`, defaults to `True`):
Enable cross-attention normalization.
qk_norm (`bool`, defaults to `True`):
Enable query/key normalization.
eps (`float`, defaults to `1e-6`):
Epsilon value for normalization layers.
add_img_emb (`bool`, defaults to `False`):
Whether to use img_emb.
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
The number of channels to use for the added key and value projections. If `None`, no projection is used.
downscale_coef (`int`, *optional*, defaults to `8`):
Coeficient for downscale controlnet input video.
out_proj_dim (`int`, *optional*, defaults to `128 * 12`):
Output projection dimention for last linear layers.
"""
_supports_gradient_checkpointing = True
_skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"]
_no_split_modules = ["WanTransformerBlock"]
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
@register_to_config
def __init__(
self,
patch_size: Tuple[int] = (1, 2, 2),
num_attention_heads: int = 40,
attention_head_dim: int = 128,
in_channels: int = 3,
vae_channels: int = 16,
text_dim: int = 4096,
freq_dim: int = 256,
ffn_dim: int = 13824,
num_layers: int = 20,
cross_attn_norm: bool = True,
qk_norm: Optional[str] = "rms_norm_across_heads",
eps: float = 1e-6,
image_dim: Optional[int] = None,
added_kv_proj_dim: Optional[int] = None,
rope_max_seq_len: int = 1024,
downscale_coef: int = 8,
out_proj_dim: int = 128 * 12,
) -> None:
super().__init__()
start_channels = in_channels * (downscale_coef ** 2)
input_channels = [start_channels, start_channels // 2, start_channels // 4]
self.control_encoder = nn.ModuleList([
## Spatial compression with time awareness
nn.Sequential(
nn.Conv3d(
in_channels,
input_channels[0],
kernel_size=(3, downscale_coef + 1, downscale_coef + 1),
stride=(1, downscale_coef, downscale_coef),
padding=(1, downscale_coef // 2, downscale_coef // 2)
),
nn.GELU(approximate="tanh"),
nn.GroupNorm(2, input_channels[0]),
),
## Spatio-Temporal compression with spatial awareness
nn.Sequential(
nn.Conv3d(input_channels[0], input_channels[1], kernel_size=3, stride=(2, 1, 1), padding=1),
nn.GELU(approximate="tanh"),
nn.GroupNorm(2, input_channels[1]),
),
## Temporal compression with spatial awareness
nn.Sequential(
nn.Conv3d(input_channels[1], input_channels[2], kernel_size=3, stride=(2, 1, 1), padding=1),
nn.GELU(approximate="tanh"),
nn.GroupNorm(2, input_channels[2]),
)
])
inner_dim = num_attention_heads * attention_head_dim
# 1. Patch & position embedding
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size)
# 2. Condition embeddings
# image_embedding_dim=1280 for I2V model
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=freq_dim,
time_proj_dim=inner_dim * 6,
text_embed_dim=text_dim,
image_embed_dim=image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList(
[
WanTransformerBlock(
inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim
)
for _ in range(num_layers)
]
)
# 4 Controlnet modules
self.controlnet_blocks = nn.ModuleList([])
for _ in range(len(self.blocks)):
controlnet_block = nn.Linear(inner_dim, out_proj_dim)
self.controlnet_blocks.append(controlnet_block)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_hidden_states: torch.Tensor,
controlnet_states: torch.Tensor,
encoder_hidden_states_image: Optional[torch.Tensor] = None,
return_dict: bool = True,
attention_kwargs: Optional[Dict[str, Any]] = None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if attention_kwargs is not None:
attention_kwargs = attention_kwargs.copy()
lora_scale = attention_kwargs.pop("scale", 1.0)
else:
lora_scale = 1.0
if USE_PEFT_BACKEND:
# weight the lora layers by setting `lora_scale` for each PEFT layer
scale_lora_layers(self, lora_scale)
else:
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
logger.warning(
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
)
rotary_emb = self.rope(hidden_states)
# 0. Controlnet encoder
for control_encoder_block in self.control_encoder:
controlnet_states = control_encoder_block(controlnet_states)
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
## 1. Patch embedding and stack
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.ndim == 2:
## for ComfyUI workflow
if hidden_states.shape[1] != timestep.shape[1]:
timestep = timestep.repeat_interleave(hidden_states.shape[1] // timestep.shape[1], dim=1)
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten() # batch_size * seq_len
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len
)
if ts_seq_len is not None:
# batch_size, seq_len, 6, inner_dim
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
# batch_size, 6, inner_dim
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 4. Transformer blocks
controlnet_hidden_states = ()
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block, controlnet_block in zip(self.blocks, self.controlnet_blocks):
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb
)
controlnet_hidden_states += (controlnet_block(hidden_states),)
else:
for block, controlnet_block in zip(self.blocks, self.controlnet_blocks):
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
controlnet_hidden_states += (controlnet_block(hidden_states),)
if USE_PEFT_BACKEND:
# remove `lora_scale` from each PEFT layer
unscale_lora_layers(self, lora_scale)
if not return_dict:
return (controlnet_hidden_states,)
return Transformer2DModelOutput(sample=controlnet_hidden_states)
-281
View File
@@ -1,281 +0,0 @@
import torch
import torch.nn as nn
from accelerate import init_empty_weights
from .gguf.gguf_utils import GGUFParameter, dequantize_gguf_tensor
@torch.library.custom_op("wanvideo::apply_lora", mutates_args=())
def apply_lora(weight: torch.Tensor, lora_diff_0: torch.Tensor, lora_diff_1: torch.Tensor, lora_diff_2: float, lora_strength: torch.Tensor) -> torch.Tensor:
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape)
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0
scale = lora_strength * alpha
return weight + patch_diff * scale
@apply_lora.register_fake
def _(weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
# Return weight with same metadata
return weight.clone()
@torch.library.custom_op("wanvideo::apply_single_lora", mutates_args=())
def apply_single_lora(weight: torch.Tensor, lora_diff: torch.Tensor, lora_strength: torch.Tensor) -> torch.Tensor:
return weight + lora_diff * lora_strength
@apply_single_lora.register_fake
def _(weight, lora_diff, lora_strength):
# Return weight with same metadata
return weight.clone()
@torch.library.custom_op("wanvideo::linear_forward", mutates_args=())
def linear_forward(input: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor:
return torch.nn.functional.linear(input, weight, bias)
@linear_forward.register_fake
def _(input, weight, bias):
# Calculate output shape: (..., out_features)
out_features = weight.shape[0]
output_shape = list(input.shape[:-1]) + [out_features]
return input.new_empty(output_shape)
#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, compile_args=None, modules_to_not_convert=[]):
has_children = list(model.children())
if not has_children:
return
allow_compile = False
for name, module in model.named_children():
if compile_args is not None:
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "."
module_prefix = module_prefix.replace("_orig_mod.", "")
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert)
if isinstance(module, nn.Linear) and "loras" not in module_prefix and "dual_controller" not in module_prefix and name not in modules_to_not_convert:
weight_key = module_prefix + "weight"
if weight_key not in state_dict:
continue
in_features = state_dict[weight_key].shape[1]
out_features = state_dict[weight_key].shape[0]
is_gguf = isinstance(state_dict[weight_key], GGUFParameter)
scale_weight = None
if not is_gguf and scale_weights is not None:
scale_key = f"{module_prefix}scale_weight"
scale_weight = scale_weights.get(scale_key)
with init_empty_weights():
model._modules[name] = CustomLinear(
in_features,
out_features,
module.bias is not None,
compute_dtype=compute_dtype,
scale_weight=scale_weight,
allow_compile=allow_compile,
is_gguf=is_gguf
)
model._modules[name].source_cls = type(module)
model._modules[name].requires_grad_(False)
return model
def set_lora_params(module, patches, module_prefix="", device=torch.device("cpu")):
remove_lora_from_module(module)
# Recursively set lora_diffs and lora_strengths for all CustomLinear layers
for name, child in module.named_children():
params = list(child.parameters())
if params:
device = params[0].device
else:
device = torch.device("cpu")
child_prefix = (f"{module_prefix}{name}.")
set_lora_params(child, patches, child_prefix, device)
if isinstance(module, CustomLinear):
key = f"diffusion_model.{module_prefix}weight"
patch = patches.get(key, [])
#print(f"Processing LoRA patches for {key}: {len(patch)} patches found")
if len(patch) == 0:
key = key.replace("_orig_mod.", "")
patch = patches.get(key, [])
#print(f"Processing LoRA patches for {key}: {len(patch)} patches found")
if len(patch) != 0:
lora_diffs = []
for p in patch:
lora_obj = p[1]
if "head" in key:
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]
module.set_lora_diffs(lora_diffs, device=device)
module.set_lora_strengths(lora_strengths, device=device)
module._step.fill_(0) # Initialize step for LoRA scheduling
class CustomLinear(nn.Linear):
def __init__(
self,
in_features,
out_features,
bias=False,
compute_dtype=None,
device=None,
scale_weight=None,
allow_compile=False,
is_gguf=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
self.lora_diffs = []
self.register_buffer("_step", torch.zeros((), dtype=torch.long))
self.scale_weight = scale_weight
self.lora_strengths = []
self.allow_compile = allow_compile
self.is_gguf = is_gguf
if not allow_compile:
self._apply_lora_impl = self._apply_lora_custom_op
self._apply_single_lora_impl = self._apply_single_lora_custom_op
self._linear_forward_impl = self._linear_forward_custom_op
else:
self._apply_lora_impl = self._apply_lora_direct
self._apply_single_lora_impl = self._apply_single_lora_direct
self._linear_forward_impl = self._linear_forward_direct
# Direct implementations (no custom ops)
def _apply_lora_direct(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
patch_diff = torch.mm(
lora_diff_0.flatten(start_dim=1),
lora_diff_1.flatten(start_dim=1)
).reshape(weight.shape) + 0
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0
scale = lora_strength * alpha
return weight + patch_diff * scale
def _apply_single_lora_direct(self, weight, lora_diff, lora_strength):
return weight + lora_diff * lora_strength
def _linear_forward_direct(self, input, weight, bias):
return torch.nn.functional.linear(input, weight, bias)
# Custom op implementations
def _apply_lora_custom_op(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
return torch.ops.wanvideo.apply_lora(weight, lora_diff_0, lora_diff_1,
float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
)
def _apply_single_lora_custom_op(self, weight, lora_diff, lora_strength):
return torch.ops.wanvideo.apply_single_lora(weight, lora_diff, lora_strength)
def _linear_forward_custom_op(self, input, weight, bias):
return torch.ops.wanvideo.linear_forward(input, weight, bias)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = []
for i, diff in enumerate(lora_diffs):
if len(diff) > 1:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
setattr(self, f"lora_diff_{i}_2", diff[2])
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
else:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.lora_diffs.append(f"lora_diff_{i}_0")
def set_lora_strengths(self, lora_strengths, device=torch.device("cpu")):
self._lora_strength_tensors = []
self._lora_strength_is_scheduled = []
self._step = self._step.to(device)
for i, strength in enumerate(lora_strengths):
if isinstance(strength, list):
tensor = torch.tensor(strength, dtype=self.compute_dtype, device=device)
self.register_buffer(f"_lora_strength_{i}", tensor)
self._lora_strength_is_scheduled.append(True)
else:
tensor = torch.tensor([strength], dtype=self.compute_dtype, device=device)
self.register_buffer(f"_lora_strength_{i}", tensor)
self._lora_strength_is_scheduled.append(False)
def _get_lora_strength(self, idx):
strength_tensor = getattr(self, f"_lora_strength_{idx}")
if self._lora_strength_is_scheduled[idx]:
return strength_tensor.index_select(0, self._step).squeeze(0)
return strength_tensor[0]
def _get_weight_with_lora(self, weight):
"""Apply LoRA using custom ops to avoid graph breaks"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for idx, lora_diff_names in enumerate(self.lora_diffs):
lora_strength = self._get_lora_strength(idx)
if isinstance(lora_diff_names, tuple):
lora_diff_0 = getattr(self, lora_diff_names[0])
lora_diff_1 = getattr(self, lora_diff_names[1])
lora_diff_2 = getattr(self, lora_diff_names[2])
weight = self._apply_lora_impl(
weight, lora_diff_0, lora_diff_1,
float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
)
else:
lora_diff = getattr(self, lora_diff_names)
weight = self._apply_single_lora_impl(weight, lora_diff, lora_strength)
return weight
def _prepare_weight(self, input):
"""Prepare weight tensor - handles both regular and GGUF weights"""
if self.is_gguf:
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
else:
weight = self.weight.to(input)
return weight
def forward(self, input):
weight = self._prepare_weight(input)
if self.bias is not None:
bias = self.bias.to(input if not self.is_gguf else self.compute_dtype)
else:
bias = None
# Only apply scale_weight for non-GGUF models
if not self.is_gguf and self.scale_weight is not None:
if weight.numel() < input.numel():
weight = weight * self.scale_weight
else:
input = input * self.scale_weight
weight = self._get_weight_with_lora(weight)
out = self._linear_forward_impl(input, weight, bias)
del weight, input, bias
return out
def update_lora_step(module, step):
for name, submodule in module.named_modules():
if isinstance(submodule, CustomLinear) and hasattr(submodule, "_step"):
submodule._step.fill_(step)
def remove_lora_from_module(module):
for name, submodule in module.named_modules():
if hasattr(submodule, "lora_diffs"):
for i in range(len(submodule.lora_diffs)):
if hasattr(submodule, f"lora_diff_{i}_0"):
delattr(submodule, f"lora_diff_{i}_0")
if hasattr(submodule, f"lora_diff_{i}_1"):
delattr(submodule, f"lora_diff_{i}_1")
if hasattr(submodule, f"lora_diff_{i}_2"):
delattr(submodule, f"lora_diff_{i}_2")
+1 -1
View File
@@ -75,7 +75,7 @@ def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict,
for name, module in model.named_children(): for name, module in model.named_children():
for source_module, target_module in module_map.items(): for source_module, target_module in module_map.items():
if isinstance(module, source_module): if isinstance(module, source_module):
if "rope_embedder" in name or "patch_embedding" in name or "emb_pos" in name: if "rope_embedder" in name or "patch_embedding" in name:
continue continue
num_param = sum(p.numel() for p in module.parameters()) num_param = sum(p.numel() for p in module.parameters())
-104
View File
@@ -1,104 +0,0 @@
import torch
from comfy.model_management import get_autocast_device, get_torch_device
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_z(x, grid_sizes, freqs, inner_t, shift=6):
n, c = x.size(2), x.size(3) // 2
# loop over samples
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(
x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)
)
start_ind = [sum(inner_t[i][:_]) for _ in range(len(inner_t[i]))]
end_ind = [sum(inner_t[i][:_+1]) for _ in range(len(inner_t[i]))]
freq_select = []
for shot_ind, (s, e) in enumerate(zip(start_ind, end_ind)):
freq_select += [shot_ind * shift] * (e - s)
shot_freqs = freqs[freq_select]
freqs_i = shot_freqs.view(f, 1, 1, -1).expand(f, h, w, -1).reshape(seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
# append to collection
output.append(x_i)
return torch.stack(output).float()
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_c(x, freqs, inner_c, shift=6):
b, s, n, c = x.size(0), x.size(1), x.size(2), x.size(3) // 2
# loop over samples
output = []
for i in range(b):
# precompute multipliers
x_i = torch.view_as_complex(
x[i].to(torch.float64).reshape(s, n, -1, 2)
)
freq_select = []
for shot_ind, c_len in enumerate(inner_c[i]):
freq_select += [shot_ind * shift] * c_len
freq_select += [shot_ind+10] * (s-len(freq_select)) # extra suppression for the empty token
shot_freqs = freqs[freq_select]
freqs_i = shot_freqs.view(s, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
# append to collection
output.append(x_i)
return torch.stack(output).float()
@torch.autocast(device_type=get_autocast_device(get_torch_device()), enabled=False)
@torch.compiler.disable()
def rope_apply_echoshot(x, grid_sizes, freqs, inner_t, shift=4):
n, c = x.size(2), x.size(3) // 2
# split freqs
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
# loop over samples
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
# precompute multipliers
x_i = torch.view_as_complex(
x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2)
)
start_ind = [sum(inner_t[i][:_]) for _ in range(len(inner_t[i]))]
end_ind = [sum(inner_t[i][:_+1]) for _ in range(len(inner_t[i]))]
freq_select = []
for shot_ind, (s, e) in enumerate(zip(start_ind, end_ind)):
freq_select += list(range(shot_ind * shift + s, shot_ind * shift + e))
t_freqs = freqs[0][freq_select]
freqs_i = torch.cat([
# freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
t_freqs.view(f, 1, 1, -1).expand(f, h, w, -1), ###
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
], dim=-1).reshape(seq_len, 1, -1)
# apply rotary embedding
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
# append to collection
output.append(x_i)
return torch.stack(output).float()
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
Binary file not shown.
Binary file not shown.

Before

Width:  |  Height:  |  Size: 192 KiB

Binary file not shown.
File diff suppressed because it is too large Load Diff
@@ -1,8 +1,8 @@
{ {
"id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1", "id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1",
"revision": 0, "revision": 0,
"last_node_id": 206, "last_node_id": 204,
"last_link_id": 341, "last_link_id": 336,
"nodes": [ "nodes": [
{ {
"id": 42, "id": 42,
@@ -101,7 +101,7 @@
200 200
], ],
"flags": {}, "flags": {},
"order": 22, "order": 21,
"mode": 2, "mode": 2,
"inputs": [ "inputs": [
{ {
@@ -258,7 +258,7 @@
174 174
], ],
"flags": {}, "flags": {},
"order": 39, "order": 34,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -371,7 +371,7 @@
200 200
], ],
"flags": {}, "flags": {},
"order": 23, "order": 22,
"mode": 2, "mode": 2,
"inputs": [ "inputs": [
{ {
@@ -413,7 +413,7 @@
86 86
], ],
"flags": {}, "flags": {},
"order": 25, "order": 24,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -647,7 +647,9 @@
"ver": "0.3.27", "ver": "0.3.27",
"Node name for S&R": "PreviewImage" "Node name for S&R": "PreviewImage"
}, },
"widgets_values": [] "widgets_values": [
""
]
}, },
{ {
"id": 125, "id": 125,
@@ -661,12 +663,15 @@
190.28567504882812 190.28567504882812
], ],
"flags": {}, "flags": {},
"order": 41, "order": 40,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
"name": "text", "name": "text",
"type": "STRING", "type": "STRING",
"widget": {
"name": "text"
},
"link": 215 "link": 215
} }
], ],
@@ -684,7 +689,7 @@
"Node name for S&R": "ShowText|pysssss" "Node name for S&R": "ShowText|pysssss"
}, },
"widgets_values": [ "widgets_values": [
"A man in a suit and tie walking down a hallway. He has a friendly expression and is looking directly at the camera. The hallway has beige walls adorned with framed black and white photographs. There is a door on the left side of the hallway and a poster on the wall. The lighting is soft and natural. The image is high quality and has a watermark in the bottom right corner.", "",
"A man in a suit and tie walking down a hallway. He has a friendly expression and is looking directly at the camera. The hallway has beige walls adorned with framed black and white photographs. There is a door on the left side of the hallway and a poster on the wall. The lighting is soft and natural. The image is high quality and has a watermark in the bottom right corner." "A man in a suit and tie walking down a hallway. He has a friendly expression and is looking directly at the camera. The hallway has beige walls adorned with framed black and white photographs. There is a door on the left side of the hallway and a poster on the wall. The lighting is soft and natural. The image is high quality and has a watermark in the bottom right corner."
] ]
}, },
@@ -740,7 +745,7 @@
261.5306701660156 261.5306701660156
], ],
"flags": {}, "flags": {},
"order": 42, "order": 41,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -837,7 +842,7 @@
"flags": { "flags": {
"collapsed": true "collapsed": true
}, },
"order": 24, "order": 23,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -909,7 +914,7 @@
266 266
], ],
"flags": {}, "flags": {},
"order": 29, "order": 28,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -917,22 +922,28 @@
"type": "IMAGE", "type": "IMAGE",
"link": 244 "link": 244
}, },
{
"name": "get_image_size",
"shape": 7,
"type": "IMAGE",
"link": null
},
{ {
"name": "width_input", "name": "width_input",
"shape": 7, "shape": 7,
"type": "INT", "type": "INT",
"widget": {
"name": "width_input"
},
"link": null "link": null
}, },
{ {
"name": "height_input", "name": "height_input",
"shape": 7, "shape": 7,
"type": "INT", "type": "INT",
"link": null "widget": {
}, "name": "height_input"
{ },
"name": "get_image_size",
"shape": 7,
"type": "IMAGE",
"link": null "link": null
} }
], ],
@@ -967,6 +978,8 @@
"lanczos", "lanczos",
false, false,
16, 16,
0,
0,
"center" "center"
] ]
}, },
@@ -1012,6 +1025,50 @@
"color": "#2a363b", "color": "#2a363b",
"bgcolor": "#3f5159" "bgcolor": "#3f5159"
}, },
{
"id": 74,
"type": "WidgetToString",
"pos": [
2128.5166015625,
-432.6599426269531
],
"size": [
315,
154
],
"flags": {},
"order": 29,
"mode": 0,
"inputs": [
{
"name": "any_input",
"shape": 7,
"type": "*",
"link": 102
}
],
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
212
]
}
],
"properties": {
"cnr_id": "comfyui-kjnodes",
"ver": "d57154c3a808b8a3f232ed293eaa2d000867c884",
"Node name for S&R": "WidgetToString"
},
"widgets_values": [
0,
"camera_type",
false,
"",
2
]
},
{ {
"id": 58, "id": 58,
"type": "WanVideoEncode", "type": "WanVideoEncode",
@@ -1082,7 +1139,7 @@
"flags": { "flags": {
"collapsed": true "collapsed": true
}, },
"order": 43, "order": 42,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -1120,7 +1177,7 @@
555.8994140625 555.8994140625
], ],
"flags": {}, "flags": {},
"order": 40, "order": 35,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -1135,7 +1192,9 @@
"ver": "0.3.27", "ver": "0.3.27",
"Node name for S&R": "PreviewImage" "Node name for S&R": "PreviewImage"
}, },
"widgets_values": [] "widgets_values": [
""
]
}, },
{ {
"id": 128, "id": 128,
@@ -1233,7 +1292,7 @@
274 274
], ],
"flags": {}, "flags": {},
"order": 44, "order": 39,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -1245,6 +1304,9 @@
"name": "caption", "name": "caption",
"shape": 7, "shape": 7,
"type": "STRING", "type": "STRING",
"widget": {
"name": "caption"
},
"link": null "link": null
}, },
{ {
@@ -1279,7 +1341,8 @@
"black", "black",
"FreeMonoBoldOblique.otf", "FreeMonoBoldOblique.otf",
"input", "input",
"up" "up",
""
] ]
}, },
{ {
@@ -1290,7 +1353,7 @@
-1155.6121826171875 -1155.6121826171875
], ],
"size": [ "size": [
421.6000061035156, 390.5999755859375,
202 202
], ],
"flags": {}, "flags": {},
@@ -1320,6 +1383,82 @@
128 128
] ]
}, },
{
"id": 138,
"type": "ReCamMasterPoseVisualizer",
"pos": [
1597.2598876953125,
177.93458557128906
],
"size": [
349.6756591796875,
130
],
"flags": {},
"order": 31,
"mode": 0,
"inputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"link": 242
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
243
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "2dc25c150ec4288e0b8689fc49fe2c9c8ab01999",
"Node name for S&R": "ReCamMasterPoseVisualizer"
},
"widgets_values": [
0.10000000000000002,
0.20000000000000004,
0.4000000000000001,
0.5000000000000001
]
},
{
"id": 157,
"type": "GetNode",
"pos": [
1394.278076171875,
-105.60762786865234
],
"size": [
210,
60
],
"flags": {
"collapsed": true
},
"order": 14,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
274
]
}
],
"title": "Get_InputLatents",
"properties": {},
"widgets_values": [
"InputLatents"
],
"color": "#323",
"bgcolor": "#535"
},
{ {
"id": 127, "id": 127,
"type": "WanVideoExperimentalArgs", "type": "WanVideoExperimentalArgs",
@@ -1332,7 +1471,7 @@
130 130
], ],
"flags": {}, "flags": {},
"order": 14, "order": 15,
"mode": 0, "mode": 0,
"inputs": [], "inputs": [],
"outputs": [ "outputs": [
@@ -1370,7 +1509,7 @@
"flags": { "flags": {
"collapsed": true "collapsed": true
}, },
"order": 15, "order": 16,
"mode": 0, "mode": 0,
"inputs": [], "inputs": [],
"outputs": [ "outputs": [
@@ -1404,7 +1543,7 @@
"flags": { "flags": {
"collapsed": true "collapsed": true
}, },
"order": 16, "order": 17,
"mode": 0, "mode": 0,
"inputs": [], "inputs": [],
"outputs": [ "outputs": [
@@ -1436,7 +1575,7 @@
46 46
], ],
"flags": {}, "flags": {},
"order": 28, "order": 27,
"mode": 2, "mode": 2,
"inputs": [ "inputs": [
{ {
@@ -1478,7 +1617,7 @@
"flags": { "flags": {
"collapsed": true "collapsed": true
}, },
"order": 45, "order": 44,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -1518,7 +1657,7 @@
"flags": { "flags": {
"collapsed": true "collapsed": true
}, },
"order": 17, "order": 18,
"mode": 0, "mode": 0,
"inputs": [], "inputs": [],
"outputs": [ "outputs": [
@@ -1550,7 +1689,7 @@
688.150634765625 688.150634765625
], ],
"flags": {}, "flags": {},
"order": 34, "order": 30,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -1587,9 +1726,9 @@
"link": null "link": null
}, },
{ {
"name": "cache_args", "name": "teacache_args",
"shape": 7, "shape": 7,
"type": "CACHEARGS", "type": "TEACACHEARGS",
"link": 335 "link": 335
}, },
{ {
@@ -1615,12 +1754,6 @@
"shape": 7, "shape": 7,
"type": "EXPERIMENTALARGS", "type": "EXPERIMENTALARGS",
"link": 334 "link": 334
},
{
"name": "sigmas",
"shape": 7,
"type": "SIGMAS",
"link": null
} }
], ],
"outputs": [ "outputs": [
@@ -1648,7 +1781,8 @@
0, 0,
1, 1,
false, false,
"comfy" "comfy",
""
] ]
}, },
{ {
@@ -1663,13 +1797,13 @@
178 178
], ],
"flags": {}, "flags": {},
"order": 18, "order": 19,
"mode": 0, "mode": 0,
"inputs": [], "inputs": [],
"outputs": [ "outputs": [
{ {
"name": "cache_args", "name": "teacache_args",
"type": "CACHEARGS", "type": "TEACACHEARGS",
"links": [ "links": [
335 335
] ]
@@ -1698,10 +1832,10 @@
], ],
"size": [ "size": [
908.9017944335938, 908.9017944335938,
334 912.1107788085938
], ],
"flags": {}, "flags": {},
"order": 46, "order": 43,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -1766,6 +1900,53 @@
} }
} }
}, },
{
"id": 56,
"type": "WanVideoReCamMasterCameraEmbed",
"pos": [
1379.0372314453125,
-36.21258544921875
],
"size": [
356.0601806640625,
79.98188781738281
],
"flags": {},
"order": 25,
"mode": 0,
"inputs": [
{
"name": "latents",
"type": "LATENT",
"link": 274
}
],
"outputs": [
{
"name": "camera_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"links": [
102,
272
]
},
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": [
242
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "11e9166d0b00fe3b1e6ebb0a3d1db50ce7a56d58",
"Node name for S&R": "WanVideoReCamMasterCameraEmbed"
},
"widgets_values": [
"arc_right"
]
},
{ {
"id": 22, "id": 22,
"type": "WanVideoModelLoader", "type": "WanVideoModelLoader",
@@ -1778,7 +1959,7 @@
234 234
], ],
"flags": {}, "flags": {},
"order": 19, "order": 20,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
@@ -1836,248 +2017,6 @@
], ],
"color": "#223", "color": "#223",
"bgcolor": "#335" "bgcolor": "#335"
},
{
"id": 138,
"type": "ReCamMasterPoseVisualizer",
"pos": [
1758.8209228515625,
168.25579833984375
],
"size": [
349.6756591796875,
130
],
"flags": {},
"order": 35,
"mode": 0,
"inputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"link": 242
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
243
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "2dc25c150ec4288e0b8689fc49fe2c9c8ab01999",
"Node name for S&R": "ReCamMasterPoseVisualizer"
},
"widgets_values": [
0.10000000000000002,
0.20000000000000004,
0.4000000000000001,
0.5000000000000001
]
},
{
"id": 157,
"type": "GetNode",
"pos": [
1339.9281005859375,
-96.67338562011719
],
"size": [
210,
60
],
"flags": {
"collapsed": true
},
"order": 20,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "LATENT",
"type": "LATENT",
"links": [
274,
338
]
}
],
"title": "Get_InputLatents",
"properties": {},
"widgets_values": [
"InputLatents"
],
"color": "#323",
"bgcolor": "#535"
},
{
"id": 56,
"type": "WanVideoReCamMasterCameraEmbed",
"pos": [
1338.8333740234375,
-36.212589263916016
],
"size": [
356.0601806640625,
79.98188781738281
],
"flags": {},
"order": 30,
"mode": 0,
"inputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"link": 340
},
{
"name": "latents",
"type": "LATENT",
"link": 274
}
],
"outputs": [
{
"name": "camera_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"links": [
272
]
},
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": [
242
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "11e9166d0b00fe3b1e6ebb0a3d1db50ce7a56d58",
"Node name for S&R": "WanVideoReCamMasterCameraEmbed"
},
"widgets_values": []
},
{
"id": 206,
"type": "WanVideoReCamMasterGenerateOrbitCamera",
"pos": [
1305.5355224609375,
185.5100555419922
],
"size": [
384.8144836425781,
82
],
"flags": {},
"order": 21,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": []
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "8257cd1f8abaa6504248b946f31c5173c0228b3d",
"Node name for S&R": "WanVideoReCamMasterGenerateOrbitCamera"
},
"widgets_values": [
81,
90
]
},
{
"id": 74,
"type": "WidgetToString",
"pos": [
2128.5166015625,
-432.6599426269531
],
"size": [
315,
154
],
"flags": {},
"order": 31,
"mode": 0,
"inputs": [
{
"name": "any_input",
"shape": 7,
"type": "*",
"link": 341
}
],
"outputs": [
{
"name": "STRING",
"type": "STRING",
"links": [
212
]
}
],
"properties": {
"cnr_id": "comfyui-kjnodes",
"ver": "d57154c3a808b8a3f232ed293eaa2d000867c884",
"Node name for S&R": "WidgetToString"
},
"widgets_values": [
0,
"camera_type",
false,
"",
2
]
},
{
"id": 205,
"type": "WanVideoReCamMasterDefaultCamera",
"pos": [
1317.4481201171875,
-241.10047912597656
],
"size": [
388.8835754394531,
58
],
"flags": {},
"order": 27,
"mode": 0,
"inputs": [
{
"name": "latents",
"type": "LATENT",
"link": 338
}
],
"outputs": [
{
"name": "camera_poses",
"type": "CAMERAPOSES",
"links": [
340,
341
]
}
],
"properties": {
"cnr_id": "ComfyUI-WanVideoWrapper",
"ver": "8257cd1f8abaa6504248b946f31c5173c0228b3d",
"Node name for S&R": "WanVideoReCamMasterDefaultCamera"
},
"widgets_values": [
"pan_right"
]
} }
], ],
"links": [ "links": [
@@ -2121,6 +2060,14 @@
1, 1,
"CONDITIONING" "CONDITIONING"
], ],
[
102,
56,
0,
74,
0,
"*"
],
[ [
210, 210,
28, 28,
@@ -2310,7 +2257,7 @@
157, 157,
0, 0,
56, 56,
1, 0,
"LATENT" "LATENT"
], ],
[ [
@@ -2344,30 +2291,6 @@
155, 155,
6, 6,
"TEACACHEARGS" "TEACACHEARGS"
],
[
338,
157,
0,
205,
0,
"LATENT"
],
[
340,
205,
0,
56,
0,
"CAMERAPOSES"
],
[
341,
205,
0,
74,
0,
"*"
] ]
], ],
"groups": [ "groups": [
@@ -2414,12 +2337,13 @@
"config": {}, "config": {},
"extra": { "extra": {
"ds": { "ds": {
"scale": 0.611590904484162, "scale": 0.7400249944258357,
"offset": [ "offset": [
1176.5764579377562, 1311.6629502036258,
1095.9393193240473 1366.2150672288358
] ]
}, },
"linkExtensions": [],
"node_versions": { "node_versions": {
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd", "ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
"comfy-core": "0.3.26", "comfy-core": "0.3.26",
@@ -2428,8 +2352,7 @@
"VHS_latentpreview": true, "VHS_latentpreview": true,
"VHS_latentpreviewrate": 0, "VHS_latentpreviewrate": 0,
"VHS_MetadataImage": true, "VHS_MetadataImage": true,
"VHS_KeepIntermediate": true, "VHS_KeepIntermediate": true
"frontendVersion": "1.16.7"
}, },
"version": 0.4 "version": 0.4
} }
File diff suppressed because it is too large Load Diff
@@ -176,8 +176,8 @@
"inputs": [], "inputs": [],
"outputs": [ "outputs": [
{ {
"name": "cache_args", "name": "teacache_args",
"type": "CACHEARGS", "type": "TEACACHEARGS",
"links": [ "links": [
103 103
] ]
@@ -826,9 +826,9 @@
"link": null "link": null
}, },
{ {
"name": "cache_args", "name": "teacache_args",
"shape": 7, "shape": 7,
"type": "CACHEARGS", "type": "TEACACHEARGS",
"link": 103 "link": 103
}, },
{ {
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -794,8 +794,8 @@
"inputs": [], "inputs": [],
"outputs": [ "outputs": [
{ {
"name": "cache_args", "name": "teacache_args",
"type": "CACHEARGS", "type": "TEACACHEARGS",
"links": [ "links": [
56 56
] ]
@@ -1719,9 +1719,9 @@
"link": null "link": null
}, },
{ {
"name": "cache_args", "name": "teacache_args",
"shape": 7, "shape": 7,
"type": "CACHEARGS", "type": "TEACACHEARGS",
"link": 56 "link": 56
}, },
{ {

Some files were not shown because too many files have changed in this diff Show More