Merge branch 'main' into dev
This commit is contained in:
+24
-78
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,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
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user