Merge branch 'main' into dev

This commit is contained in:
kijai
2025-08-23 21:00:45 +03:00
12 changed files with 525 additions and 802 deletions
+24 -78
View File
@@ -1,11 +1,8 @@
from diffusers import ModelMixin, ConfigMixin
from einops import rearrange, repeat
import torch
import torch.nn as nn
from ..wanvideo.modules.attention import attention
from comfy import model_management as mm
def timestep_transform(
t,
shift=5.0,
@@ -65,7 +62,6 @@ def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', att
x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype)
for class_idx, ref_target_mask in enumerate(ref_target_masks):
mm.soft_empty_cache()
ref_target_mask = ref_target_mask[None, None, None, ...]
x_ref_attnmap = x_ref_attn_map_source * ref_target_mask
x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens
@@ -78,13 +74,11 @@ def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', att
x_ref_attn_maps.append(x_ref_attnmap)
del attn
del x_ref_attn_map_source
mm.soft_empty_cache()
del attn, x_ref_attn_map_source
return torch.concat(x_ref_attn_maps, dim=0)
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2, enable_sp=False):
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2):
"""Args:
query (torch.tensor): B M H K
key (torch.tensor): B M H K
@@ -145,7 +139,7 @@ class RotaryPositionalEmbedding1D(nn.Module):
return x_.type_as(x)
class AudioProjModel(ModelMixin, ConfigMixin):
class AudioProjModel(nn.Module):
def __init__(
self,
seq_len=5,
@@ -217,11 +211,6 @@ class SingleStreamAttention(nn.Module):
encoder_hidden_states_dim: int,
num_heads: int,
qkv_bias: bool,
qk_norm: bool,
norm_layer: nn.Module,
attn_drop: float = 0.0,
proj_drop: float = 0.0,
eps: float = 1e-6,
attention_mode: str = 'sdpa',
) -> None:
super().__init__()
@@ -230,74 +219,46 @@ class SingleStreamAttention(nn.Module):
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.qk_norm = qk_norm
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
self.q_norm = norm_layer(self.head_dim, eps=eps) if qk_norm else nn.Identity()
self.k_norm = norm_layer(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.add_q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor:
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
self.proj = nn.Linear(dim, dim)
self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None) -> torch.Tensor:
N_t, N_h, N_w = shape
expected_tokens = N_t * N_h * N_w
actual_tokens = x.shape[1]
x_extra = None
if x.shape[0] * N_t != encoder_hidden_states.shape[0]:
if actual_tokens != expected_tokens:
x_extra = x[:, -N_h * N_w:, :]
x = x[:, :-N_h * N_w, :]
N_t = N_t - 1
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
B = x.shape[0]
S = N_h * N_w
x = x.view(B * N_t, S, self.dim)
# 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))
if self.qk_norm:
q = self.q_norm(q)
q = self.q_linear(x).view(B * N_t, S, self.num_heads, self.head_dim)
# get kv from encoder_hidden_states
_, N_a, _ = encoder_hidden_states.shape
encoder_kv = self.kv_linear(encoder_hidden_states)
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)
# get kv from encoder_hidden_states # shape: (B, N, num_heads, head_dim)
kv = self.kv_linear(encoder_hidden_states)
encoder_k, encoder_v = kv.view(B * N_t, encoder_hidden_states.shape[1], 2, self.num_heads, self.head_dim).unbind(2)
if self.qk_norm:
encoder_k = self.add_k_norm(encoder_k)
x = attention(
q.transpose(1, 2),
encoder_k.transpose(1, 2),
encoder_v.transpose(1, 2),
attention_mode=self.attention_mode
)
x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode)
# 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)
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
x = self.proj(x.reshape(B * N_t, S, self.dim))
x = x.view(B, N_t * S, self.dim)
if x_extra is not None:
x = torch.cat([x, torch.zeros_like(x_extra)], dim=1)
return x
class SingleStreamMultiAttention(SingleStreamAttention):
"""Multi-speaker rotary-position cross-attention.
@@ -315,11 +276,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
encoder_hidden_states_dim: int,
num_heads: int,
qkv_bias: bool,
qk_norm: bool,
norm_layer: nn.Module,
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',
@@ -329,11 +285,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
encoder_hidden_states_dim=encoder_hidden_states_dim,
num_heads=num_heads,
qkv_bias=qkv_bias,
qk_norm=qk_norm,
norm_layer=norm_layer,
attn_drop=attn_drop,
proj_drop=proj_drop,
eps=eps,
attention_mode=attention_mode,
)
@@ -376,8 +327,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
B, N, C = x.shape
q = self.q_linear(x)
q = q.view(B, N, self.num_heads, self.head_dim).permute(0, 2, 1, 3)
if self.qk_norm:
q = self.q_norm(q)
if human_num == 2:
# Use `class_range` logic for exactly 2 speakers
@@ -441,8 +390,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
encoder_kv = self.kv_linear(encoder_hidden_states)
encoder_kv = encoder_kv.view(B, N_a, 2, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)
encoder_k, encoder_v = encoder_kv.unbind(0)
if self.qk_norm:
encoder_k = self.add_k_norm(encoder_k)
# Rotary for keys – assign centre of each speaker bucket to its context tokens
if human_num == 2:
@@ -478,7 +425,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
# Linear projection
x = x.reshape(B, N, C)
x = self.proj(x)
x = self.proj_drop(x)
# Restore original layout
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
-1
View File
@@ -2,7 +2,6 @@ import folder_paths
from comfy import model_management as mm
from comfy.utils import load_torch_file, common_upscale
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
import torch
from ..utils import log
+140 -152
View File
@@ -1576,7 +1576,37 @@ class WanVideoScheduler: #WIP
def process(self, scheduler):
return (scheduler,)
rope_functions = ["default", "comfy", "comfy_chunked"]
class WanVideoRoPEFunction: #WIP
@classmethod
def INPUT_TYPES(s):
return {"required": {
"rope_function": (rope_functions, {"default": "comfy"}),
"ntk_scale_f": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}),
"ntk_scale_h": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}),
"ntk_scale_w": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}),
},
}
RETURN_TYPES = (rope_functions, )
RETURN_NAMES = ("rope_function",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def process(self, rope_function, ntk_scale_f, ntk_scale_h, ntk_scale_w):
if ntk_scale_f != 1.0 or ntk_scale_h != 1.0 or ntk_scale_w != 1.0:
rope_func_dict = {
"rope_function": rope_function,
"ntk_scale_f": ntk_scale_f,
"ntk_scale_h": ntk_scale_h,
"ntk_scale_w": ntk_scale_w,
}
return (rope_func_dict,)
return (rope_function,)
#region Sampler
class WanVideoSampler:
@classmethod
@@ -1603,7 +1633,7 @@ class WanVideoSampler:
"flowedit_args": ("FLOWEDITARGS", ),
"batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Batch cond and uncond for faster sampling, possibly faster on some hardware, uses more memory"}),
"slg_args": ("SLGARGS", ),
"rope_function": (["default", "comfy", "comfy_chunked"], {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile. Chunked version has reduced peak VRAM usage when not using torch.compile"}),
"rope_function": (rope_functions, {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile. Chunked version has reduced peak VRAM usage when not using torch.compile"}),
"loop_args": ("LOOPARGS", ),
"experimental_args": ("EXPERIMENTALARGS", ),
"sigmas": ("SIGMAS", ),
@@ -1701,7 +1731,7 @@ class WanVideoSampler:
log.info(f"sigmas: {sample_scheduler.sigmas}")
else:
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
total_steps = steps
steps = len(timesteps)
if end_step != -1 and start_step >= end_step:
@@ -1713,8 +1743,6 @@ class WanVideoSampler:
start_step = steps - int(steps * denoise_strength) - 1
add_noise_to_samples = True #for now to not break old workflows
first_sampler = (end_step != -1 or end_step >= steps)
noise_pred_flipped = None
if isinstance(cfg, list):
@@ -1724,7 +1752,7 @@ class WanVideoSampler:
else:
cfg = [cfg] * (steps + 1)
if first_sampler:
if end_step != -1:
timesteps = timesteps[:end_step]
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
log.info(f"Sampling until step {end_step}, timestep: {timesteps[-1]}")
@@ -1964,14 +1992,15 @@ class WanVideoSampler:
dwpose_data = torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2)
dwpose_data = transformer.dwpose_embedding(dwpose_data)
log.info(f"UniAnimate pose embed shape: {dwpose_data.shape}")
if dwpose_data.shape[2] > latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is longer than the video length {latent_video_length}, truncating")
dwpose_data = dwpose_data[:,:, :latent_video_length]
elif dwpose_data.shape[2] < latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is shorter than the video length {latent_video_length}, padding with last pose")
pad_len = latent_video_length - dwpose_data.shape[2]
pad = dwpose_data[:,:,:1].repeat(1,1,pad_len,1,1)
dwpose_data = torch.cat([dwpose_data, pad], dim=2)
if not multitalk_sampling:
if dwpose_data.shape[2] > latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is longer than the video length {latent_video_length}, truncating")
dwpose_data = dwpose_data[:,:, :latent_video_length]
elif dwpose_data.shape[2] < latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is shorter than the video length {latent_video_length}, padding with last pose")
pad_len = latent_video_length - dwpose_data.shape[2]
pad = dwpose_data[:,:,:1].repeat(1,1,pad_len,1,1)
dwpose_data = torch.cat([dwpose_data, pad], dim=2)
dwpose_data_flat = rearrange(dwpose_data, 'b c f h w -> b (f h w) c').contiguous()
random_ref_dwpose_data = None
@@ -2118,6 +2147,7 @@ class WanVideoSampler:
mtv_freqs = mtv_freqs.to(device, dtype)
# vid2vid
noise_mask=original_image=None
if samples is not None and not multitalk_sampling:
saved_generator_state = samples.get("generator_state", None)
if saved_generator_state is not None:
@@ -2127,30 +2157,28 @@ class WanVideoSampler:
input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1)
if add_noise_to_samples:
latent_timestep = timesteps[:1].to(noise)
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
latent_timestep = timesteps[:1].to(noise)
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
noise = input_samples
noise_mask = samples.get("noise_mask", None)
if noise_mask is not None:
log.info(f"Latent noise_mask shape: {noise_mask.shape}")
original_image = input_samples.to(device)
original_image = samples.get("original_image", None)
if original_image is None:
original_image = input_samples
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
if noise_mask.shape[0] < noise.shape[1]:
noise_mask = noise_mask.repeat(noise.shape[1] // noise_mask.shape[0], 1, 1)
noise_mask = torch.nn.functional.interpolate(
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
mode='trilinear',
align_corners=False
).squeeze(0) # Remove batch dim, keep channel dim
# Add batch & channel dims for final output
noise_mask = noise_mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1)
if noise_mask.shape[2] != noise.shape[1]:
noise_mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - noise_mask.shape[2], noise.shape[2], noise.shape[3]), noise_mask], dim=2)
).repeat(1, noise.shape[0], 1, 1, 1)
# extra latents (Pusa) and 5b
latents_to_insert = add_index = None
@@ -2196,10 +2224,7 @@ class WanVideoSampler:
set_enhance_weight(feta_args["weight"])
feta_start_percent = feta_args["start_percent"]
feta_end_percent = feta_args["end_percent"]
if context_options is not None:
set_num_frames(context_frames)
else:
set_num_frames(latent_video_length)
set_num_frames(latent_video_length) if context_options is None else set_num_frames(context_frames)
enhance_enabled = True
else:
feta_args = None
@@ -2349,6 +2374,11 @@ class WanVideoSampler:
sample_scheduler_flipped = copy.deepcopy(sample_scheduler)
#rope
ntk_alphas = [1.0, 1.0, 1.0]
if isinstance(rope_function, dict):
ntk_alphas = rope_function["ntk_scale_f"], rope_function["ntk_scale_h"], rope_function["ntk_scale_w"]
rope_function = rope_function["rope_function"]
freqs = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
@@ -2548,6 +2578,7 @@ class WanVideoSampler:
"fantasy_portrait_input": fantasy_portrait_input,
"phantom_ref": phantom_ref,
"reverse_time": reverse_time,
"ntk_alphas": ntk_alphas,
"mtv_motion_tokens": mtv_motion_tokens if mtv_input is not None else None,
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None,
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0,
@@ -2667,13 +2698,13 @@ class WanVideoSampler:
raise e
#https://github.com/WeichenFan/CFG-Zero-star/
alpha = 1.0
if use_cfg_zero_star:
alpha = optimized_scale(
noise_pred_cond.view(batch_size, -1),
noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 1)
else:
alpha = 1.0
noise_pred_uncond_scaled = noise_pred_uncond * alpha
@@ -2712,19 +2743,16 @@ class WanVideoSampler:
intermediate_device = device
# diff diff prep
# Differential diffusion prep
masks = None
if not multitalk_sampling and samples is not None and noise_mask is not None:
noise_mask = 1 - noise_mask
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = masks > thresholds
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
latent_shift_loop = False
if loop_args is not None:
latent_shift_loop = True
is_looped = True
latent_shift_loop = is_looped = True
latent_skip = loop_args["shift_skip"]
latent_shift_start_percent = loop_args["start_percent"]
latent_shift_end_percent = loop_args["end_percent"]
@@ -2770,10 +2798,8 @@ class WanVideoSampler:
# Generate new random noise
z_rand = torch.randn(z_T.shape, dtype=torch.float32, generator=seed_g, device=torch.device("cpu"))
# Apply frequency mixing
current_latent = freq_mix_3d(z_T.to(torch.float32), z_rand.to(device), LPF=freq_filter)
current_latent = current_latent.to(dtype)
current_latent = (freq_mix_3d(z_T.to(torch.float32), z_rand.to(device), LPF=freq_filter)).to(dtype)
# Store initial noise for first iteration
if freeinit_args is not None and iter_idx == 0:
@@ -3157,7 +3183,6 @@ class WanVideoSampler:
}
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
loop_pbar = tqdm(total=estimated_iterations, desc="Total progress", position=1, leave=True)
callback = prepare_callback(patcher, estimated_iterations)
audio_embedding = multitalk_audio_embedding
@@ -3224,6 +3249,8 @@ class WanVideoSampler:
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
if noise_mask is not None:
original_image = input_samples.to(device)
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
@@ -3234,13 +3261,24 @@ class WanVideoSampler:
noise = input_samples
# diff diff prep
masks = None
noise_mask = samples.get("noise_mask", None)
if noise_mask is not None:
noise_mask = 1 - noise_mask
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
if noise_mask.shape[0] < noise.shape[1]:
noise_mask = noise_mask.repeat(noise.shape[1] // noise_mask.shape[0], 1, 1)
else:
noise_mask = noise_mask[latent_start_idx:latent_end_idx]
noise_mask = torch.nn.functional.interpolate(
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
mode='trilinear',
align_corners=False
).repeat(1, noise.shape[0], 1, 1, 1)
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = masks > thresholds
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
window_vace_data = None
if vace_data is not None:
@@ -3280,15 +3318,11 @@ class WanVideoSampler:
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)
if mode == "multitalk":
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
else:
if is_first_clip:
latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
else:
latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
latent_motion_frames = latent_motion_frames[0]
cond_ = cond_image if is_first_clip else cond_frame
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
vae.to(offload_device)
y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W
mm.soft_empty_cache()
@@ -3311,7 +3345,7 @@ class WanVideoSampler:
if start_step != 0:
raise ValueError("start_step must be 0 when denoise_strength is used")
start_step = steps - int(steps * denoise_strength) - 1
if (end_step != -1 or end_step >= steps):
if end_step != -1:
timesteps = timesteps[:end_step]
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
if start_step > 0:
@@ -3338,8 +3372,7 @@ class WanVideoSampler:
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
_, T_m, _, _ = add_latent.shape
latent[:, :T_m] = add_latent
latent[:, :add_latent.shape[1]] = add_latent
if offload:
#blockswap init
@@ -3382,6 +3415,18 @@ class WanVideoSampler:
else:
positive = text_embeds["prompt_embeds"]
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c')
partial_unianim_data = {
"dwpose": partial_dwpose_flat,
"random_ref": unianim_data["random_ref"],
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_percent"]
}
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
@@ -3390,44 +3435,32 @@ class WanVideoSampler:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input,
cfg[idx],
positive,
text_embeds["negative_prompt_embeds"],
timestep, idx, y, clip_embeds, control_latents, window_vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
latent_model_input, cfg[i], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs)
sampling_pbar.update(1)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
del callback_latent
sampling_pbar.update(1)
step_iteration_count += 1
# update latent
if scheduler == "multitalk":
noise_pred = -noise_pred
dt = timesteps[i] - timesteps[i + 1]
dt = dt / 1000
dt = (timesteps[i] - timesteps[i + 1]) / 1000
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = latent.to(intermediate_device)
temp_x0 = sample_scheduler.step(
noise_pred.unsqueeze(0),
timestep,
latent.unsqueeze(0),
**scheduler_step_args)[0]
latent = temp_x0.squeeze(0)
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
del noise_pred, latent_model_input, timestep
# differential diffusion inpaint
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image, torch.tensor([noise_timestep]), noise.to(device)
)
mask = masks[idx].to(latent)
if i < len(timesteps) - 1:
image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1])
mask = masks[i].to(latent)
latent = image_latent * mask + latent * (1-mask)
# injecting motion frames
@@ -3435,62 +3468,56 @@ class WanVideoSampler:
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
_, T_m, _, _ = add_latent.shape
latent[:, :T_m] = add_latent
latent[:, :add_latent.shape[1]] = add_latent
else:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
x0 = latent.to(device)
del latent_model_input, timestep
del noise, y, msk, latent_motion_frames
if offload:
offload_transformer(transformer)
vae.to(device)
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae, pbar=False)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device)
sampling_pbar.close()
# cache generated samples
videos = torch.stack(videos).cpu() # B C T H W
# optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled":
videos = videos[0].permute(1, 2, 3, 0).cpu().float().numpy()
videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher
cm = ColorMatcher()
cm_result_list = []
for img in videos:
if mode == "multitalk":
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch)
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result))
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
videos = torch.stack(cm_result_list, dim=0).to(torch.float32).permute(3, 0, 1, 2).unsqueeze(0)
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
# cache generated samples
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
if is_first_clip:
gen_video_list.append(videos)
else:
gen_video_list.append(videos[:, :, cur_motion_frames_num:])
current_condframe_index += 1
iteration_count += 1
# decide whether is done
if arrive_last_frame:
loop_pbar.update(estimated_iterations - iteration_count)
loop_pbar.close()
break
# update next condition frames
is_first_clip = False
cur_motion_frames_num = motion_frame
cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0)
if mode == "infinitetalk":
cond_frame = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
cond_frame = cond_
else:
cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
cond_image = cond_
# Update progress bar
iteration_count += 1
loop_pbar.update(1)
del videos, latent
mm.soft_empty_cache()
# Repeat audio emb
if multitalk_embeds is not None:
@@ -3515,9 +3542,8 @@ class WanVideoSampler:
miss_length = 1
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
gen_video_samples = torch.cat(gen_video_list, dim=2).to(torch.float32)
del noise, latent
gen_video_samples = torch.cat(gen_video_list, dim=1)
if force_offload:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
@@ -3526,7 +3552,7 @@ class WanVideoSampler:
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return {"video": gen_video_samples[0].permute(1, 2, 3, 0).cpu()},
return {"video": gen_video_samples.permute(1, 2, 3, 0)},
#region normal inference
else:
@@ -3607,7 +3633,7 @@ class WanVideoSampler:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image, torch.tensor([noise_timestep]), noise.to(device)
original_image.to(device), torch.tensor([noise_timestep]), noise.to(device)
)
mask = masks[idx].to(latent)
latent = image_latent * mask + latent * (1-mask)
@@ -3659,6 +3685,7 @@ class WanVideoSampler:
"has_ref": has_ref,
"drop_last": drop_last,
"generator_state": seed_g.get_state(),
"original_image": original_image.cpu() if original_image is not None else None
},{
"samples": callback_latent.unsqueeze(0).cpu() if callback is not None else None,
})
@@ -3704,9 +3731,9 @@ class WanVideoDecode:
mm.soft_empty_cache()
video = samples.get("video", None)
if video is not None:
video = torch.clamp(video, -1.0, 1.0)
video = (video + 1.0) / 2.0
return video.cpu(),
video.clamp_(-1.0, 1.0)
video.add_(1.0).div_(2.0)
return video.cpu().float(),
latents = samples["samples"]
end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False)
@@ -3819,45 +3846,6 @@ class WanVideoEncode:
return ({"samples": latents, "noise_mask": mask},)
class WanVideoLatentReScale:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"samples": ("LATENT",),
"direction": (["comfy_to_wrapper", "wrapper_to_comfy"], {"tooltip": "Direction to rescale latents, from comfy to wrapper or vice versa"}),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "encode"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Rescale latents to match the expected range for encoding or decoding. Can be used to "
def encode(self, samples, direction):
samples = samples.copy()
latents = samples["samples"]
mean = [
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
]
std = [
2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743,
3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160
]
mean = torch.tensor(mean).view(1, latents.shape[1], 1, 1, 1)
std = torch.tensor(std).view(1, latents.shape[1], 1, 1, 1)
inv_std = (1.0 / std).view(1, latents.shape[1], 1, 1, 1)
if direction == "comfy_to_wrapper":
latents = (latents - mean.to(latents)) * inv_std.to(latents)
elif direction == "wrapper_to_comfy":
latents = latents / inv_std.to(latents) + mean.to(latents)
samples["samples"] = latents
return (samples,)
NODE_CLASS_MAPPINGS = {
"WanVideoSampler": WanVideoSampler,
"WanVideoDecode": WanVideoDecode,
@@ -3886,11 +3874,11 @@ NODE_CLASS_MAPPINGS = {
"WanVideoBlockList": WanVideoBlockList,
"WanVideoTextEncodeCached": WanVideoTextEncodeCached,
"WanVideoAddExtraLatent": WanVideoAddExtraLatent,
"WanVideoLatentReScale": WanVideoLatentReScale,
"WanVideoScheduler": WanVideoScheduler,
"WanVideoAddStandInLatent": WanVideoAddStandInLatent,
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,,
"WanVideoRoPEFunction": WanVideoRoPEFunction
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -3921,8 +3909,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoBlockList": "WanVideo Block List",
"WanVideoTextEncodeCached": "WanVideo TextEncode Cached",
"WanVideoAddExtraLatent": "WanVideo Add Extra Latent",
"WanVideoLatentReScale": "WanVideo Latent ReScale",
"WanVideoAddStandInLatent": "WanVideo Add StandIn Latent",
"WanVideoAddControlEmbeds": "WanVideo Add Control Embeds",
"WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion"
"WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion",
"WanVideoRoPEFunction": "WanVideo RoPE Function"
}
+169 -44
View File
@@ -1221,7 +1221,7 @@ class WanVideoModelLoader:
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
# init audio module
from .multitalk.multitalk import SingleStreamMultiAttention
from .wanvideo.modules.model import WanRMSNorm, WanLayerNorm
from .wanvideo.modules.model import WanLayerNorm
norm_input_visual = True #dunno what this is
for block in transformer.blocks:
@@ -1229,10 +1229,7 @@ class WanVideoModelLoader:
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
qk_norm=False,
qkv_bias=True,
eps=transformer.eps,
norm_layer=WanRMSNorm,
class_range=24,
class_interval=4,
attention_mode=attention_mode,
@@ -1269,50 +1266,178 @@ class WanVideoModelLoader:
patcher.model.is_patched = False
scale_weights = {}
if "fp8" in quantization:
for k, v in sd.items():
if k.endswith(".scale_weight"):
scale_weights[k] = v.to(base_dtype)
if "fp8_e4m3fn" in quantization:
weight_dtype = torch.float8_e4m3fn
elif "fp8_e5m2" in quantization:
weight_dtype = torch.float8_e5m2
else:
weight_dtype = base_dtype
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
control_lora = False
if not merge_loras and control_lora:
log.warning("Control-LoRA patching is only supported with merge_loras=True")
if lora is not None:
patcher, control_lora = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras)
if not gguf:
if merge_loras and lora is not None:
if not lora_low_mem_load:
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
if control_lora:
patch_control_lora(patcher.model.diffusion_model, device)
patcher.model.is_patched = True
log.info("Merging LoRA to the model...")
patcher = apply_lora(
patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights,)
if not control_lora:
scale_weights.clear()
patcher.patches.clear()
transformer.patched_linear = False
else:
if "fp8" in quantization:
for k, v in sd.items():
if k.endswith(".scale_weight"):
scale_weights[k] = v
if not merge_loras:
from .custom_linear import _replace_linear
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
transformer.patched_linear = True
if "fp8_e4m3fn" in quantization:
dtype = torch.float8_e4m3fn
elif "fp8_e5m2" in quantization:
dtype = torch.float8_e5m2
else:
dtype = base_dtype
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
if not lora_low_mem_load:
log.info("Using accelerate to load and assign model weights to device...")
param_count = sum(1 for _ in transformer.named_parameters())
pbar = ProgressBar(param_count)
cnt = 0
for name, param in tqdm(transformer.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}",
total=param_count,
leave=True):
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
dtype_to_use = dtype if sd[name].dtype == dtype else dtype_to_use
if "modulation" in name or "norm" in name or "bias" in name:
dtype_to_use = base_dtype
if "patch_embedding" in name:
dtype_to_use = torch.float32
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
cnt += 1
if cnt % 100 == 0:
pbar.update(100)
#for name, param in transformer.named_parameters():
# print(name, param.dtype, param.device, param.shape)
pbar.update_absolute(param_count)
pbar.update_absolute(0)
comfy_model.diffusion_model = transformer
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
patcher.model.is_patched = False
control_lora = False
if lora is not None:
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
lora_strength = l["strength"]
if isinstance(lora_strength, list):
if merge_loras:
raise ValueError("LoRA strength should be a single value when merge_loras=True")
transformer.lora_scheduling_enabled = True
if lora_strength == 0:
log.warning(f"LoRA {lora_path} has strength 0, skipping...")
continue
lora_sd = load_torch_file(lora_path, safe_load=True)
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
from .unianimate.nodes import update_transformer
log.info("Unianimate LoRA detected, patching model...")
transformer = update_transformer(transformer, lora_sd)
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
if not any('img' in k for k in sd.keys()):
lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
#spacepxl's control LoRA patch
# for key in lora_sd.keys():
# print(key)
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
log.info("Control-LoRA detected, patching model...")
if not merge_loras:
log.warning("Control-LoRA patching is only supported with merge_loras=True, setting it to True")
merge_loras = True
control_lora = True
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
old_in_dim = transformer.in_dim # 16
new_in_dim = lora_sd["diffusion_model.patch_embedding.lora_A.weight"].shape[1]
assert new_in_dim == 32
new_in = in_cls(
new_in_dim,
transformer.patch_embedding.out_channels,
transformer.patch_embedding.kernel_size,
transformer.patch_embedding.stride,
transformer.patch_embedding.padding,
).to(device=device, dtype=torch.float32)
new_in.weight.zero_()
new_in.bias.zero_()
new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight)
new_in.bias.copy_(transformer.patch_embedding.bias)
transformer.patch_embedding = new_in
transformer.expanded_patch_embedding = new_in
if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
log.info("Stand-In LoRA detected")
for block in transformer.blocks:
block.self_attn.q_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
block.self_attn.k_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
block.self_attn.v_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
for lora in [block.self_attn.q_loras, block.self_attn.k_loras, block.self_attn.v_loras]:
for param in lora.parameters():
param.requires_grad = False
for name, param in transformer.named_parameters():
if "lora" in name:
param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype))
else:
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
del lora_sd
if not gguf and merge_loras:
log.info("Patching LoRA to the model...")
patcher = apply_lora(
patcher, device, transformer_load_device,
params_to_keep=params_to_keep, dtype=dtype, base_dtype=base_dtype, state_dict=sd,
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights)
scale_weights.clear()
patcher.patches.clear()
if gguf:
#from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
log.info("Using GGUF to load and assign model weights to device...")
param_count = sum(1 for _ in transformer.named_parameters())
out_features = sd["blocks.0.self_attn.k.weight"].shape[1]
patcher.model.diffusion_model = _replace_with_gguf_linear(patcher.model.diffusion_model, base_dtype, sd, patches=patcher.patches)
pbar = ProgressBar(param_count)
cnt = 0
for name, param in tqdm(patcher.model.diffusion_model.named_parameters(),
desc=f"Loading transformer parameters to {transformer_load_device}",
total=param_count,
leave=True):
if "loras" in name:
continue
#print(name, param.dtype, param.device, param.shape)
if isinstance(param, GGUFParameter):
dtype_to_use = torch.uint8
elif "patch_embedding" in name:
dtype_to_use = torch.float32
else:
dtype_to_use = base_dtype
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
cnt += 1
if cnt % 100 == 0:
pbar.update(100)
#for name, param in transformer.named_parameters():
# print(name, param.dtype, param.device, param.shape)
#patcher.load(device, full_load=True)
pbar.update_absolute(param_count)
patcher.model.is_patched = True
patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False)
if "fast" in quantization:
if lora is not None and not merge_loras:
raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs")
+146 -5
View File
@@ -2,6 +2,11 @@ import torch
import numpy as np
from comfy.utils import common_upscale
try:
from server import PromptServer
except:
PromptServer = None
VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2)
@@ -188,7 +193,10 @@ class CreateCFGScheduleFloatList:
"interpolation": (["linear", "ease_in", "ease_out"], {"default": "linear", "tooltip": "Interpolation method to use for the cfg scale"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "Start percent of the steps to apply cfg"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "End percent of the steps to apply cfg"}),
}
},
"hidden": {
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = ("FLOAT", )
@@ -197,8 +205,8 @@ class CreateCFGScheduleFloatList:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Helper node to generate a list of floats that can be used to schedule cfg scale for the steps, outside the set range cfg is set to 1.0"
def process(self, steps, cfg_scale_start, cfg_scale_end, interpolation, start_percent, end_percent):
def process(self, steps, cfg_scale_start, cfg_scale_end, interpolation, start_percent, end_percent, unique_id):
# Create a list of floats for the cfg schedule
cfg_list = [1.0] * steps
start_idx = min(int(steps * start_percent), steps - 1)
@@ -226,6 +234,78 @@ class CreateCFGScheduleFloatList:
if start_percent > 0:
cfg_list[0] = 1.0
if unique_id and PromptServer is not None:
try:
PromptServer.instance.send_progress_text(
f"{cfg_list}",
unique_id
)
except:
pass
return (cfg_list,)
class CreateScheduleFloatList:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"steps": ("INT", {"default": 30, "min": 2, "max": 1000, "step": 1, "tooltip": "Number of steps to schedule cfg for"} ),
"start_value": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}),
"end_value": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}),
"default_value": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 0.01, "round": 0.01, "tooltip": "Default value to use for the steps"}),
"interpolation": (["linear", "ease_in", "ease_out"], {"default": "linear", "tooltip": "Interpolation method to use for the cfg scale"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "Start percent of the steps to apply cfg"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "End percent of the steps to apply cfg"}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
},
}
RETURN_TYPES = ("FLOAT", )
RETURN_NAMES = ("float_list",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Helper node to generate a list of floats that can be used to schedule things like cfg and lora scale per step"
def process(self, steps, start_value, end_value, default_value,interpolation, start_percent, end_percent, unique_id):
# Create a list of floats for the cfg schedule
cfg_list = [default_value] * steps
start_idx = min(int(steps * start_percent), steps - 1)
end_idx = min(int(steps * end_percent), steps - 1)
for i in range(start_idx, end_idx + 1):
if i >= steps:
break
if end_idx == start_idx:
t = 0
else:
t = (i - start_idx) / (end_idx - start_idx)
if interpolation == "linear":
factor = t
elif interpolation == "ease_in":
factor = t * t
elif interpolation == "ease_out":
factor = t * (2 - t)
cfg_list[i] = round(start_value + factor * (end_value - start_value), 2)
# If start_percent > 0, always include the first step
if start_percent > 0:
cfg_list[0] = default_value
if unique_id and PromptServer is not None:
try:
PromptServer.instance.send_progress_text(
f"{cfg_list}",
unique_id
)
except:
pass
return (cfg_list,)
@@ -254,17 +334,78 @@ class DummyComfyWanModelObject:
return None
return (DummyModel(),)
class WanVideoLatentReScale:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"samples": ("LATENT",),
"direction": (["comfy_to_wrapper", "wrapper_to_comfy"], {"tooltip": "Direction to rescale latents, from comfy to wrapper or vice versa"}),
}
}
RETURN_TYPES = ("LATENT",)
RETURN_NAMES = ("samples",)
FUNCTION = "encode"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Rescale latents to match the expected range for encoding or decoding between native ComfyUI VAE and the WanVideoWrapper VAE."
def encode(self, samples, direction):
samples = samples.copy()
latents = samples["samples"]
if latents.shape[1] == 48:
mean = [
-0.2289, -0.0052, -0.1323, -0.2339, -0.2799, 0.0174, 0.1838, 0.1557,
-0.1382, 0.0542, 0.2813, 0.0891, 0.1570, -0.0098, 0.0375, -0.1825,
-0.2246, -0.1207, -0.0698, 0.5109, 0.2665, -0.2108, -0.2158, 0.2502,
-0.2055, -0.0322, 0.1109, 0.1567, -0.0729, 0.0899, -0.2799, -0.1230,
-0.0313, -0.1649, 0.0117, 0.0723, -0.2839, -0.2083, -0.0520, 0.3748,
0.0152, 0.1957, 0.1433, -0.2944, 0.3573, -0.0548, -0.1681, -0.0667,
]
std = [
0.4765, 1.0364, 0.4514, 1.1677, 0.5313, 0.4990, 0.4818, 0.5013,
0.8158, 1.0344, 0.5894, 1.0901, 0.6885, 0.6165, 0.8454, 0.4978,
0.5759, 0.3523, 0.7135, 0.6804, 0.5833, 1.4146, 0.8986, 0.5659,
0.7069, 0.5338, 0.4889, 0.4917, 0.4069, 0.4999, 0.6866, 0.4093,
0.5709, 0.6065, 0.6415, 0.4944, 0.5726, 1.2042, 0.5458, 1.6887,
0.3971, 1.0600, 0.3943, 0.5537, 0.5444, 0.4089, 0.7468, 0.7744
]
else:
mean = [
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
]
std = [
2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743,
3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160
]
mean = torch.tensor(mean).view(1, latents.shape[1], 1, 1, 1)
std = torch.tensor(std).view(1, latents.shape[1], 1, 1, 1)
inv_std = (1.0 / std).view(1, latents.shape[1], 1, 1, 1)
if direction == "comfy_to_wrapper":
latents = (latents - mean.to(latents)) * inv_std.to(latents)
elif direction == "wrapper_to_comfy":
latents = latents / inv_std.to(latents) + mean.to(latents)
samples["samples"] = latents
return (samples,)
NODE_CLASS_MAPPINGS = {
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
"ExtractStartFramesForContinuations": ExtractStartFramesForContinuations,
"CreateCFGScheduleFloatList": CreateCFGScheduleFloatList,
"DummyComfyWanModelObject": DummyComfyWanModelObject
"DummyComfyWanModelObject": DummyComfyWanModelObject,
"WanVideoLatentReScale": WanVideoLatentReScale,
"CreateScheduleFloatList": CreateScheduleFloatList
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
"WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame",
"ExtractStartFramesForContinuations": "Extract Start Frames For Continuations",
"CreateCFGScheduleFloatList": "Create CFG Schedule Float List",
"DummyComfyWanModelObject": "Dummy Comfy Wan Model Object"
"DummyComfyWanModelObject": "Dummy Comfy Wan Model Object",
"WanVideoLatentReScale": "WanVideo Latent ReScale",
"CreateScheduleFloatList": "Create Schedule Float List"
}
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo"
version = "1.2.9"
version = "1.3.0"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.15.0", "ftfy", "gguf >= 0.14.0", "pyloudnorm"]
-127
View File
@@ -1,127 +0,0 @@
import cv2
import numpy as np
import onnxruntime
def nms(boxes, scores, nms_thr):
"""Single class NMS implemented in Numpy."""
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
order = scores.argsort()[::-1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0.0, xx2 - xx1 + 1)
h = np.maximum(0.0, yy2 - yy1 + 1)
inter = w * h
ovr = inter / (areas[i] + areas[order[1:]] - inter)
inds = np.where(ovr <= nms_thr)[0]
order = order[inds + 1]
return keep
def multiclass_nms(boxes, scores, nms_thr, score_thr):
"""Multiclass NMS implemented in Numpy. Class-aware version."""
final_dets = []
num_classes = scores.shape[1]
for cls_ind in range(num_classes):
cls_scores = scores[:, cls_ind]
valid_score_mask = cls_scores > score_thr
if valid_score_mask.sum() == 0:
continue
else:
valid_scores = cls_scores[valid_score_mask]
valid_boxes = boxes[valid_score_mask]
keep = nms(valid_boxes, valid_scores, nms_thr)
if len(keep) > 0:
cls_inds = np.ones((len(keep), 1)) * cls_ind
dets = np.concatenate(
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
)
final_dets.append(dets)
if len(final_dets) == 0:
return None
return np.concatenate(final_dets, 0)
def demo_postprocess(outputs, img_size, p6=False):
grids = []
expanded_strides = []
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
hsizes = [img_size[0] // stride for stride in strides]
wsizes = [img_size[1] // stride for stride in strides]
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
grids.append(grid)
shape = grid.shape[:2]
expanded_strides.append(np.full((*shape, 1), stride))
grids = np.concatenate(grids, 1)
expanded_strides = np.concatenate(expanded_strides, 1)
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
return outputs
def preprocess(img, input_size, swap=(2, 0, 1)):
if len(img.shape) == 3:
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
else:
padded_img = np.ones(input_size, dtype=np.uint8) * 114
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
resized_img = cv2.resize(
img,
(int(img.shape[1] * r), int(img.shape[0] * r)),
interpolation=cv2.INTER_LINEAR,
).astype(np.uint8)
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
padded_img = padded_img.transpose(swap)
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
return padded_img, r
def inference_detector(session, oriImg):
input_shape = (640,640)
img, ratio = preprocess(oriImg, input_shape)
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
output = session.run(None, ort_inputs)
predictions = demo_postprocess(output[0], input_shape)[0]
boxes = predictions[:, :4]
scores = predictions[:, 4:5] * predictions[:, 5:]
boxes_xyxy = np.ones_like(boxes)
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
boxes_xyxy /= ratio
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
if dets is not None:
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
isscore = final_scores>0.3
iscat = final_cls_inds == 0
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
final_boxes = final_boxes[isbbox]
else:
final_boxes = np.array([])
return final_boxes
-360
View File
@@ -1,360 +0,0 @@
from typing import List, Tuple
import cv2
import numpy as np
import onnxruntime as ort
def preprocess(
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Do preprocessing for RTMPose model inference.
Args:
img (np.ndarray): Input image in shape.
input_size (tuple): Input image size in shape (w, h).
Returns:
tuple:
- resized_img (np.ndarray): Preprocessed image.
- center (np.ndarray): Center of image.
- scale (np.ndarray): Scale of image.
"""
# get shape of image
img_shape = img.shape[:2]
out_img, out_center, out_scale = [], [], []
if len(out_bbox) == 0:
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
for i in range(len(out_bbox)):
x0 = out_bbox[i][0]
y0 = out_bbox[i][1]
x1 = out_bbox[i][2]
y1 = out_bbox[i][3]
bbox = np.array([x0, y0, x1, y1])
# get center and scale
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
# do affine transformation
resized_img, scale = top_down_affine(input_size, scale, center, img)
# normalize image
mean = np.array([123.675, 116.28, 103.53])
std = np.array([58.395, 57.12, 57.375])
resized_img = (resized_img - mean) / std
out_img.append(resized_img)
out_center.append(center)
out_scale.append(scale)
return out_img, out_center, out_scale
def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray:
"""Inference RTMPose model.
Args:
sess (ort.InferenceSession): ONNXRuntime session.
img (np.ndarray): Input image in shape.
Returns:
outputs (np.ndarray): Output of RTMPose model.
"""
all_out = []
# build input
for i in range(len(img)):
input = [img[i].transpose(2, 0, 1)]
# build output
sess_input = {sess.get_inputs()[0].name: input}
sess_output = []
for out in sess.get_outputs():
sess_output.append(out.name)
# run model
outputs = sess.run(sess_output, sess_input)
all_out.append(outputs)
return all_out
def postprocess(outputs: List[np.ndarray],
model_input_size: Tuple[int, int],
center: Tuple[int, int],
scale: Tuple[int, int],
simcc_split_ratio: float = 2.0
) -> Tuple[np.ndarray, np.ndarray]:
"""Postprocess for RTMPose model output.
Args:
outputs (np.ndarray): Output of RTMPose model.
model_input_size (tuple): RTMPose model Input image size.
center (tuple): Center of bbox in shape (x, y).
scale (tuple): Scale of bbox in shape (w, h).
simcc_split_ratio (float): Split ratio of simcc.
Returns:
tuple:
- keypoints (np.ndarray): Rescaled keypoints.
- scores (np.ndarray): Model predict scores.
"""
all_key = []
all_score = []
for i in range(len(outputs)):
# use simcc to decode
simcc_x, simcc_y = outputs[i]
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
# rescale keypoints
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
all_key.append(keypoints[0])
all_score.append(scores[0])
return np.array(all_key), np.array(all_score)
def bbox_xyxy2cs(bbox: np.ndarray,
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
"""Transform the bbox format from (x,y,w,h) into (center, scale)
Args:
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
as (left, top, right, bottom)
padding (float): BBox padding factor that will be multilied to scale.
Default: 1.0
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
(n, 2)
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
(n, 2)
"""
# convert single bbox from (4, ) to (1, 4)
dim = bbox.ndim
if dim == 1:
bbox = bbox[None, :]
# get bbox center and scale
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
scale = np.hstack([x2 - x1, y2 - y1]) * padding
if dim == 1:
center = center[0]
scale = scale[0]
return center, scale
def _fix_aspect_ratio(bbox_scale: np.ndarray,
aspect_ratio: float) -> np.ndarray:
"""Extend the scale to match the given aspect ratio.
Args:
scale (np.ndarray): The image scale (w, h) in shape (2, )
aspect_ratio (float): The ratio of ``w/h``
Returns:
np.ndarray: The reshaped image scale in (2, )
"""
w, h = np.hsplit(bbox_scale, [1])
bbox_scale = np.where(w > h * aspect_ratio,
np.hstack([w, w / aspect_ratio]),
np.hstack([h * aspect_ratio, h]))
return bbox_scale
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
"""Rotate a point by an angle.
Args:
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
angle_rad (float): rotation angle in radian
Returns:
np.ndarray: Rotated point in shape (2, )
"""
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
rot_mat = np.array([[cs, -sn], [sn, cs]])
return rot_mat @ pt
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
"""To calculate the affine matrix, three pairs of points are required. This
function is used to get the 3rd point, given 2D points a & b.
The 3rd point is defined by rotating vector `a - b` by 90 degrees
anticlockwise, using b as the rotation center.
Args:
a (np.ndarray): The 1st point (x,y) in shape (2, )
b (np.ndarray): The 2nd point (x,y) in shape (2, )
Returns:
np.ndarray: The 3rd point.
"""
direction = a - b
c = b + np.r_[-direction[1], direction[0]]
return c
def get_warp_matrix(center: np.ndarray,
scale: np.ndarray,
rot: float,
output_size: Tuple[int, int],
shift: Tuple[float, float] = (0., 0.),
inv: bool = False) -> np.ndarray:
"""Calculate the affine transformation matrix that can warp the bbox area
in the input image to the output size.
Args:
center (np.ndarray[2, ]): Center of the bounding box (x, y).
scale (np.ndarray[2, ]): Scale of the bounding box
wrt [width, height].
rot (float): Rotation angle (degree).
output_size (np.ndarray[2, ] | list(2,)): Size of the
destination heatmaps.
shift (0-100%): Shift translation ratio wrt the width/height.
Default (0., 0.).
inv (bool): Option to inverse the affine transform direction.
(inv=False: src->dst or inv=True: dst->src)
Returns:
np.ndarray: A 2x3 transformation matrix
"""
shift = np.array(shift)
src_w = scale[0]
dst_w = output_size[0]
dst_h = output_size[1]
# compute transformation matrix
rot_rad = np.deg2rad(rot)
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
dst_dir = np.array([0., dst_w * -0.5])
# get four corners of the src rectangle in the original image
src = np.zeros((3, 2), dtype=np.float32)
src[0, :] = center + scale * shift
src[1, :] = center + src_dir + scale * shift
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
# get four corners of the dst rectangle in the input image
dst = np.zeros((3, 2), dtype=np.float32)
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
if inv:
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
else:
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
return warp_mat
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get the bbox image as the model input by affine transform.
Args:
input_size (dict): The input size of the model.
bbox_scale (dict): The bbox scale of the img.
bbox_center (dict): The bbox center of the img.
img (np.ndarray): The original image.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: img after affine transform.
- np.ndarray[float32]: bbox scale after affine transform.
"""
w, h = input_size
warp_size = (int(w), int(h))
# reshape bbox to fixed aspect ratio
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
# get the affine matrix
center = bbox_center
scale = bbox_scale
rot = 0
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
# do affine transform
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
return img, bbox_scale
def get_simcc_maximum(simcc_x: np.ndarray,
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get maximum response location and value from simcc representations.
Note:
instance number: N
num_keypoints: K
heatmap height: H
heatmap width: W
Args:
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
Returns:
tuple:
- locs (np.ndarray): locations of maximum heatmap responses in shape
(K, 2) or (N, K, 2)
- vals (np.ndarray): values of maximum heatmap responses in shape
(K,) or (N, K)
"""
N, K, Wx = simcc_x.shape
simcc_x = simcc_x.reshape(N * K, -1)
simcc_y = simcc_y.reshape(N * K, -1)
# get maximum value locations
x_locs = np.argmax(simcc_x, axis=1)
y_locs = np.argmax(simcc_y, axis=1)
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
max_val_x = np.amax(simcc_x, axis=1)
max_val_y = np.amax(simcc_y, axis=1)
# get maximum value across x and y axis
mask = max_val_x > max_val_y
max_val_x[mask] = max_val_y[mask]
vals = max_val_x
locs[vals <= 0.] = -1
# reshape
locs = locs.reshape(N, K, 2)
vals = vals.reshape(N, K)
return locs, vals
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
"""Modulate simcc distribution with Gaussian.
Args:
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
simcc_split_ratio (int): The split ratio of simcc.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
- np.ndarray[float32]: scores in shape (K,) or (n, K)
"""
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
keypoints /= simcc_split_ratio
return keypoints, scores
def inference_pose(session, out_bbox, oriImg):
h, w = session.get_inputs()[0].shape[2:]
model_input_size = (w, h)
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
outputs = inference(session, resized_img)
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
return keypoints, scores
+6 -10
View File
@@ -180,12 +180,11 @@ def draw_body_and_foot(canvas, candidate, subset, score, stick_width=4, draw_bod
# Append head elements based on the condition
limbSeq_and_colors += head_elements
for limb_info in limbSeq_and_colors[:17]:
for limb_info in limbSeq_and_colors[:19]:
limbSeq, color = limb_info
for n in range(len(subset)):
index = subset[n][np.array(limbSeq) - 1]
conf = score[n][np.array(limbSeq) - 1]
if conf[0] < 0.3 or conf[1] < 0.3:
if index[0] < 0.3 or index[1] < 0.3:
continue
Y = candidate[index.astype(int), 0] * float(W)
X = candidate[index.astype(int), 1] * float(H)
@@ -194,11 +193,11 @@ def draw_body_and_foot(canvas, candidate, subset, score, stick_width=4, draw_bod
length = np.sqrt((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2)
angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1]))
polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stick_width), int(angle), 0, 360, 1)
cv2.fillConvexPoly(canvas, polygon, alpha_blend_color(color, conf[0] * conf[1]))
cv2.fillConvexPoly(canvas, polygon, alpha_blend_color(color, index[0] * index[1]))
canvas = (canvas * 0.6).astype(np.uint8)
for limb_info in limbSeq_and_colors[:18]:
for limb_info in limbSeq_and_colors[:19]:
limbSeq, color = limb_info
for i in limbSeq:
for n in range(len(subset)):
@@ -209,12 +208,9 @@ def draw_body_and_foot(canvas, candidate, subset, score, stick_width=4, draw_bod
conf = score[n][i - 1]
if not np.isfinite(x) or not np.isfinite(y):
continue
x = int(np.clip(x * W, 0, W - 1
))
x = int(np.clip(x * W, 0, W - 1))
y = int(np.clip(y * H, 0, H - 1))
# x = int(x * W)
# y = int(y * H)
cv2.circle(canvas, (x, y), 4, alpha_blend_color(color, conf), thickness=-1)
cv2.circle(canvas, (x, y), body_keypoint_size, alpha_blend_color(color, conf), thickness=-1)
return canvas
-1
View File
@@ -1,7 +1,6 @@
import numpy as np
from .jit_det import inference_detector as inference_jit_yolox
from .jit_pose import inference_pose as inference_jit_pose
import os
class Wholebody:
+18 -16
View File
@@ -1,9 +1,14 @@
import torch
import torch.nn as nn
import os, copy, math
import numpy as np
from tqdm import tqdm
from ..utils import log
import comfy.model_management as mm
from comfy.utils import ProgressBar
from tqdm import tqdm
def update_transformer(transformer, state_dict):
@@ -37,17 +42,18 @@ def update_transformer(transformer, state_dict):
nn.SiLU(),
nn.Conv2d(concat_dim * 4, randomref_dim, 3, stride=2, padding=1),
)
unianimate_sd = {}
state_dict_new = {}
for key in list(state_dict.keys()):
if "dwpose_embedding" in key:
state_dict_new[key.split("dwpose_embedding.")[1]] = state_dict.pop(key)
transformer.dwpose_embedding.load_state_dict(state_dict_new, strict=True)
state_dict_new = {}
state_dict_new[key] = state_dict.pop(key)
unianimate_sd.update(state_dict_new)
for key in list(state_dict.keys()):
if "randomref_embedding_pose" in key:
state_dict_new[key.split("randomref_embedding_pose.")[1]] = state_dict.pop(key)
transformer.randomref_embedding_pose.load_state_dict(state_dict_new,strict=True)
return transformer
state_dict_new[key] = state_dict.pop(key)
unianimate_sd.update(state_dict_new)
del state_dict_new
return transformer, unianimate_sd
# Openpose
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
@@ -55,14 +61,6 @@ def update_transformer(transformer, state_dict):
# 3rd Edited by ControlNet
# 4th Edited by ControlNet (added face and correct hands)
import os
import torch
import numpy as np
import copy
import torch
import numpy as np
import math
from .dwpose.wholebody import Wholebody
def smoothing_factor(t_e, cutoff):
@@ -771,10 +769,14 @@ class WanVideoUniAnimateDWPoseDetector:
ref = reference_pose_image
ref_np = ref.cpu().numpy() * 255
prev_fuser_state = torch._C._jit_texpr_fuser_enabled()
torch._C._jit_set_texpr_fuser_enabled(False) # removes warmup delay, may want to enable later
poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width, score_threshold, stick_width=stick_width,
draw_body=draw_body, body_keypoint_size=body_keypoint_size, draw_feet=draw_feet,
draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, handle_not_detected=handle_not_detected, draw_head=draw_head)
poses = poses / 255.0
torch._C._jit_set_texpr_fuser_enabled(prev_fuser_state)
if reference_pose_image is not None:
reference_pose = reference_pose.unsqueeze(0) / 255.0
else:
+21 -7
View File
@@ -99,18 +99,21 @@ def apply_rope_comfy_chunked(xq, xk, freqs_cis, num_chunks=4):
return xq_out, xk_out
def rope_riflex(pos, dim, theta, L_test, k, temporal):
def rope_riflex(pos, dim, i, theta, L_test, k, ntk_factor=1.0):
assert dim % 2 == 0
if mm.is_device_mps(pos.device) or mm.is_intel_xpu() or mm.is_directml_enabled():
device = torch.device("cpu")
else:
device = pos.device
if ntk_factor != 1.0:
theta *= ntk_factor
scale = torch.linspace(0, (dim - 2) / dim, steps=dim//2, dtype=torch.float64, device=device)
omega = 1.0 / (theta**scale)
# RIFLEX modification - adjust last frequency component if L_test and k are provided
if temporal and k > 0 and L_test:
if i==0 and k > 0 and L_test:
omega[k-1] = 0.9 * 2 * torch.pi / L_test
out = torch.einsum("...n,d->...nd", pos.to(dtype=torch.float32, device=device), omega)
@@ -127,10 +130,18 @@ class EmbedND_RifleX(nn.Module):
self.num_frames = num_frames
self.k = k
def forward(self, ids):
def forward(self, ids, ntk_factor=[1.0,1.0,1.0]):
n_axes = ids.shape[-1]
emb = torch.cat(
[rope_riflex(ids[..., i], self.axes_dim[i], self.theta, self.num_frames, self.k, temporal=True if i == 0 else False) for i in range(n_axes)],
[rope_riflex(
ids[..., i],
self.axes_dim[i],
i, #f h w
self.theta,
self.num_frames,
self.k,
ntk_factor[i])
for i in range(n_axes)],
dim=-3,
)
return emb.unsqueeze(1)
@@ -1594,6 +1605,7 @@ class WanModel(torch.nn.Module):
fantasy_portrait_input=None,
phantom_ref=None,
reverse_time=False,
ntk_alphas = [1.0, 1.0, 1.0],
mtv_motion_tokens=None,
mtv_motion_rotary_emb=None,
mtv_freqs=None,
@@ -1755,7 +1767,8 @@ class WanModel(torch.nn.Module):
if (self.cached_freqs is not None and
self.cached_shape == current_shape and
self.cached_cond == has_cond and
self.cached_rope_k == self.rope_embedder.k
self.cached_rope_k == self.rope_embedder.k and
self.cached_ntk_alphas == ntk_alphas
):
freqs = self.cached_freqs
else:
@@ -1786,15 +1799,16 @@ class WanModel(torch.nn.Module):
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
# Generate RoPE frequencies for the combined positions
freqs = self.rope_embedder(combined_img_ids).movedim(1, 2)
freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2)
else:
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
freqs = self.rope_embedder(img_ids).movedim(1, 2)
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
self.cached_freqs = freqs
self.cached_shape = current_shape
self.cached_cond = has_cond
self.cached_rope_k = self.rope_embedder.k
self.cached_ntk_alphas = ntk_alphas
# Stand-In RoPE frequencies
if x_ip is not None: