Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fdb23dec7d | ||
|
|
07d7d8ca8e | ||
|
|
01869d4bf5 | ||
|
|
bf1d77fe15 | ||
|
|
cced7fefbe | ||
|
|
36bb0c73ee | ||
|
|
b982b4ef0c | ||
|
|
3a7100bc39 | ||
|
|
55c672028b | ||
|
|
be41f67fae | ||
|
|
351829a2b6 | ||
|
|
b551ec9e31 | ||
|
|
19bcee67ed | ||
|
|
2fe4834178 | ||
|
|
486564060f | ||
|
|
3730ccf603 | ||
|
|
1eab022bb0 | ||
|
|
6fd4c6640c | ||
|
|
7a6efc1456 | ||
|
|
e855726f10 | ||
|
|
a896101ec8 | ||
|
|
fd818faa08 | ||
|
|
b132a82f7a | ||
|
|
4b709a7a04 | ||
|
|
220aac2771 | ||
|
|
027bed8c3d | ||
|
|
b2c520ca44 | ||
|
|
20942b8fd9 | ||
|
|
74f337e06c | ||
|
|
264212dddb | ||
|
|
f988d19fdb | ||
|
|
95255c7ffa | ||
|
|
c42bf94b07 | ||
|
|
dcae850b96 | ||
|
|
9f019d7dfb | ||
|
|
c5d3fb450c | ||
|
|
f28e7da442 | ||
|
|
ac7d8cab98 | ||
|
|
fc5322fae4 | ||
|
|
e75f814312 | ||
|
|
222fc70eb7 | ||
|
|
8509236da1 | ||
|
|
1c24ef50f8 | ||
|
|
5360eeb345 | ||
|
|
3e45021422 | ||
|
|
41683a0423 | ||
|
|
0ac7401271 | ||
|
|
fed3b22bc9 | ||
|
|
392d0305fc | ||
|
|
be411bdef6 | ||
|
|
2ee9e2f356 | ||
|
|
ef82826161 | ||
|
|
93f7af6dc8 | ||
|
|
ae6fe0853e | ||
|
|
c80fed01fe | ||
|
|
95097fefc2 | ||
|
|
c49fe98e55 | ||
|
|
ebceb165cc | ||
|
|
e6bd1b413a | ||
|
|
4dbba4d06d | ||
|
|
a9e21f164c | ||
|
|
a19d6bf12d | ||
|
|
3aae54f220 | ||
|
|
3611341339 | ||
|
|
0fa5383106 | ||
|
|
78e3e1857c | ||
|
|
164a6bbebd | ||
|
|
4a3ab6958a | ||
|
|
91413b33f5 | ||
|
|
60b3ec57dd | ||
|
|
f6e2d6dd47 | ||
|
|
68d65979d6 | ||
|
|
66a21e544a | ||
|
|
7273470faa | ||
|
|
c36e56875e | ||
|
|
f9287e3ecd | ||
|
|
2790c532cc | ||
|
|
369043c5d8 | ||
|
|
eb6dd96049 | ||
|
|
38a48c670a | ||
|
|
dd58511d4e | ||
|
|
4c97d27583 | ||
|
|
bcdc0c0661 | ||
|
|
83c25644ef | ||
|
|
113b6df04d | ||
|
|
d000fbc645 | ||
|
|
cc58964027 | ||
|
|
fa57681424 | ||
|
|
0a65354247 | ||
|
|
a27a4892b6 | ||
|
|
8b037bce2e | ||
|
|
e867e642d4 | ||
|
|
00971fc121 | ||
|
|
0e81adb843 | ||
|
|
7f9560eb33 | ||
|
|
2369cdbbe9 | ||
|
|
123c9ca312 | ||
|
|
2b62866945 | ||
|
|
b5415573c5 | ||
|
|
9e005fab90 | ||
|
|
f85ea5a48a | ||
|
|
ee1bb5c5c7 | ||
|
|
041c31fe2e | ||
|
|
4e7b8dd92c | ||
|
|
e2333d0f04 | ||
|
|
2dd4c6bf15 | ||
|
|
9d0a21228b | ||
|
|
001bc0e24a | ||
|
|
7189922f36 | ||
|
|
faf9c5927d | ||
|
|
6c5dc0ba2a | ||
|
|
995804166d | ||
|
|
4e2bf0f1fa | ||
|
|
b06c7d2d6d | ||
|
|
014e711972 | ||
|
|
c47b1ded69 | ||
|
|
7bc45daaf2 | ||
|
|
aebeeb9160 | ||
|
|
e60eb995ee | ||
|
|
b9f6c9aa50 | ||
|
|
c1fbc93521 | ||
|
|
c4ca252fea | ||
|
|
c4db00609a | ||
|
|
5a52d6b92f | ||
|
|
a652c55bf7 | ||
|
|
196d39695f | ||
|
|
e5be3e5263 | ||
|
|
a6071c7be5 | ||
|
|
0cba1edd4e | ||
|
|
a9cd073f29 | ||
|
|
1e9e2be622 | ||
|
|
99c3978da4 | ||
|
|
30bd7d46eb | ||
|
|
0c9d5b8dcc | ||
|
|
66d44ec8db | ||
|
|
394c7c13d2 | ||
|
|
c9931364b3 | ||
|
|
e54fa5d059 | ||
|
|
772642b4f1 | ||
|
|
472ed70757 | ||
|
|
44feb24290 | ||
|
|
fa7a967ee7 | ||
|
|
ec161373f4 | ||
|
|
0e3fd0b491 | ||
|
|
f872460285 | ||
|
|
b826642a83 | ||
|
|
aa9f474958 | ||
|
|
e3c2a1431b | ||
|
|
4e31081262 | ||
|
|
ff26836cab | ||
|
|
22037243ab | ||
|
|
e926f7a069 | ||
|
|
e01e34da1f | ||
|
|
47514f678d | ||
|
|
de3c9c895a | ||
|
|
4576ddb35e | ||
|
|
68392684b5 | ||
|
|
d3f33a9f09 | ||
|
|
d0ef3b5601 | ||
|
|
6a37c0b2d6 | ||
|
|
f653544f7b | ||
|
|
425035d810 |
+149
-2
@@ -2,6 +2,10 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch
|
||||
import math
|
||||
from einops import rearrange
|
||||
|
||||
from ..wanvideo.modules.model import WanRMSNorm, attention
|
||||
from ..multitalk.multitalk import RotaryPositionalEmbedding1D, normalize_and_scale
|
||||
|
||||
class FeedForwardSwiGLU(nn.Module):
|
||||
def __init__(
|
||||
@@ -22,7 +26,7 @@ class FeedForwardSwiGLU(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
@@ -62,4 +66,147 @@ class TimestepEmbedder(nn.Module):
|
||||
if t_freq.dtype != dtype:
|
||||
t_freq = t_freq.to(dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
return t_emb
|
||||
|
||||
|
||||
class SingleStreamAttention(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
encoder_hidden_states_dim: int,
|
||||
num_heads: int,
|
||||
qkv_bias: bool,
|
||||
qk_norm: bool,
|
||||
attn_drop: float = 0.0,
|
||||
proj_drop: float = 0.0,
|
||||
eps: float = 1e-6,
|
||||
class_range: int = 24,
|
||||
class_interval: int = 4,
|
||||
attention_mode: str = "sdpa",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0, "dim should be divisible by num_heads"
|
||||
self.dim = dim
|
||||
self.encoder_hidden_states_dim = encoder_hidden_states_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.scale = self.head_dim**-0.5
|
||||
|
||||
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.q_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
|
||||
self.k_norm = WanRMSNorm(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
self.attention_mode = attention_mode
|
||||
|
||||
# multitalk related params
|
||||
self.class_interval = class_interval
|
||||
self.class_range = class_range
|
||||
self.rope_h1 = (0, self.class_interval)
|
||||
self.rope_h2 = (self.class_range - self.class_interval, self.class_range)
|
||||
self.rope_bak = int(self.class_range // 2)
|
||||
self.rope_1d = RotaryPositionalEmbedding1D(self.head_dim)
|
||||
|
||||
def _process_cross_attn(self, x, cond, frames_num=None, x_ref_attn_map=None):
|
||||
|
||||
N_t = frames_num
|
||||
out_dtype = x.dtype
|
||||
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
|
||||
|
||||
# get q for hidden_state
|
||||
B, N, C = x.shape
|
||||
q = self.q_linear(x)
|
||||
q_shape = (B, N, self.num_heads, self.head_dim)
|
||||
q = q.view(q_shape).permute((0, 2, 1, 3)) # [B, H, N, D]
|
||||
q = self.q_norm(q.to(self.q_norm.weight.dtype)).to(q.dtype)
|
||||
|
||||
# multitalk with rope1d pe
|
||||
if x_ref_attn_map is not None:
|
||||
max_values = x_ref_attn_map.max(1).values[:, None, None]
|
||||
min_values = x_ref_attn_map.min(1).values[:, None, None]
|
||||
max_min_values = torch.cat([max_values, min_values], dim=2)
|
||||
human1_max_value, human1_min_value = max_min_values[0, :, 0].max(), max_min_values[0, :, 1].min()
|
||||
human2_max_value, human2_min_value = max_min_values[1, :, 0].max(), max_min_values[1, :, 1].min()
|
||||
|
||||
human1 = normalize_and_scale(x_ref_attn_map[0], (human1_min_value, human1_max_value), (self.rope_h1[0], self.rope_h1[1]))
|
||||
human2 = normalize_and_scale(x_ref_attn_map[1], (human2_min_value, human2_max_value), (self.rope_h2[0], self.rope_h2[1]))
|
||||
back = torch.full((x_ref_attn_map.size(1),), self.rope_bak, dtype=human1.dtype).to(human1.device)
|
||||
max_indices = x_ref_attn_map.argmax(dim=0)
|
||||
normalized_map = torch.stack([human1, human2, back], dim=1)
|
||||
normalized_pos = normalized_map[range(x_ref_attn_map.size(1)), max_indices]
|
||||
|
||||
q = rearrange(q, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
|
||||
q = self.rope_1d(q, normalized_pos)
|
||||
q = rearrange(q, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
|
||||
|
||||
# get kv from encoder_hidden_states
|
||||
_, N_a, _ = cond.shape
|
||||
encoder_kv = self.kv_linear(cond)
|
||||
encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim)
|
||||
encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4))
|
||||
|
||||
encoder_k, encoder_v = encoder_kv.unbind(0)
|
||||
encoder_k = self.k_norm(encoder_k.to(self.k_norm.weight.dtype)).to(encoder_k.dtype)
|
||||
|
||||
|
||||
# multitalk with rope1d pe
|
||||
if x_ref_attn_map is not None:
|
||||
per_frame = torch.zeros(N_a, dtype=encoder_k.dtype).to(encoder_k.device)
|
||||
per_frame[:per_frame.size(0)//2] = (self.rope_h1[0] + self.rope_h1[1]) / 2
|
||||
per_frame[per_frame.size(0)//2:] = (self.rope_h2[0] + self.rope_h2[1]) / 2
|
||||
encoder_pos = torch.concat([per_frame]*N_t, dim=0)
|
||||
encoder_k = rearrange(encoder_k, "(B N_t) H S C -> B H (N_t S) C", N_t=N_t)
|
||||
encoder_k = self.rope_1d(encoder_k, encoder_pos)
|
||||
encoder_k = rearrange(encoder_k, "B H (N_t S) C -> (B N_t) H S C", N_t=N_t)
|
||||
|
||||
# Input tensors must be in format ``[B, M, H, K]``, where B is the batch size, M \
|
||||
# the sequence length, H the number of heads, and K the embeding size per head
|
||||
|
||||
q = rearrange(q, "B H M K -> B M H K")
|
||||
encoder_k = rearrange(encoder_k, "B H M K -> B M H K")
|
||||
encoder_v = rearrange(encoder_v, "B H M K -> B M H K")
|
||||
x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode)
|
||||
x = rearrange(x, "B M H K -> B H M K")
|
||||
|
||||
# linear transform
|
||||
x_output_shape = (B, N, C)
|
||||
x = x.transpose(1, 2)
|
||||
x = x.reshape(x_output_shape)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
|
||||
# reshape x to origin shape
|
||||
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
|
||||
|
||||
return x.type(out_dtype)
|
||||
|
||||
def forward(self, x, cond, num_latent_frames=None, num_cond_latents=None, x_ref_attn_map=None, human_num=None):
|
||||
|
||||
B, N, C = x.shape
|
||||
if (num_cond_latents is None or num_cond_latents == 0):
|
||||
# text to video
|
||||
output = self._process_cross_attn(x, cond, num_latent_frames, x_ref_attn_map)
|
||||
return None, output
|
||||
elif num_cond_latents is not None and num_cond_latents > 0:
|
||||
# image to video or video continuation
|
||||
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
|
||||
x_noise = x[:, num_cond_latents_thw:]
|
||||
cond = rearrange(cond, "(B N_t) M C -> B N_t M C", B=B)
|
||||
cond = cond[:, num_cond_latents:]
|
||||
cond = rearrange(cond, "B N_t M C -> (B N_t) M C")
|
||||
frames_num = num_latent_frames - num_cond_latents
|
||||
if human_num is not None and human_num == 2:
|
||||
# multitalk mode
|
||||
output_noise = self._process_cross_attn(x_noise, cond, frames_num, x_ref_attn_map)
|
||||
else:
|
||||
# singletalk mode
|
||||
output_noise = self._process_cross_attn(x_noise, cond, frames_num)
|
||||
output_cond = torch.zeros((B, num_cond_latents_thw, C), dtype=output_noise.dtype, device=output_noise.device)
|
||||
return output_cond, output_noise
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
import torch
|
||||
from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
from comfy_api.latest import io
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
|
||||
class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="WanVideoLongCatAvatarExtendEmbeds",
|
||||
category="WanVideoWrapper",
|
||||
inputs=[
|
||||
io.Latent.Input("prev_latents", tooltip="Full previous latents to be used to continue generation, continuation frames are selected based on 'overlap' parameter"),
|
||||
io.Custom("MULTITALK_EMBEDS").Input("audio_embeds", tooltip="Full length audio embeddings"),
|
||||
io.Int.Input("num_frames", default=93, min=1, max=256, step=1, tooltip="Number of new frames to generate"),
|
||||
io.Int.Input("overlap", default=13, min=0, max=16, step=1, tooltip="Number of overlapping frames from previous latents for video continuation, set to 0 for T2V"),
|
||||
io.Int.Input("frames_processed", default=0, min=0, max=10000, step=1, tooltip="Number of frames already processed in the video, used to select audio features"),
|
||||
io.Combo.Input("if_not_enough_audio", ["pad_with_start", "mirror_from_end"], default="pad_with_start", tooltip="What to do if there are not enough frames in pose_images for the window"),
|
||||
io.Int.Input("ref_frame_index", default=10, min=0, max=1000, step=1, tooltip="Values between 0 - 24 ensures better consistency, while selecting other ranges (e.g., -10 or 30) helps reduce repeated actions"),
|
||||
io.Int.Input("ref_mask_frame_range", default=3, min=0, max=20, step=1, tooltip="Larger range can further help mitigate repeated actions, but excessively large values may introduce artifacts"),
|
||||
io.Latent.Input("ref_latent", optional=True, tooltip="Reference latent used for consistency, generally should be either the init image, or first latent from first generation"),
|
||||
io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"),
|
||||
],
|
||||
outputs=[
|
||||
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
|
||||
io.Latent.Output(display_name="samples_slice", tooltip="Sliced latent samples for the new frames"),
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed, ref_frame_index, ref_mask_frame_range, ref_latent=None, samples=None) -> io.NodeOutput:
|
||||
|
||||
new_audio_embed = audio_embeds.copy()
|
||||
|
||||
audio_features = torch.stack(new_audio_embed["audio_features"])
|
||||
num_audio_features = audio_features.shape[1]
|
||||
if audio_features.shape[1] < frames_processed + num_frames:
|
||||
deficit = frames_processed + num_frames - audio_features.shape[1]
|
||||
if if_not_enough_audio == "pad_with_start":
|
||||
pad = audio_features[:, :1].repeat(1, deficit, 1, 1)
|
||||
audio_features = torch.cat([audio_features, pad], dim=1)
|
||||
elif if_not_enough_audio == "mirror_from_end":
|
||||
to_add = audio_features[:, -deficit:, :].flip(dims=[1])
|
||||
audio_features = torch.cat([audio_features, to_add], dim=1)
|
||||
log.warning(f"Not enough audio features, padded with strategy '{if_not_enough_audio}' from {num_audio_features} to {audio_features.shape[1]} frames")
|
||||
|
||||
ref_target_masks = new_audio_embed.get("ref_target_masks", None)
|
||||
if ref_target_masks is not None:
|
||||
new_audio_embed["ref_target_masks"] = ref_target_masks[:, frames_processed:frames_processed+num_frames, :]
|
||||
|
||||
prev_samples = prev_latents["samples"].clone()
|
||||
if overlap != 0:
|
||||
latent_overlap = (overlap - 1) // 4 + 1
|
||||
prev_samples = prev_samples[:, :, -latent_overlap:]
|
||||
|
||||
ref_sample = None
|
||||
if ref_latent is not None:
|
||||
ref_sample = ref_latent["samples"][0, :, :1].clone()
|
||||
log.info(f"Previous latents shape: {prev_samples.shape}, using last {latent_overlap} latent frames for overlap.")
|
||||
|
||||
new_latent_frames = (num_frames - 1) // 4 + 1
|
||||
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
|
||||
|
||||
audio_stride = 2
|
||||
indices = torch.arange(2 * 2 + 1) - 2
|
||||
|
||||
if frames_processed == 0:
|
||||
audio_start_idx = 0
|
||||
else:
|
||||
audio_start_idx = (frames_processed - overlap) * audio_stride
|
||||
audio_end_idx = audio_start_idx + num_frames * audio_stride
|
||||
|
||||
log.info(f"Extracting audio embeddings from index {audio_start_idx} to {audio_end_idx}")
|
||||
|
||||
audio_embs = []
|
||||
for human_idx in range(len(audio_features)):
|
||||
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
|
||||
center_indices = torch.clamp(center_indices, min=0, max=audio_features[human_idx].shape[0] - 1)
|
||||
|
||||
audio_emb = audio_features[human_idx][center_indices].unsqueeze(0).to(device)
|
||||
audio_embs.append(audio_emb)
|
||||
audio_emb = torch.cat(audio_embs, dim=0)
|
||||
|
||||
new_audio_embed["audio_features"] = None
|
||||
new_audio_embed["audio_emb_slice"] = audio_emb
|
||||
|
||||
longcat_avatar_options = {
|
||||
"longcat_ref_latent": ref_sample,
|
||||
"ref_frame_index": ref_frame_index,
|
||||
"ref_mask_frame_range": ref_mask_frame_range,
|
||||
}
|
||||
|
||||
embeds = {
|
||||
"target_shape": target_shape,
|
||||
"num_frames": num_frames,
|
||||
"extra_latents": [{"samples": prev_samples, "index": 0}] if overlap != 0 else None,
|
||||
"multitalk_embeds": new_audio_embed,
|
||||
"longcat_avatar_options": longcat_avatar_options,
|
||||
}
|
||||
|
||||
samples_slice = None
|
||||
if samples is not None:
|
||||
latent_start_index = (frames_processed - 1) // 4 + 1 if frames_processed > 0 else 0
|
||||
latent_end_index = latent_start_index + new_latent_frames
|
||||
samples_slice = samples.copy()
|
||||
samples_slice["samples"] = samples["samples"][:, :, latent_start_index:latent_end_index].clone()
|
||||
|
||||
return io.NodeOutput(embeds, samples_slice)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from ..wanvideo.modules.attention import attention
|
||||
|
||||
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
|
||||
return (x * (1 + scale) + shift)
|
||||
|
||||
|
||||
def sinusoidal_embedding_1d(dim, position):
|
||||
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
|
||||
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
|
||||
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
return x.to(position.dtype)
|
||||
|
||||
|
||||
def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0):
|
||||
# 3d rope precompute
|
||||
f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta)
|
||||
h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
|
||||
w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
|
||||
return f_freqs_cis, h_freqs_cis, w_freqs_cis
|
||||
|
||||
|
||||
def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0):
|
||||
# 1d rope precompute
|
||||
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
|
||||
[: (dim // 2)].double() / dim))
|
||||
freqs = torch.outer(torch.arange(end, device=freqs.device), freqs)
|
||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
|
||||
return freqs_cis
|
||||
|
||||
|
||||
def rope_apply(x, freqs, num_heads):
|
||||
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
|
||||
x_out = torch.view_as_complex(x.to(torch.float64).reshape(
|
||||
x.shape[0], x.shape[1], x.shape[2], -1, 2))
|
||||
x_out = torch.view_as_real(x_out * freqs).flatten(2)
|
||||
return x_out.to(x.dtype)
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
def forward(self, x):
|
||||
dtype = x.dtype
|
||||
return self.norm(x.float()).to(dtype) * self.weight
|
||||
|
||||
|
||||
class AttentionModule(nn.Module):
|
||||
def __init__(self, num_heads, head_dim):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = head_dim
|
||||
|
||||
def forward(self, q, k, v):
|
||||
b, n, d = q.size(0), self.num_heads, self.head_dim
|
||||
x = attention(
|
||||
q.view(b, -1, n, d),
|
||||
k.view(b, -1, n, d),
|
||||
v.view(b, -1, n, d)
|
||||
)
|
||||
return x.flatten(2)
|
||||
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
self.q = nn.Linear(dim, dim)
|
||||
self.k = nn.Linear(dim, dim)
|
||||
self.v = nn.Linear(dim, dim)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.norm_q = RMSNorm(dim, eps=eps)
|
||||
self.norm_k = RMSNorm(dim, eps=eps)
|
||||
|
||||
self.attn = AttentionModule(self.num_heads, self.head_dim)
|
||||
|
||||
def forward(self, x, freqs):
|
||||
q = self.norm_q(self.q(x))
|
||||
k = self.norm_k(self.k(x))
|
||||
v = self.v(x)
|
||||
q = rope_apply(q, freqs, self.num_heads)
|
||||
k = rope_apply(k, freqs, self.num_heads)
|
||||
x = self.attn(q, k, v)
|
||||
return self.o(x)
|
||||
|
||||
|
||||
class CrossAttention(nn.Module):
|
||||
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
|
||||
self.q = nn.Linear(dim, dim)
|
||||
self.k = nn.Linear(dim, dim)
|
||||
self.v = nn.Linear(dim, dim)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.norm_q = RMSNorm(dim, eps=eps)
|
||||
self.norm_k = RMSNorm(dim, eps=eps)
|
||||
|
||||
|
||||
self.k_img = nn.Linear(dim, dim)
|
||||
self.v_img = nn.Linear(dim, dim)
|
||||
self.norm_k_img = RMSNorm(dim, eps=eps)
|
||||
|
||||
self.attn = AttentionModule(self.num_heads, self.head_dim)
|
||||
|
||||
def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None):
|
||||
ctx = y
|
||||
q = self.norm_q(self.q(x))
|
||||
k = self.norm_k(self.k(ctx))
|
||||
v = self.v(ctx)
|
||||
x = self.attn(q, k, v)
|
||||
if clip_fea is not None:
|
||||
k_img = self.norm_k_img(self.k_img(clip_fea))
|
||||
v_img = self.v_img(clip_fea)
|
||||
y = self.attn(q, k_img, v_img)
|
||||
x = x + y
|
||||
return self.o(x)
|
||||
|
||||
|
||||
class GateModule(nn.Module):
|
||||
def __init__(self,):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, x, gate, residual):
|
||||
return x + gate * residual
|
||||
|
||||
class DiTBlock(nn.Module):
|
||||
def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.ffn_dim = ffn_dim
|
||||
|
||||
self.self_attn = SelfAttention(dim, num_heads, eps)
|
||||
self.cross_attn = CrossAttention(dim, num_heads, eps)
|
||||
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
||||
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
||||
self.norm3 = nn.LayerNorm(dim, eps=eps)
|
||||
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(
|
||||
approximate='tanh'), nn.Linear(ffn_dim, dim))
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
self.gate = GateModule()
|
||||
|
||||
def forward(self, x, context, t_mod, freqs, clip_fea=None):
|
||||
has_seq = len(t_mod.shape) == 4
|
||||
chunk_dim = 2 if has_seq else 1
|
||||
# msa: multi-head self-attention mlp: multi-layer perceptron
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)
|
||||
if has_seq:
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
|
||||
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
|
||||
)
|
||||
input_x = modulate(self.norm1(x), shift_msa, scale_msa)
|
||||
x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
|
||||
x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea)
|
||||
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
|
||||
x = self.gate(x, gate_mlp, self.ffn(input_x))
|
||||
return x
|
||||
|
||||
|
||||
class WanModelDualControl(torch.nn.Module):
|
||||
def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12):
|
||||
super().__init__()
|
||||
self.control_layers = control_layers
|
||||
self.control_blocks_dense = nn.ModuleList([
|
||||
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
|
||||
for _ in range(self.control_layers)
|
||||
])
|
||||
|
||||
self.control_blocks_sparse = nn.ModuleList([
|
||||
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
|
||||
for _ in range(self.control_layers)
|
||||
])
|
||||
|
||||
self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2)
|
||||
self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2)
|
||||
|
||||
self.control_text_linear = torch.nn.Linear(dim, dim//2)
|
||||
self.control_t_mod = torch.nn.Linear(dim, dim//2)
|
||||
|
||||
self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)])
|
||||
head_dim = dim // num_heads
|
||||
self.freqs = precompute_freqs_cis_3d(head_dim)
|
||||
@@ -0,0 +1,88 @@
|
||||
import torch
|
||||
from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
class WanVideoAddDualControlEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||
"first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}),
|
||||
},
|
||||
"optional": {
|
||||
"dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}),
|
||||
"sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}),
|
||||
"prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None):
|
||||
updated = dict(embeds)
|
||||
updated.setdefault("dual_control", {})
|
||||
|
||||
if dense is None and sparse is None:
|
||||
raise ValueError("At least one of dense or sparse inputs must be provided.")
|
||||
|
||||
num_frames = dense.shape[0] if dense is not None else sparse.shape[0]
|
||||
height = dense.shape[1] if dense is not None else sparse.shape[1]
|
||||
width = dense.shape[2] if dense is not None else sparse.shape[2]
|
||||
msk = torch.ones(1, num_frames, height//8, width//8, device=device)
|
||||
msk[:, 1:] = 0
|
||||
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
|
||||
msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
|
||||
msk = msk.transpose(1, 2)
|
||||
|
||||
dense_input_latent = sparse_input_latent = None
|
||||
|
||||
vae.to(device)
|
||||
if dense is not None:
|
||||
dense_images = 1 - dense[..., :3] # Invert colors for depth to match the usual range in comfy
|
||||
dense_images = dense_images.permute(3, 0, 1, 2) * 2 - 1
|
||||
dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False)
|
||||
dense_first = (dense_images[:, :1]).to(device, vae.dtype)
|
||||
vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
|
||||
dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False)
|
||||
dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1)
|
||||
dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1)
|
||||
if sparse is not None:
|
||||
sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1
|
||||
sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False)
|
||||
sparse_first = (sparse_images[:, :1]).to(device, vae.dtype)
|
||||
vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
|
||||
sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False)
|
||||
sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1)
|
||||
sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1)
|
||||
|
||||
if prev_images is not None:
|
||||
prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1
|
||||
prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False)
|
||||
updated["dual_control"]["prev_latent"] = prev_video_latent[0]
|
||||
|
||||
vae.to(offload_device)
|
||||
updated["dual_control"]["dense_input_latent"] = dense_input_latent
|
||||
updated["dual_control"]["sparse_input_latent"] = sparse_input_latent
|
||||
updated["dual_control"]["strength"] = strength
|
||||
updated["dual_control"]["start_percent"] = start_percent
|
||||
updated["dual_control"]["end_percent"] = end_percent
|
||||
updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level
|
||||
return (updated,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds",
|
||||
}
|
||||
+83
-12
@@ -31,7 +31,7 @@ def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3
|
||||
|
||||
def get_pose_images(smpl_data, offset):
|
||||
pose_images = []
|
||||
for data in smpl_data:
|
||||
for data in smpl_data:
|
||||
if isinstance(data, np.ndarray):
|
||||
joints3d = data
|
||||
else:
|
||||
@@ -43,28 +43,33 @@ def get_pose_images(smpl_data, offset):
|
||||
return pose_images
|
||||
|
||||
|
||||
def get_control_conditions(poses, h, w):
|
||||
video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
|
||||
def get_control_conditions(poses, h, w, stick_width=1.0, point_radius=2, style="original"):
|
||||
control_images = []
|
||||
for idx, pose in enumerate(poses):
|
||||
canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
|
||||
try:
|
||||
joints3d = p3d_to_p2d(pose, h, w)
|
||||
canvas = draw_3d_points(
|
||||
canvas,
|
||||
joints3d[0],
|
||||
stickwidth=int(h / 350),
|
||||
)
|
||||
if style == "original":
|
||||
canvas = draw_3d_points(
|
||||
canvas,
|
||||
joints3d[0],
|
||||
stickwidth=int(h / 350 * stick_width),
|
||||
r=point_radius,
|
||||
)
|
||||
elif style == "scail":
|
||||
canvas = draw_3d_points_scail(
|
||||
canvas,
|
||||
joints3d[0],
|
||||
stickwidth=int(h / 350 * stick_width),
|
||||
r=point_radius,
|
||||
)
|
||||
resized_canvas = cv2.resize(canvas, (w, h))
|
||||
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
|
||||
control_images.append(resized_canvas)
|
||||
except Exception as e:
|
||||
print("wrong:", e)
|
||||
except Exception:
|
||||
control_images.append(Image.fromarray(canvas))
|
||||
control_pixel_values = np.array(control_images)
|
||||
control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
|
||||
print("control_pixel_values.shape", control_pixel_values.shape)
|
||||
#control_pixel_values = video_transforms(control_pixel_values)
|
||||
return control_pixel_values
|
||||
|
||||
|
||||
@@ -140,3 +145,69 @@ def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True):
|
||||
cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
|
||||
|
||||
return canvas
|
||||
|
||||
def draw_3d_points_scail(canvas, points, stickwidth=2, r=2, draw_line=True):
|
||||
|
||||
connetions = [
|
||||
[15,12],[12, 16],[16, 18],[18, 20],[20, 22], # 0-4: Left arm chain
|
||||
[12,17],[17,19],[19,21], # 5-7: Right arm chain
|
||||
[21,23], # 8: Right hand
|
||||
[12,1],[1,4],[4,7], # 9-11: Neck to left leg (hip, thigh, shin)
|
||||
[12,2],[2,5],[5,8], # 12-14: Neck to right leg (hip, thigh, shin)
|
||||
]
|
||||
|
||||
# Warm colors for right side, cool colors for left side
|
||||
connection_colors = [
|
||||
[180, 180, 180], # 0: [15,12] - L. clavicle (Bright Cyan)
|
||||
[0, 200, 255], # 1: [12,16] - L. shoulder (Bright Cyan)
|
||||
[0, 120, 255], # 2: [16,18] - L. upper arm (Bright Blue)
|
||||
[0, 60, 255], # 3: [18,20] - L. forearm (Deep Blue)
|
||||
[60, 0, 255], # 4: [20,22] - L. hand (Blue-Purple)
|
||||
[255, 0, 0], # 5: [12,17] - R. clavicle (Bright Red)
|
||||
[255, 100, 0], # 6: [17,19] - R. upper arm (Bright Orange)
|
||||
[255, 180, 0], # 7: [19,21] - R. forearm (Golden Orange)
|
||||
[255, 255, 0], # 8: [21,23] - R. hand (Bright Yellow)
|
||||
[30, 27, 160], # 9: [12,1] - Neck to L. hip (purple-blue)
|
||||
[73, 27, 177], # 10: [1,4] - L. thigh (purple)
|
||||
[145, 27, 194], # 11: [4,7] - L. shin (magenta)
|
||||
[200, 255, 100], # 12: [12,2] - Neck to R. hip (yellow)
|
||||
[54, 201, 52], # 13: [2,5] - R. thigh (green)
|
||||
[30, 176, 85], # 14: [5,8] - R. shin (green)
|
||||
]
|
||||
|
||||
# draw line
|
||||
if draw_line:
|
||||
# Collect all joints that are part of connections
|
||||
joints_in_use = set()
|
||||
for connection in connetions:
|
||||
joints_in_use.add(connection[0])
|
||||
joints_in_use.add(connection[1])
|
||||
|
||||
for i in range(len(connetions)):
|
||||
point1_idx, point2_idx = connetions[i][0:2]
|
||||
point1 = points[point1_idx]
|
||||
point2 = points[point2_idx]
|
||||
x1, y1 = int(point1[0]), int(point1[1])
|
||||
x2, y2 = int(point2[0]), int(point2[1])
|
||||
cv2.line(canvas, (x1, y1), (x2, y2), connection_colors[i], stickwidth)
|
||||
|
||||
# draw points for joints that have connections
|
||||
joints_in_use = set()
|
||||
for connection in connetions:
|
||||
joints_in_use.add(connection[0])
|
||||
joints_in_use.add(connection[1])
|
||||
|
||||
for joint_idx in joints_in_use:
|
||||
if joint_idx >= len(points):
|
||||
continue
|
||||
x, y = points[joint_idx][0:2]
|
||||
x, y = int(x), int(y)
|
||||
# Use the color from the first connection involving this joint
|
||||
joint_color = [180, 180, 180] # default grey
|
||||
for i, connection in enumerate(connetions):
|
||||
if connection[0] == joint_idx or connection[1] == joint_idx:
|
||||
joint_color = connection_colors[i]
|
||||
break
|
||||
cv2.circle(canvas, (x, y), r, joint_color, thickness=-1)
|
||||
|
||||
return canvas
|
||||
|
||||
+146
-48
@@ -1,34 +1,56 @@
|
||||
import os
|
||||
import torch
|
||||
import gc
|
||||
from ..utils import log, dict_to_device
|
||||
from ..utils import log
|
||||
import numpy as np
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import load_torch_file
|
||||
import folder_paths
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript")
|
||||
folder_paths.add_model_folder_path("nlf", os.path.join(folder_paths.models_dir, "nlf"))
|
||||
|
||||
from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
|
||||
from .mtv import prepare_motion_embeddings
|
||||
|
||||
def check_jit_script_function():
|
||||
if torch.jit.script.__name__ != "script":
|
||||
# Get more details about what modified it
|
||||
module = torch.jit.script.__module__
|
||||
qualname = getattr(torch.jit.script, '__qualname__', 'unknown')
|
||||
code_file = None
|
||||
try:
|
||||
code_file = torch.jit.script.__code__.co_filename
|
||||
code_line = torch.jit.script.__code__.co_firstlineno
|
||||
log.warning(f"torch.jit.script has been modified by another custom node.\n"
|
||||
f" Function name: {torch.jit.script.__name__}\n"
|
||||
f" Module: {module}\n"
|
||||
f" Qualified name: {qualname}\n"
|
||||
f" Defined in: {code_file}:{code_line}\n"
|
||||
f"This may cause issues with the NLF model.")
|
||||
except:
|
||||
log.warning("--------------------------------")
|
||||
log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, "
|
||||
f"this has been modified by another custom node. This may cause issues with the NLF model.")
|
||||
log.warning("--------------------------------")
|
||||
|
||||
model_list = [
|
||||
"https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript",
|
||||
"https://github.com/isarandi/nlf/releases/download/v0.2.2/nlf_l_multi_0.2.2.torchscript",
|
||||
]
|
||||
|
||||
class DownloadAndLoadNLFModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"url": (
|
||||
[
|
||||
"https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"
|
||||
],
|
||||
)
|
||||
"url": (model_list, {"default": "https://github.com/isarandi/nlf/releases/download/v0.3.2/nlf_l_multi_0.3.2.torchscript"}),
|
||||
},
|
||||
"optional": {
|
||||
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -37,8 +59,11 @@ class DownloadAndLoadNLFModel:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, url):
|
||||
|
||||
def loadmodel(self, url, warmup=True):
|
||||
if url not in model_list:
|
||||
raise ValueError(f"URL {url} is not in the list of allowed models.")
|
||||
check_jit_script_function()
|
||||
|
||||
if not os.path.exists(local_model_path):
|
||||
log.info(f"Downloading NLF model to: {local_model_path}")
|
||||
import requests
|
||||
@@ -52,6 +77,20 @@ class DownloadAndLoadNLFModel:
|
||||
|
||||
model = torch.jit.load(local_model_path).eval()
|
||||
|
||||
if warmup:
|
||||
log.info("Warming up NLF model...")
|
||||
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
|
||||
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
|
||||
try:
|
||||
for _ in range(2):
|
||||
_ = model.detect_smpl_batched(dummy_input)
|
||||
finally:
|
||||
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
|
||||
|
||||
log.info("NLF model warmed up")
|
||||
|
||||
model = model.to(offload_device)
|
||||
|
||||
return (model,)
|
||||
|
||||
class LoadNLFModel:
|
||||
@@ -59,8 +98,12 @@ class LoadNLFModel:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"path": ("STRING", {"default": local_model_path}),
|
||||
"nlf_model": (folder_paths.get_filename_list("nlf"), {"tooltip": "These models are loaded from the 'ComfyUI/models/nlf' -folder",}),
|
||||
|
||||
},
|
||||
"optional": {
|
||||
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NLFMODEL",)
|
||||
@@ -68,8 +111,22 @@ class LoadNLFModel:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, path):
|
||||
model = torch.jit.load(path).eval()
|
||||
def loadmodel(self, nlf_model, warmup=True):
|
||||
check_jit_script_function()
|
||||
model = torch.jit.load(folder_paths.get_full_path_or_raise("nlf", nlf_model)).eval()
|
||||
|
||||
if warmup:
|
||||
log.info("Warming up NLF model...")
|
||||
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
|
||||
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
|
||||
try:
|
||||
for _ in range(2):
|
||||
_ = model.detect_smpl_batched(dummy_input)
|
||||
finally:
|
||||
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
|
||||
log.info("NLF model warmed up")
|
||||
|
||||
model = model.to(offload_device)
|
||||
|
||||
return model,
|
||||
|
||||
@@ -108,7 +165,7 @@ class LoadVQVAE:
|
||||
frame_upsample_rate=[2.0, 2.0],
|
||||
joint_upsample_rate=[1.0, 1.0]
|
||||
)
|
||||
|
||||
|
||||
vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device)
|
||||
vqvae.load_state_dict(vae_sd, strict=True)
|
||||
|
||||
@@ -131,15 +188,6 @@ class MTVCrafterEncodePoses:
|
||||
|
||||
def encode(self, vqvae, poses):
|
||||
|
||||
# import pickle
|
||||
# with open(os.path.join(script_directory, "data", "sampled_data.pkl"), 'rb') as f:
|
||||
# data_list = pickle.load(f)
|
||||
# if not isinstance(data_list, list):
|
||||
# data_list = [data_list]
|
||||
# print(data_list)
|
||||
|
||||
# smpl_poses = data_list[1]['pose']
|
||||
|
||||
global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3)
|
||||
global_std = np.load(os.path.join(script_directory, "data", "std.npy"))
|
||||
|
||||
@@ -153,7 +201,7 @@ class MTVCrafterEncodePoses:
|
||||
|
||||
vqvae.to(device)
|
||||
motion_tokens, vq_loss = vqvae(norm_poses.to(device), return_vq=True)
|
||||
|
||||
|
||||
recon_motion = vqvae(norm_poses.to(device))[0][0].to(dtype=torch.float32).cpu().detach() * global_std + global_mean
|
||||
vqvae.to(offload_device)
|
||||
|
||||
@@ -162,7 +210,7 @@ class MTVCrafterEncodePoses:
|
||||
'global_mean': global_mean,
|
||||
'global_std': global_std
|
||||
}
|
||||
|
||||
|
||||
return poses_dict, recon_motion
|
||||
|
||||
|
||||
@@ -173,32 +221,74 @@ class NLFPredict:
|
||||
"model": ("NLFMODEL",),
|
||||
"images": ("IMAGE", {"tooltip": "Input images for the model"}),
|
||||
},
|
||||
"optional": {
|
||||
"per_batch": ("INT", {"default": -1, "min": -1, "max": 10000, "step": 1, "tooltip": "How many images to process at once. -1 means all at once."}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NLFPRED", )
|
||||
RETURN_NAMES = ("pose_results",)
|
||||
RETURN_TYPES = ("NLFPRED", "BBOX",)
|
||||
RETURN_NAMES = ("pose_results", "bboxes")
|
||||
FUNCTION = "predict"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def predict(self, model, images):
|
||||
|
||||
model.to(device)
|
||||
pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
|
||||
model.to(offload_device)
|
||||
def predict(self, model, images, per_batch=-1):
|
||||
|
||||
pred = dict_to_device(pred, offload_device)
|
||||
check_jit_script_function()
|
||||
model = model.to(device)
|
||||
|
||||
num_images = images.shape[0]
|
||||
|
||||
# Determine batch size
|
||||
if per_batch == -1:
|
||||
batch_size = num_images
|
||||
else:
|
||||
batch_size = per_batch
|
||||
|
||||
# Initialize result containers
|
||||
all_boxes = []
|
||||
all_joints3d_nonparam = []
|
||||
|
||||
# Process in batches
|
||||
for i in range(0, num_images, batch_size):
|
||||
end_idx = min(i + batch_size, num_images)
|
||||
batch_images = images[i:end_idx]
|
||||
|
||||
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
|
||||
try:
|
||||
pred = model.detect_smpl_batched(batch_images.permute(0, 3, 1, 2).to(device))
|
||||
finally:
|
||||
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
|
||||
|
||||
# Collect boxes and joints from this batch
|
||||
if 'boxes' in pred:
|
||||
all_boxes.extend(pred['boxes'])
|
||||
if 'joints3d_nonparam' in pred:
|
||||
all_joints3d_nonparam.extend(pred['joints3d_nonparam'])
|
||||
|
||||
model = model.to(offload_device)
|
||||
|
||||
# Move collected results to offload device
|
||||
all_boxes = [box.to(offload_device) for box in all_boxes]
|
||||
all_joints3d_nonparam = [joints.to(offload_device) for joints in all_joints3d_nonparam]
|
||||
|
||||
# Maintain the original nested format: wrap in a list to match expected structure
|
||||
pose_results = {
|
||||
'joints3d_nonparam': [],
|
||||
'joints3d_nonparam': [all_joints3d_nonparam],
|
||||
}
|
||||
# Collect pose data
|
||||
for key in pose_results.keys():
|
||||
if key in pred:
|
||||
pose_results[key].append(pred[key])
|
||||
|
||||
# Convert bboxes to list format: [x_min, y_min, x_max, y_max] for each detection
|
||||
# Each box tensor is shape (1, 5) with [x_min, y_min, x_max, y_max, confidence]
|
||||
formatted_boxes = []
|
||||
for box in all_boxes:
|
||||
# Handle empty detections (no person detected in frame)
|
||||
if box.numel() == 0 or box.shape[0] == 0:
|
||||
formatted_boxes.append([0.0, 0.0, 0.0, 0.0])
|
||||
else:
|
||||
pose_results[key].append(None)
|
||||
|
||||
return (pose_results,)
|
||||
# Extract first 4 values (x_min, y_min, x_max, y_max), drop confidence
|
||||
bbox_values = box[0, :4].cpu().tolist()
|
||||
formatted_boxes.append(bbox_values)
|
||||
|
||||
return (pose_results, formatted_boxes)
|
||||
|
||||
class DrawNLFPoses:
|
||||
@classmethod
|
||||
@@ -208,25 +298,32 @@ class DrawNLFPoses:
|
||||
"width": ("INT", {"default": 512}),
|
||||
"height": ("INT", {"default": 512}),
|
||||
},
|
||||
}
|
||||
"optional": {
|
||||
"stick_width": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 1000.0, "step": 0.01, "tooltip": "Stick width multiplier"}),
|
||||
"point_radius": ("INT", {"default": 5, "min": 1, "max": 10, "step": 1, "tooltip": "Point radius for drawing the pose"}),
|
||||
"style": (["original", "scail"], {"default": "original", "tooltip": "style of the pose drawing"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "predict"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def predict(self, poses, width, height):
|
||||
def predict(self, poses, width, height, stick_width=1.0, point_radius=2, style="original"):
|
||||
from .draw_pose import get_control_conditions
|
||||
print(type(poses))
|
||||
|
||||
if isinstance(poses, dict):
|
||||
pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
|
||||
else:
|
||||
pose_input = poses
|
||||
control_conditions = get_control_conditions(pose_input, height, width)
|
||||
|
||||
control_conditions = get_control_conditions(pose_input, height, width, stick_width=stick_width, point_radius=point_radius, style=style)
|
||||
|
||||
return (control_conditions,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LoadNLFModel": LoadNLFModel,
|
||||
"DownloadAndLoadNLFModel": DownloadAndLoadNLFModel,
|
||||
"NLFPredict": NLFPredict,
|
||||
"DrawNLFPoses": DrawNLFPoses,
|
||||
@@ -234,6 +331,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"MTVCrafterEncodePoses": MTVCrafterEncodePoses
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LoadNLFModel": "Load NLF Model",
|
||||
"DownloadAndLoadNLFModel": "(Download)Load NLF Model",
|
||||
"NLFPredict": "NLF Predict",
|
||||
"DrawNLFPoses": "Draw NLF Poses",
|
||||
|
||||
+13
-3
@@ -8,6 +8,8 @@ from .vae.autoencoder import AutoEncoderModule
|
||||
from .vae.distributions import DiagonalGaussianDistribution
|
||||
import torchaudio
|
||||
|
||||
from ..utils import log
|
||||
|
||||
from comfy import model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
@@ -216,9 +218,11 @@ class WanVideoOviCFG:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"original_text_embeds": ("WANVIDEOTEXTEMBEDS",),
|
||||
"ovi_negative_text_embeds": ("WANVIDEOTEXTEMBEDS",),
|
||||
"ovi_audio_cfg": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 100.0, "step": 0.01}),
|
||||
},
|
||||
"optional": {
|
||||
"ovi_negative_text_embeds": ("WANVIDEOTEXTEMBEDS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
|
||||
@@ -227,10 +231,16 @@ class WanVideoOviCFG:
|
||||
CATEGORY = "WanVideoWrapper/Ovi"
|
||||
DESCRIPTION = "Adds Ovi negative text embeddings and audio CFG scale to the text embeddings dictionary"
|
||||
|
||||
def process(self, original_text_embeds, ovi_negative_text_embeds, ovi_audio_cfg):
|
||||
negative_text_embeds = ovi_negative_text_embeds.get("negative_prompt_embeds", None)
|
||||
def process(self, original_text_embeds, ovi_audio_cfg, ovi_negative_text_embeds=None):
|
||||
negative_text_embeds = None
|
||||
if ovi_negative_text_embeds is not None:
|
||||
negative_text_embeds = ovi_negative_text_embeds.get("prompt_embeds", None)
|
||||
if negative_text_embeds is None:
|
||||
negative_text_embeds = original_text_embeds["prompt_embeds"]
|
||||
log.info("WanVideoOviCFG: Ovi negative text embeddings not provided, using original prompt embeddings as negative embeddings")
|
||||
else:
|
||||
log.info("WanVideoOviCFG: Using provided Ovi audio negative text embeddings")
|
||||
log.info("WanVideoOviCFG: negative text embedding shape: {}".format(negative_text_embeds[0].shape))
|
||||
|
||||
prompt_embeds_dict_copy = original_text_embeds.copy()
|
||||
prompt_embeds_dict_copy.update({
|
||||
|
||||
+13
-6
@@ -75,14 +75,21 @@ class VAE(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
if data_dim == 80:
|
||||
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_80D, dtype=torch.float32))
|
||||
self.data_std = nn.Buffer(torch.tensor(DATA_STD_80D, dtype=torch.float32))
|
||||
data_mean = torch.tensor(DATA_MEAN_80D, dtype=torch.float32)
|
||||
data_std = torch.tensor(DATA_STD_80D, dtype=torch.float32)
|
||||
elif data_dim == 128:
|
||||
self.data_mean = nn.Buffer(torch.tensor(DATA_MEAN_128D, dtype=torch.float32))
|
||||
self.data_std = nn.Buffer(torch.tensor(DATA_STD_128D, dtype=torch.float32))
|
||||
data_mean = torch.tensor(DATA_MEAN_128D, dtype=torch.float32)
|
||||
data_std = torch.tensor(DATA_STD_128D, dtype=torch.float32)
|
||||
else:
|
||||
raise ValueError(f"Unsupported data_dim={data_dim}, expected 80 or 128")
|
||||
|
||||
self.data_mean = self.data_mean.view(1, -1, 1)
|
||||
self.data_std = self.data_std.view(1, -1, 1)
|
||||
# match old shape: (1, channels, 1)
|
||||
data_mean = data_mean.view(1, -1, 1)
|
||||
data_std = data_std.view(1, -1, 1)
|
||||
|
||||
# register as buffers so they move with .to(device) / .cuda()
|
||||
self.register_buffer("data_mean", data_mean)
|
||||
self.register_buffer("data_std", data_std)
|
||||
|
||||
self.encoder = Encoder1D(
|
||||
dim=hidden_dim,
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
import torch
|
||||
from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
class WanVideoAddSCAILReferenceEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||
"ref_image": ("IMAGE",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||
},
|
||||
"optional": {
|
||||
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, clip_embeds=None):
|
||||
updated = dict(embeds)
|
||||
|
||||
vae.to(device)
|
||||
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
|
||||
ref_latent = vae.encode([ref_image_in], device, tiled=False)[0]
|
||||
log.info(f"SCAIL ref_latent shape: {ref_latent.shape}")
|
||||
|
||||
ref_mask = torch.ones_like(ref_latent[:4])
|
||||
ref_latent = torch.cat([ref_latent, ref_mask], dim=0)
|
||||
vae.to(offload_device)
|
||||
|
||||
updated.setdefault("scail_embeds", {})
|
||||
updated["scail_embeds"]["ref_latent_pos"] = ref_latent * strength
|
||||
updated["scail_embeds"]["ref_latent_neg"] = torch.zeros_like(ref_latent)
|
||||
updated["scail_embeds"]["ref_start_percent"] = start_percent
|
||||
updated["scail_embeds"]["ref_end_percent"] = end_percent
|
||||
updated["clip_context"] = clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None
|
||||
|
||||
return (updated,)
|
||||
|
||||
class WanVideoAddSCAILPoseEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, vae, pose_images, strength, start_percent=0.0, end_percent=1.0):
|
||||
updated = dict(embeds)
|
||||
|
||||
vae.to(device)
|
||||
pose_images_in = (pose_images[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
|
||||
pose_latent = vae.encode([pose_images_in], device, tiled=False)[0]
|
||||
pose_mask = torch.ones_like(pose_latent[:4])
|
||||
pose_latent = torch.cat([pose_latent, pose_mask], dim=0)
|
||||
log.info(f"SCAIL pose_latent shape: {pose_latent.shape}")
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
updated.setdefault("scail_embeds", {})
|
||||
updated["scail_embeds"]["pose_latent"] = pose_latent
|
||||
updated["scail_embeds"]["pose_strength"] = strength
|
||||
updated["scail_embeds"]["pose_start_percent"] = start_percent
|
||||
updated["scail_embeds"]["pose_end_percent"] = end_percent
|
||||
|
||||
return (updated,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddSCAILPoseEmbeds": WanVideoAddSCAILPoseEmbeds,
|
||||
"WanVideoAddSCAILReferenceEmbeds": WanVideoAddSCAILReferenceEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddSCAILReferenceEmbeds": "WanVideo Add SCAIL Reference Embeds",
|
||||
"WanVideoAddSCAILPoseEmbeds": "WanVideo Add SCAIL Pose Embeds",
|
||||
}
|
||||
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,207 @@
|
||||
import json
|
||||
import torch
|
||||
import torchvision.transforms.functional as TF
|
||||
from ..utils import log
|
||||
from .trajectory import create_pos_feature_map, draw_tracks_on_video, replace_feature
|
||||
import os
|
||||
from comfy import model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
VAE_STRIDE = (4, 8, 8) # t, h, w
|
||||
|
||||
class WanVideoWanDrawWanMoveTracks:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"images": ("IMAGE",),
|
||||
"tracks": ("TRACKS",),
|
||||
},
|
||||
"optional": {
|
||||
"line_resolution": ("INT", {"default": 24, "min": 4, "max": 64, "step": 1, "tooltip": "Number of points to use for each line segment"}),
|
||||
"circle_size": ("INT", {"default": 10, "min": 1, "max": 20, "step": 1, "tooltip": "Size of the circle to draw for each track point"}),
|
||||
"opacity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Opacity of the circle to draw for each track point"}),
|
||||
"line_width": ("INT", {"default": 14, "min": 1, "max": 50, "step": 1, "tooltip": "Width of the line to draw for each track"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("image",)
|
||||
FUNCTION = "execute"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def execute(self, images, tracks, line_resolution=24, circle_size=10, opacity=0.5, line_width=14):
|
||||
if tracks is None or "track_path" not in tracks:
|
||||
log.warning("WanVideoWanDrawWanMoveTracks: No tracks provided.")
|
||||
return (images.float().cpu(), )
|
||||
track = tracks["track_path"].unsqueeze(0)
|
||||
track_visibility = tracks["track_visibility"].unsqueeze(0)
|
||||
images_in = images * 255.0
|
||||
if images_in.shape[0] != track.shape[1]:
|
||||
repeat_count = track.shape[1] // images.shape[0]
|
||||
images_in = images_in.repeat(repeat_count, 1, 1, 1)
|
||||
track_video = draw_tracks_on_video(images_in, track, track_visibility, track_frame=line_resolution, circle_size=circle_size, opacity=opacity, line_width=line_width)
|
||||
track_video = torch.stack([TF.to_tensor(frame) for frame in track_video], dim=0).movedim(1, -1)
|
||||
|
||||
return (track_video.float().cpu(), )
|
||||
|
||||
|
||||
class WanVideoAddWanMoveTracks:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
|
||||
},
|
||||
"optional": {
|
||||
"track_mask": ("MASK",),
|
||||
"track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}),
|
||||
"tracks": ("TRACKS", {"tooltip": "Alternatively use Comfy Tracks dictionary"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "TRACKS")
|
||||
RETURN_NAMES = ("image_embeds", "tracks")
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, image_embeds, track_coords=None, tracks=None, strength=1.0, track_mask=None):
|
||||
updated = dict(image_embeds)
|
||||
|
||||
track_visibility = None
|
||||
|
||||
target_shape = image_embeds.get("target_shape")
|
||||
if target_shape is not None:
|
||||
height = target_shape[2] * VAE_STRIDE[1]
|
||||
width = target_shape[3] * VAE_STRIDE[2]
|
||||
else:
|
||||
height = image_embeds["lat_h"] * VAE_STRIDE[1]
|
||||
width = image_embeds["lat_w"] * VAE_STRIDE[2]
|
||||
num_frames = image_embeds["num_frames"]
|
||||
|
||||
if track_coords is not None:
|
||||
tracks_data = parse_json_tracks(track_coords)
|
||||
track_list = [
|
||||
[[track[frame]['x'], track[frame]['y']] for track in tracks_data]
|
||||
for frame in range(len(tracks_data[0]))
|
||||
]
|
||||
track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2)
|
||||
elif tracks is not None and "track_path" in tracks:
|
||||
track = tracks["track_path"]
|
||||
if track_mask is None:
|
||||
track_visibility = tracks.get("track_visibility", None)
|
||||
track = track[:num_frames]
|
||||
|
||||
num_tracks = track.shape[-2]
|
||||
if track_visibility is None:
|
||||
if track_mask is None:
|
||||
track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device)
|
||||
else:
|
||||
track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1)
|
||||
feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=num_tracks, device=device)
|
||||
|
||||
updated.setdefault("wanmove_embeds", {})
|
||||
updated["wanmove_embeds"]["track_pos"] = track_pos
|
||||
updated["wanmove_embeds"]["strength"] = strength
|
||||
|
||||
tracks_dict = {
|
||||
"track_path": track,
|
||||
"track_visibility": track_visibility,
|
||||
}
|
||||
|
||||
return (updated, tracks_dict,)
|
||||
|
||||
|
||||
def parse_json_tracks(tracks):
|
||||
tracks_data = []
|
||||
try:
|
||||
# If tracks is a string, try to parse it as JSON
|
||||
if isinstance(tracks, str):
|
||||
parsed = json.loads(tracks.replace("'", '"'))
|
||||
tracks_data.extend(parsed)
|
||||
else:
|
||||
# If tracks is a list of strings, parse each one
|
||||
for track_str in tracks:
|
||||
parsed = json.loads(track_str.replace("'", '"'))
|
||||
tracks_data.append(parsed)
|
||||
|
||||
# Check if we have a single track (dict with x,y) or a list of tracks
|
||||
if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]:
|
||||
# Single track detected, wrap it in a list
|
||||
tracks_data = [tracks_data]
|
||||
elif tracks_data and isinstance(tracks_data[0], list) and tracks_data[0] and isinstance(tracks_data[0][0], dict) and 'x' in tracks_data[0][0]:
|
||||
# Already a list of tracks, nothing to do
|
||||
pass
|
||||
else:
|
||||
# Unexpected format
|
||||
log.warning(f"Warning: Unexpected track format: {type(tracks_data[0])}")
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
log.warning(f"Error parsing tracks JSON: {e}")
|
||||
tracks_data = []
|
||||
|
||||
return tracks_data
|
||||
|
||||
import node_helpers
|
||||
|
||||
class WanMove_native:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"positive": ("CONDITIONING",),
|
||||
"track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}),
|
||||
},
|
||||
"optional": {
|
||||
"track_mask": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING", "TRACKS")
|
||||
RETURN_NAMES = ("positive", "tracks")
|
||||
FUNCTION = "patchcond"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DEPRECATED = True
|
||||
|
||||
def patchcond(self, positive, track_coords, track_mask=None):
|
||||
|
||||
concat_latent_image = positive[0][1]["concat_latent_image"]
|
||||
B, C, T, H, W = concat_latent_image.shape
|
||||
num_frames = (T-1) * 4 + 1
|
||||
width = W * 8
|
||||
height = H * 8
|
||||
|
||||
tracks_data = parse_json_tracks(track_coords)
|
||||
track_list = [
|
||||
[[track[frame]['x'], track[frame]['y']] for track in tracks_data]
|
||||
for frame in range(len(tracks_data[0]))
|
||||
]
|
||||
track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2)
|
||||
track = track[:num_frames]
|
||||
|
||||
num_tracks = track.shape[-2]
|
||||
if track_mask is None:
|
||||
track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device)
|
||||
else:
|
||||
track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1)
|
||||
|
||||
feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=num_tracks, device=device)
|
||||
wanmove_cond = replace_feature(concat_latent_image, track_pos.unsqueeze(0))
|
||||
positive = node_helpers.conditioning_set_values(positive, {"concat_latent_image": wanmove_cond})
|
||||
|
||||
tracks_dict = {
|
||||
"track_path": track,
|
||||
"track_visibility": track_visibility,
|
||||
}
|
||||
return (positive, tracks_dict)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddWanMoveTracks": WanVideoAddWanMoveTracks,
|
||||
"WanVideoWanDrawWanMoveTracks": WanVideoWanDrawWanMoveTracks,
|
||||
"WanMove_native": WanMove_native,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddWanMoveTracks": "WanVideo Add WanMove Tracks",
|
||||
"WanVideoWanDrawWanMoveTracks": "WanVideo Draw WanMove Tracks",
|
||||
"WanMove_native": "WanMove Native",
|
||||
}
|
||||
@@ -0,0 +1,340 @@
|
||||
# https://github.com/ali-vilab/Wan-Move/blob/main/wan/modules/trajectory.py
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
SKIP_ZERO = False
|
||||
|
||||
def get_pos_emb(
|
||||
pos_k: torch.Tensor,
|
||||
pos_emb_dim: int,
|
||||
theta_func: callable = lambda i, d: torch.pow(10000, torch.mul(2, torch.div(i.to(torch.float32), d))),
|
||||
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Generate batch position embeddings.
|
||||
|
||||
Args:
|
||||
pos_k (torch.Tensor): A 1D tensor containing positions for which to generate embeddings.
|
||||
pos_emb_dim (int): The dimension of position embeddings.
|
||||
theta_func (callable): Function to compute thetas based on position and embedding dimensions.
|
||||
device (torch.device): Device to store the position embeddings.
|
||||
dtype (torch.dtype): Desired data type for computations.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: The position embeddings with shape (batch_size, pos_emb_dim).
|
||||
"""
|
||||
assert pos_emb_dim % 2 == 0, "The dimension of position embeddings must be even."
|
||||
pos_k = pos_k.to(device, dtype)
|
||||
if SKIP_ZERO:
|
||||
pos_k = pos_k + 1
|
||||
batch_size = pos_k.size(0)
|
||||
|
||||
denominator = torch.arange(0, pos_emb_dim // 2, device=device, dtype=dtype)
|
||||
# Expand denominator to match the shape needed for broadcasting
|
||||
denominator_expanded = denominator.view(1, -1).expand(batch_size, -1)
|
||||
|
||||
thetas = theta_func(denominator_expanded, pos_emb_dim)
|
||||
|
||||
# Ensure pos_k is in the correct shape for broadcasting
|
||||
pos_k_expanded = pos_k.view(-1, 1).to(dtype)
|
||||
sin_thetas = torch.sin(torch.div(pos_k_expanded, thetas))
|
||||
cos_thetas = torch.cos(torch.div(pos_k_expanded, thetas))
|
||||
|
||||
# Concatenate sine and cosine embeddings along the last dimension
|
||||
pos_emb = torch.cat([sin_thetas, cos_thetas], dim=-1)
|
||||
|
||||
return pos_emb
|
||||
|
||||
def create_pos_feature_map(
|
||||
pred_tracks: torch.Tensor, # [T, N, 2]
|
||||
pred_visibility: torch.Tensor, # [T, N]
|
||||
downsample_ratios: list[int],
|
||||
height: int,
|
||||
width: int,
|
||||
pos_emb_dim: int,
|
||||
track_num: int = -1,
|
||||
t_down_strategy: str = "sample",
|
||||
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
|
||||
dtype: torch.dtype = torch.float32,
|
||||
):
|
||||
"""
|
||||
Create a feature map from the predicted tracks.
|
||||
|
||||
Args:
|
||||
- pred_tracks: torch.Tensor, the predicted tracks, [T, N, 2]
|
||||
- pred_visibility: torch.Tensor, the predicted visibility, [T, N]
|
||||
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
|
||||
- height: int, the height of the feature map
|
||||
- width: int, the width of the feature map
|
||||
- pos_emb_dim: int, the dimension of the position embeddings
|
||||
- track_num: int, the number of tracks to use
|
||||
- t_down_strategy: str, the strategy for downsampling time dimension
|
||||
- device: torch.device, the device
|
||||
- dtype: torch.dtype, the data type
|
||||
|
||||
Returns:
|
||||
- feature_map: torch.Tensor, the feature map, [T', H', W', pos_emb_dim]
|
||||
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
|
||||
"""
|
||||
|
||||
assert t_down_strategy in ["sample", "average"], "Invalid strategy for downsampling time dimension."
|
||||
|
||||
t, n, _ = pred_tracks.shape
|
||||
t_down, h_down, w_down = downsample_ratios
|
||||
feature_map = torch.zeros((t-1) // t_down + 1, height // h_down, width // w_down, pos_emb_dim, device=device, dtype=dtype)
|
||||
track_pos = - torch.ones(n, (t-1) // t_down + 1, 2, dtype=torch.long)
|
||||
|
||||
if track_num == -1:
|
||||
track_num = n
|
||||
|
||||
tracks_idx = torch.randperm(n)[:track_num]
|
||||
tracks = pred_tracks[:, tracks_idx]
|
||||
visibility = pred_visibility[:, tracks_idx]
|
||||
#tracks_embs = get_pos_emb(torch.randperm(n)[:track_num], pos_emb_dim, device=device, dtype=dtype)
|
||||
|
||||
for t_idx in range(0, t, t_down):
|
||||
if t_down_strategy == "sample" or t_idx == 0:
|
||||
cur_tracks = tracks[t_idx] # [N, 2]
|
||||
cur_visibility = visibility[t_idx] # [N]
|
||||
else:
|
||||
cur_tracks = tracks[t_idx:t_idx+t_down].mean(dim=0)
|
||||
cur_visibility = torch.any(visibility[t_idx:t_idx+t_down], dim=0)
|
||||
|
||||
for i in range(track_num):
|
||||
if not cur_visibility[i] or cur_tracks[i][0] < 0 or cur_tracks[i][1] < 0 or cur_tracks[i][0] >= width or cur_tracks[i][1] >= height:
|
||||
continue
|
||||
x, y = cur_tracks[i]
|
||||
x, y = int(x // w_down), int(y // h_down)
|
||||
#feature_map[t_idx // t_down, y, x] += tracks_embs[i]
|
||||
track_pos[i, t_idx // t_down, 0], track_pos[i, t_idx // t_down, 1] = y, x
|
||||
|
||||
return feature_map, track_pos
|
||||
|
||||
|
||||
def replace_feature(
|
||||
vae_feature: torch.Tensor, # [B, C', T', H', W']
|
||||
track_pos: torch.Tensor, # [B, N, T', 2]
|
||||
strength: float = 1.0,
|
||||
) -> torch.Tensor:
|
||||
b, _, t, h, w = vae_feature.shape
|
||||
assert b == track_pos.shape[0], "Batch size mismatch."
|
||||
n = track_pos.shape[1]
|
||||
|
||||
# Shuffle the trajectory order
|
||||
track_pos = track_pos[:, torch.randperm(n), :, :]
|
||||
|
||||
# Extract coordinates at time steps ≥ 1 and generate a valid mask
|
||||
current_pos = track_pos[:, :, 1:, :] # [B, N, T-1, 2]
|
||||
mask = (current_pos[..., 0] >= 0) & (current_pos[..., 1] >= 0) # [B, N, T-1]
|
||||
|
||||
# Get all valid indices
|
||||
valid_indices = mask.nonzero(as_tuple=False) # [num_valid, 3]
|
||||
num_valid = valid_indices.shape[0]
|
||||
|
||||
if num_valid == 0:
|
||||
return vae_feature
|
||||
|
||||
# Decompose valid indices into each dimension
|
||||
batch_idx = valid_indices[:, 0]
|
||||
track_idx = valid_indices[:, 1]
|
||||
t_rel = valid_indices[:, 2]
|
||||
t_target = t_rel + 1 # Convert to original time step indices
|
||||
|
||||
# Extract target position coordinates
|
||||
h_target = current_pos[batch_idx, track_idx, t_rel, 0].long() # Ensure integer indices
|
||||
w_target = current_pos[batch_idx, track_idx, t_rel, 1].long()
|
||||
|
||||
# Extract source position coordinates (t=0)
|
||||
h_source = track_pos[batch_idx, track_idx, 0, 0].long()
|
||||
w_source = track_pos[batch_idx, track_idx, 0, 1].long()
|
||||
|
||||
# Get source features and assign to target positions
|
||||
src_features = vae_feature[batch_idx, :, 0, h_source, w_source]
|
||||
dst_features = vae_feature[batch_idx, :, t_target, h_target, w_target]
|
||||
|
||||
vae_feature[batch_idx, :, t_target, h_target, w_target] = dst_features + (src_features - dst_features) * strength
|
||||
|
||||
return vae_feature
|
||||
|
||||
def get_video_track_video(
|
||||
model,
|
||||
video_tensor: torch.Tensor, # [T, C, H, W]
|
||||
downsample_ratios: list[int],
|
||||
pos_emb_dim: int,
|
||||
grid_size: int = 32,
|
||||
track_num: int = -1,
|
||||
t_down_strategy: str = "sample",
|
||||
device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Get the track video from the video tensor.
|
||||
|
||||
Args:
|
||||
- model: torch.nn.Module, the model for tracking, CoTracker
|
||||
- video_tensor: torch.Tensor, the video tensor, [T, C, H, W]
|
||||
- downsample_ratios: list[int], the ratios for downsampling time, height, and width
|
||||
- height: int, the height of the feature map
|
||||
- width: int, the width of the feature map
|
||||
- pos_emb_dim: int, the dimension of the position embeddings
|
||||
- grid_size: int, the size of the grid
|
||||
- track_num: int, the number of tracks to use
|
||||
- t_down_strategy: str, the strategy for downsampling time dimension
|
||||
- device: torch.device, the device
|
||||
- dtype: torch.dtype, the data type
|
||||
|
||||
Returns:
|
||||
- track_video: torch.Tensor, the track video, [pos_emb_dim, T', H', W']
|
||||
- track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width
|
||||
- pred_tracks: the predicted point trajectories
|
||||
- pred_visibility: visibility of the predicted point trajectories
|
||||
"""
|
||||
|
||||
t, c, height, width = video_tensor.shape
|
||||
with (
|
||||
torch.autocast(device_type=device.type, dtype=dtype),
|
||||
torch.no_grad(),
|
||||
):
|
||||
pred_tracks, pred_visibility = model(
|
||||
video_tensor.unsqueeze(0),
|
||||
grid_size=grid_size,
|
||||
backward_tracking=False,
|
||||
)
|
||||
|
||||
track_video, track_pos = create_pos_feature_map(
|
||||
pred_tracks[0], pred_visibility[0], downsample_ratios, height, width, pos_emb_dim, track_num, t_down_strategy, device, dtype
|
||||
)
|
||||
|
||||
return track_video.permute(3, 0, 1, 2), track_pos, pred_tracks, pred_visibility
|
||||
|
||||
# ---------------------------
|
||||
# Visualize functions
|
||||
# --------------------------
|
||||
|
||||
def add_weighted(rgb, track):
|
||||
rgb = np.array(rgb) # [H, W, C] "RGB"
|
||||
track = np.array(track) # [H, W, C] "RGBA"
|
||||
|
||||
# Compute weights from the alpha channel
|
||||
alpha = track[:, :, 3] / 255.0
|
||||
|
||||
# Expand alpha to 3 channels to match RGB
|
||||
alpha = np.stack([alpha] * 3, axis=-1)
|
||||
|
||||
# Blend the two images
|
||||
blend_img = track[:, :, :3] * alpha + rgb * (1 - alpha)
|
||||
|
||||
return Image.fromarray(blend_img.astype(np.uint8))
|
||||
|
||||
def draw_tracks_on_video(video, tracks, visibility=None, track_frame=24, circle_size=12, opacity=0.5, line_width=16):
|
||||
color_map = [(102, 153, 255), (0, 255, 255), (255, 255, 0), (255, 102, 204), (0, 255, 0)]
|
||||
|
||||
video = video.byte().cpu().numpy() # (81, 480, 832, 3)
|
||||
tracks = tracks[0].long().detach().cpu().numpy()
|
||||
if visibility is not None:
|
||||
visibility = visibility[0].detach().cpu().numpy()
|
||||
|
||||
num_frames, height, width = video.shape[:3]
|
||||
num_tracks = tracks.shape[1]
|
||||
alpha_opacity = int(255 * opacity)
|
||||
|
||||
output_frames = []
|
||||
for t in range(num_frames):
|
||||
frame_rgb = video[t].astype(np.float32)
|
||||
|
||||
# Create a single RGBA overlay for all tracks in this frame
|
||||
overlay = Image.new("RGBA", (width, height), (0, 0, 0, 0))
|
||||
draw_overlay = ImageDraw.Draw(overlay)
|
||||
|
||||
polyline_data = []
|
||||
|
||||
# Draw all circles on a single overlay
|
||||
for n in range(num_tracks):
|
||||
if visibility is not None and visibility[t, n] == 0:
|
||||
continue
|
||||
|
||||
track_coord = tracks[t, n]
|
||||
color = color_map[n % len(color_map)]
|
||||
circle_color = color + (alpha_opacity,)
|
||||
|
||||
draw_overlay.ellipse(
|
||||
(
|
||||
track_coord[0] - circle_size,
|
||||
track_coord[1] - circle_size,
|
||||
track_coord[0] + circle_size,
|
||||
track_coord[1] + circle_size
|
||||
),
|
||||
fill=circle_color
|
||||
)
|
||||
|
||||
# Store polyline data for batch processing
|
||||
tracks_coord = tracks[max(t - track_frame, 0):t + 1, n]
|
||||
if len(tracks_coord) > 1:
|
||||
polyline_data.append((tracks_coord, color))
|
||||
|
||||
# Blend circles overlay once
|
||||
overlay_np = np.array(overlay)
|
||||
alpha = overlay_np[:, :, 3:4] / 255.0
|
||||
frame_rgb = overlay_np[:, :, :3] * alpha + frame_rgb * (1 - alpha)
|
||||
|
||||
# Draw all polylines on a single overlay
|
||||
if polyline_data:
|
||||
polyline_overlay = Image.new("RGBA", (width, height), (0, 0, 0, 0))
|
||||
for tracks_coord, color in polyline_data:
|
||||
_draw_gradient_polyline_on_overlay(polyline_overlay, line_width, tracks_coord, color, opacity)
|
||||
|
||||
# Blend polylines overlay once
|
||||
polyline_np = np.array(polyline_overlay)
|
||||
alpha = polyline_np[:, :, 3:4] / 255.0
|
||||
frame_rgb = polyline_np[:, :, :3] * alpha + frame_rgb * (1 - alpha)
|
||||
|
||||
output_frames.append(Image.fromarray(frame_rgb.astype(np.uint8)))
|
||||
|
||||
return output_frames
|
||||
|
||||
|
||||
def _draw_gradient_polyline_on_overlay(overlay, line_width, points, start_color, opacity=1.0):
|
||||
"""
|
||||
Draw a gradient polyline directly onto an existing RGBA overlay image.
|
||||
This is an optimized version that doesn't create new images.
|
||||
"""
|
||||
draw = ImageDraw.Draw(overlay, 'RGBA')
|
||||
points = points[::-1]
|
||||
|
||||
# Compute total length
|
||||
total_length = 0
|
||||
segment_lengths = []
|
||||
for i in range(len(points) - 1):
|
||||
dx = points[i + 1][0] - points[i][0]
|
||||
dy = points[i + 1][1] - points[i][1]
|
||||
length = (dx * dx + dy * dy) ** 0.5
|
||||
segment_lengths.append(length)
|
||||
total_length += length
|
||||
|
||||
if total_length == 0:
|
||||
return
|
||||
|
||||
accumulated_length = 0
|
||||
|
||||
# Draw the gradient polyline
|
||||
for idx, (start_point, end_point) in enumerate(zip(points[:-1], points[1:])):
|
||||
segment_length = segment_lengths[idx]
|
||||
steps = max(int(segment_length), 1)
|
||||
|
||||
for i in range(steps):
|
||||
current_length = accumulated_length + (i / steps) * segment_length
|
||||
ratio = current_length / total_length
|
||||
|
||||
alpha = int(255 * (1 - ratio) * opacity)
|
||||
color = (*start_color, alpha)
|
||||
|
||||
x = int(start_point[0] + (end_point[0] - start_point[0]) * i / steps)
|
||||
y = int(start_point[1] + (end_point[1] - start_point[1]) * i / steps)
|
||||
|
||||
dynamic_line_width = max(int(line_width * (1 - ratio)), 1)
|
||||
draw.line([(x, y), (x + 1, y)], fill=color, width=dynamic_line_width)
|
||||
|
||||
accumulated_length += segment_length
|
||||
+60
-113
@@ -1,128 +1,75 @@
|
||||
try:
|
||||
from .utils import check_duplicate_nodes, log
|
||||
from .utils import check_duplicate_nodes, log, color_text
|
||||
duplicate_dirs = check_duplicate_nodes()
|
||||
if duplicate_dirs:
|
||||
warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n"
|
||||
for dir_path in duplicate_dirs:
|
||||
warning_msg += f" - {dir_path}\n"
|
||||
log.warning(warning_msg + "Please remove duplicates to avoid possible conflicts.")
|
||||
warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
|
||||
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
|
||||
except:
|
||||
pass
|
||||
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .nodes_sampler import NODE_CLASS_MAPPINGS as SAMPLER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SAMPLER_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .fun_camera.nodes import NODE_CLASS_MAPPINGS as FUN_CAMERA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .uni3c.nodes import NODE_CLASS_MAPPINGS as UNI3C_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNI3C_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .controlnet.nodes import NODE_CLASS_MAPPINGS as CONTROLNET_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .ATI.nodes import NODE_CLASS_MAPPINGS as ATI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ATI_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .multitalk.nodes import NODE_CLASS_MAPPINGS as MULTITALK_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MULTITALK_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .nodes_model_loading import NODE_CLASS_MAPPINGS as MODEL_LOADING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .nodes_utility import NODE_CLASS_MAPPINGS as UTILITY_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UTILITY_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODE_CACHE_DISPLAY_NAME_MAPPINGS
|
||||
from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as S2V_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .FlashVSR.flashvsr_nodes import NODE_CLASS_MAPPINGS as FLASHVSR_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .mocha.nodes import NODE_CLASS_MAPPINGS as MOCHA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MOCHA_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .utils import log
|
||||
|
||||
try:
|
||||
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: Qwen nodes not available due to error in importing them: {e}")
|
||||
QWEN_NODE_CLASS_MAPPINGS = {}
|
||||
QWEN_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
# Required modules (will raise on import failure)
|
||||
REQUIRED_MODULES = [
|
||||
(".nodes", "Main"),
|
||||
(".nodes_sampler", "Sampler"),
|
||||
(".nodes_model_loading", "ModelLoading"),
|
||||
(".nodes_utility", "Utility"),
|
||||
(".cache_methods.nodes_cache", "Cache"),
|
||||
]
|
||||
|
||||
try:
|
||||
from .fantasyportrait.nodes import NODE_CLASS_MAPPINGS as FANTASYPORTRAIT_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: FantasyPortrait nodes not available due to error in importing them: {e}")
|
||||
FANTASYPORTRAIT_NODE_CLASS_MAPPINGS = {}
|
||||
FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
# Optional modules (will warn on import failure)
|
||||
OPTIONAL_MODULES = [
|
||||
(".nodes_deprecated", "Deprecated"),
|
||||
(".s2v.nodes", "S2V"),
|
||||
(".FlashVSR.flashvsr_nodes", "FlashVSR"),
|
||||
(".mocha.nodes", "Mocha"),
|
||||
(".fun_camera.nodes", "FunCamera"),
|
||||
(".uni3c.nodes", "Uni3C"),
|
||||
(".controlnet.nodes", "ControlNet"),
|
||||
(".ATI.nodes", "ATI"),
|
||||
(".multitalk.nodes", "MultiTalk"),
|
||||
(".recammaster.nodes", "RecamMaster"),
|
||||
(".skyreels.nodes", "SkyReels"),
|
||||
(".fantasytalking.nodes", "FantasyTalking"),
|
||||
(".qwen.qwen", "Qwen"),
|
||||
(".fantasyportrait.nodes", "FantasyPortrait"),
|
||||
(".unianimate.nodes", "UniAnimate"),
|
||||
(".MTV.nodes", "MTV"),
|
||||
(".HuMo.nodes", "HuMo"),
|
||||
(".lynx.nodes", "Lynx"),
|
||||
(".Ovi.nodes_ovi", "Ovi"),
|
||||
(".steadydancer.nodes", "SteadyDancer"),
|
||||
(".onetoall.nodes", "OneToAll"),
|
||||
(".WanMove.nodes", "WanMove"),
|
||||
(".SCAIL.nodes", "SCAIL"),
|
||||
(".LongCat.nodes", "LongCat"),
|
||||
(".LongVie2.nodes", "LongVie2"),
|
||||
]
|
||||
|
||||
try:
|
||||
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: UniAnimate nodes not available due to error in importing them: {e}")
|
||||
UNIANIMATE_NODE_CLASS_MAPPINGS = {}
|
||||
UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
def register_nodes(module_path: str, name: str, optional: bool) -> None:
|
||||
"""Import and register nodes from a module."""
|
||||
try:
|
||||
import importlib
|
||||
module = importlib.import_module(module_path, package=__package__)
|
||||
NODE_CLASS_MAPPINGS.update(getattr(module, "NODE_CLASS_MAPPINGS", {}))
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(getattr(module, "NODE_DISPLAY_NAME_MAPPINGS", {}))
|
||||
except Exception as e:
|
||||
if optional:
|
||||
log.warning(f"WanVideoWrapper WARNING: {name} nodes not available: {e}")
|
||||
else:
|
||||
raise
|
||||
|
||||
try:
|
||||
from .MTV.nodes import NODE_CLASS_MAPPINGS as MTV_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MTV_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: MTV nodes not available due to error in importing them: {e}")
|
||||
MTV_NODE_CLASS_MAPPINGS = {}
|
||||
MTV_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
# Register all node modules
|
||||
for module_path, name in REQUIRED_MODULES:
|
||||
register_nodes(module_path, name, optional=False)
|
||||
|
||||
try:
|
||||
from .HuMo.nodes import NODE_CLASS_MAPPINGS as HUMO_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as HUMO_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: HuMo nodes not available due to error in importing them: {e}")
|
||||
HUMO_NODE_CLASS_MAPPINGS = {}
|
||||
HUMO_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
for module_path, name in OPTIONAL_MODULES:
|
||||
register_nodes(module_path, name, optional=True)
|
||||
|
||||
try:
|
||||
from .lynx.nodes import NODE_CLASS_MAPPINGS as LYNX_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as LYNX_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: Lynx nodes not available due to error in importing them: {e}")
|
||||
LYNX_NODE_CLASS_MAPPINGS = {}
|
||||
LYNX_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .Ovi.nodes_ovi import NODE_CLASS_MAPPINGS as OVI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as OVI_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: Ovi nodes not available due to error in importing them: {e}")
|
||||
OVI_NODE_CLASS_MAPPINGS = {}
|
||||
OVI_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FANTASYTALKING_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FANTASYPORTRAIT_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FUN_CAMERA_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(CONTROLNET_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(ATI_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MULTITALK_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MODEL_LOADING_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UTILITY_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(HUMO_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SAMPLER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYPORTRAIT_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNI3C_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(CONTROLNET_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(ATI_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MULTITALK_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MODEL_LOADING_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UTILITY_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(HUMO_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(SAMPLER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
@@ -28,19 +28,7 @@ aggressive values this can happen and the motion suffers. Starting later can hel
|
||||
When NOT using coefficients, the threshold value should be
|
||||
about 10 times smaller than the value used with coefficients.
|
||||
|
||||
Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1:
|
||||
|
||||
|
||||
<pre style='font-family:monospace'>
|
||||
+-------------------+--------+---------+--------+
|
||||
| Model | Low | Medium | High |
|
||||
+-------------------+--------+---------+--------+
|
||||
| Wan2.1 t2v 1.3B | 0.05 | 0.07 | 0.08 |
|
||||
| Wan2.1 t2v 14B | 0.14 | 0.15 | 0.20 |
|
||||
| Wan2.1 i2v 480P | 0.13 | 0.19 | 0.26 |
|
||||
| Wan2.1 i2v 720P | 0.18 | 0.20 | 0.30 |
|
||||
+-------------------+--------+---------+--------+
|
||||
</pre>
|
||||
Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaCache4Wan2.1
|
||||
"""
|
||||
|
||||
def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"):
|
||||
|
||||
@@ -10,18 +10,13 @@ from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscal
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.transformers.transformer_wan import (
|
||||
WanTimeTextImageEmbedding,
|
||||
WanRotaryPosEmbed,
|
||||
WanTimeTextImageEmbedding,
|
||||
WanRotaryPosEmbed,
|
||||
WanTransformerBlock
|
||||
)
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
def zero_module(module):
|
||||
for p in module.parameters():
|
||||
nn.init.zeros_(p)
|
||||
return module
|
||||
|
||||
|
||||
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
|
||||
r"""
|
||||
@@ -69,7 +64,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
||||
_no_split_modules = ["WanTransformerBlock"]
|
||||
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
|
||||
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
|
||||
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
@@ -100,10 +95,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
||||
## Spatial compression with time awareness
|
||||
nn.Sequential(
|
||||
nn.Conv3d(
|
||||
in_channels,
|
||||
input_channels[0],
|
||||
in_channels,
|
||||
input_channels[0],
|
||||
kernel_size=(3, downscale_coef + 1, downscale_coef + 1),
|
||||
stride=(1, downscale_coef, downscale_coef),
|
||||
stride=(1, downscale_coef, downscale_coef),
|
||||
padding=(1, downscale_coef // 2, downscale_coef // 2)
|
||||
),
|
||||
nn.GELU(approximate="tanh"),
|
||||
@@ -122,9 +117,9 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
||||
nn.GroupNorm(2, input_channels[2]),
|
||||
)
|
||||
])
|
||||
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
|
||||
self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size)
|
||||
@@ -153,11 +148,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
||||
|
||||
for _ in range(len(self.blocks)):
|
||||
controlnet_block = nn.Linear(inner_dim, out_proj_dim)
|
||||
controlnet_block = zero_module(controlnet_block)
|
||||
self.controlnet_blocks.append(controlnet_block)
|
||||
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -187,7 +181,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
||||
# 0. Controlnet encoder
|
||||
for control_encoder_block in self.control_encoder:
|
||||
controlnet_states = control_encoder_block(controlnet_states)
|
||||
|
||||
|
||||
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
|
||||
|
||||
## 1. Patch embedding and stack
|
||||
@@ -216,7 +210,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
|
||||
# 4. Transformer blocks
|
||||
controlnet_hidden_states = ()
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
@@ -239,43 +233,4 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
||||
return (controlnet_hidden_states,)
|
||||
|
||||
return Transformer2DModelOutput(sample=controlnet_hidden_states)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parameters = {
|
||||
"added_kv_proj_dim": None,
|
||||
"attention_head_dim": 128,
|
||||
"cross_attn_norm": True,
|
||||
"eps": 1e-06,
|
||||
"ffn_dim": 8960,
|
||||
"freq_dim": 256,
|
||||
"image_dim": None,
|
||||
"in_channels": 3,
|
||||
"num_attention_heads": 12,
|
||||
"num_layers": 2,
|
||||
"patch_size": [1, 2, 2],
|
||||
"qk_norm": "rms_norm_across_heads",
|
||||
"rope_max_seq_len": 1024,
|
||||
"text_dim": 4096,
|
||||
"downscale_coef": 8,
|
||||
"out_proj_dim": 12 * 128,
|
||||
"vae_channels": 16
|
||||
}
|
||||
controlnet = WanControlnet(**parameters)
|
||||
|
||||
hidden_states = torch.rand(1, 16, 13, 60, 90)
|
||||
timestep = torch.tensor([1000]).repeat(17550).unsqueeze(0) #torch.randint(low=0, high=1000, size=(1,), dtype=torch.long)
|
||||
encoder_hidden_states = torch.rand(1, 512, 4096)
|
||||
controlnet_states = torch.rand(1, 3, 49, 480, 720)
|
||||
|
||||
controlnet_hidden_states = controlnet(
|
||||
hidden_states=hidden_states,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
controlnet_states=controlnet_states,
|
||||
return_dict=False
|
||||
)
|
||||
print("Output states count", len(controlnet_hidden_states[0]))
|
||||
for out_hidden_states in controlnet_hidden_states[0]:
|
||||
print(out_hidden_states.shape)
|
||||
|
||||
|
||||
+153
-38
@@ -1,14 +1,52 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from accelerate import init_empty_weights
|
||||
from .gguf.gguf_utils import GGUFParameter, dequantize_gguf_tensor
|
||||
|
||||
@torch.library.custom_op("wanvideo::apply_lora", mutates_args=())
|
||||
def apply_lora(weight: torch.Tensor, lora_diff_0: torch.Tensor, lora_diff_1: torch.Tensor, lora_diff_2: float, lora_strength: torch.Tensor) -> torch.Tensor:
|
||||
patch_diff = torch.mm(
|
||||
lora_diff_0.flatten(start_dim=1),
|
||||
lora_diff_1.flatten(start_dim=1)
|
||||
).reshape(weight.shape)
|
||||
|
||||
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0
|
||||
scale = lora_strength * alpha
|
||||
|
||||
return weight + patch_diff * scale
|
||||
|
||||
@apply_lora.register_fake
|
||||
def _(weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
|
||||
# Return weight with same metadata
|
||||
return weight.clone()
|
||||
|
||||
@torch.library.custom_op("wanvideo::apply_single_lora", mutates_args=())
|
||||
def apply_single_lora(weight: torch.Tensor, lora_diff: torch.Tensor, lora_strength: torch.Tensor) -> torch.Tensor:
|
||||
return weight + lora_diff * lora_strength
|
||||
|
||||
@apply_single_lora.register_fake
|
||||
def _(weight, lora_diff, lora_strength):
|
||||
# Return weight with same metadata
|
||||
return weight.clone()
|
||||
|
||||
@torch.library.custom_op("wanvideo::linear_forward", mutates_args=())
|
||||
def linear_forward(input: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor:
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
@linear_forward.register_fake
|
||||
def _(input, weight, bias):
|
||||
# Calculate output shape: (..., out_features)
|
||||
out_features = weight.shape[0]
|
||||
output_shape = list(input.shape[:-1]) + [out_features]
|
||||
return input.new_empty(output_shape)
|
||||
|
||||
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
|
||||
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None):
|
||||
|
||||
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None, modules_to_not_convert=[]):
|
||||
|
||||
has_children = list(model.children())
|
||||
if not has_children:
|
||||
return
|
||||
|
||||
|
||||
allow_compile = False
|
||||
|
||||
for name, module in model.named_children():
|
||||
@@ -16,13 +54,22 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
||||
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
|
||||
module_prefix = prefix + name + "."
|
||||
module_prefix = module_prefix.replace("_orig_mod.", "")
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args)
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert)
|
||||
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
if scale_weights is not None:
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix and "dual_controller" not in module_prefix and name not in modules_to_not_convert:
|
||||
weight_key = module_prefix + "weight"
|
||||
if weight_key not in state_dict:
|
||||
continue
|
||||
|
||||
in_features = state_dict[weight_key].shape[1]
|
||||
out_features = state_dict[weight_key].shape[0]
|
||||
|
||||
is_gguf = isinstance(state_dict[weight_key], GGUFParameter)
|
||||
|
||||
scale_weight = None
|
||||
if not is_gguf and scale_weights is not None:
|
||||
scale_key = f"{module_prefix}scale_weight"
|
||||
scale_weight = scale_weights.get(scale_key)
|
||||
|
||||
with init_empty_weights():
|
||||
model._modules[name] = CustomLinear(
|
||||
@@ -30,8 +77,9 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
||||
out_features,
|
||||
module.bias is not None,
|
||||
compute_dtype=compute_dtype,
|
||||
scale_weight=scale_weights.get(scale_key) if scale_weights else None,
|
||||
allow_compile=allow_compile
|
||||
scale_weight=scale_weight,
|
||||
allow_compile=allow_compile,
|
||||
is_gguf=is_gguf
|
||||
)
|
||||
model._modules[name].source_cls = type(module)
|
||||
model._modules[name].requires_grad_(False)
|
||||
@@ -71,8 +119,8 @@ def set_lora_params(module, patches, module_prefix="", device=torch.device("cpu"
|
||||
continue
|
||||
lora_strengths = [p[0] for p in patch]
|
||||
module.set_lora_diffs(lora_diffs, device=device)
|
||||
module.lora_strengths = lora_strengths
|
||||
module.step = 0 # Initialize step for LoRA scheduling
|
||||
module.set_lora_strengths(lora_strengths, device=device)
|
||||
module._step.fill_(0) # Initialize step for LoRA scheduling
|
||||
|
||||
|
||||
class CustomLinear(nn.Linear):
|
||||
@@ -84,19 +132,56 @@ class CustomLinear(nn.Linear):
|
||||
compute_dtype=None,
|
||||
device=None,
|
||||
scale_weight=None,
|
||||
allow_compile=False
|
||||
allow_compile=False,
|
||||
is_gguf=False
|
||||
) -> None:
|
||||
super().__init__(in_features, out_features, bias, device)
|
||||
self.compute_dtype = compute_dtype
|
||||
self.lora_diffs = []
|
||||
self.step = 0
|
||||
self.register_buffer("_step", torch.zeros((), dtype=torch.long))
|
||||
self.scale_weight = scale_weight
|
||||
self.lora_strengths = []
|
||||
self.allow_compile = allow_compile
|
||||
self.is_gguf = is_gguf
|
||||
|
||||
if not allow_compile:
|
||||
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
|
||||
|
||||
self._apply_lora_impl = self._apply_lora_custom_op
|
||||
self._apply_single_lora_impl = self._apply_single_lora_custom_op
|
||||
self._linear_forward_impl = self._linear_forward_custom_op
|
||||
else:
|
||||
self._apply_lora_impl = self._apply_lora_direct
|
||||
self._apply_single_lora_impl = self._apply_single_lora_direct
|
||||
self._linear_forward_impl = self._linear_forward_direct
|
||||
|
||||
|
||||
# Direct implementations (no custom ops)
|
||||
def _apply_lora_direct(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
|
||||
patch_diff = torch.mm(
|
||||
lora_diff_0.flatten(start_dim=1),
|
||||
lora_diff_1.flatten(start_dim=1)
|
||||
).reshape(weight.shape) + 0
|
||||
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0
|
||||
scale = lora_strength * alpha
|
||||
return weight + patch_diff * scale
|
||||
|
||||
def _apply_single_lora_direct(self, weight, lora_diff, lora_strength):
|
||||
return weight + lora_diff * lora_strength
|
||||
|
||||
def _linear_forward_direct(self, input, weight, bias):
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
# Custom op implementations
|
||||
def _apply_lora_custom_op(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength):
|
||||
return torch.ops.wanvideo.apply_lora(weight, lora_diff_0, lora_diff_1,
|
||||
float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
|
||||
)
|
||||
|
||||
def _apply_single_lora_custom_op(self, weight, lora_diff, lora_strength):
|
||||
return torch.ops.wanvideo.apply_single_lora(weight, lora_diff, lora_strength)
|
||||
|
||||
def _linear_forward_custom_op(self, input, weight, bias):
|
||||
return torch.ops.wanvideo.linear_forward(input, weight, bias)
|
||||
|
||||
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
|
||||
self.lora_diffs = []
|
||||
for i, diff in enumerate(lora_diffs):
|
||||
@@ -109,51 +194,81 @@ class CustomLinear(nn.Linear):
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
|
||||
self.lora_diffs.append(f"lora_diff_{i}_0")
|
||||
|
||||
def set_lora_strengths(self, lora_strengths, device=torch.device("cpu")):
|
||||
self._lora_strength_tensors = []
|
||||
self._lora_strength_is_scheduled = []
|
||||
self._step = self._step.to(device)
|
||||
for i, strength in enumerate(lora_strengths):
|
||||
if isinstance(strength, list):
|
||||
tensor = torch.tensor(strength, dtype=self.compute_dtype, device=device)
|
||||
self.register_buffer(f"_lora_strength_{i}", tensor)
|
||||
self._lora_strength_is_scheduled.append(True)
|
||||
else:
|
||||
tensor = torch.tensor([strength], dtype=self.compute_dtype, device=device)
|
||||
self.register_buffer(f"_lora_strength_{i}", tensor)
|
||||
self._lora_strength_is_scheduled.append(False)
|
||||
|
||||
def _get_lora_strength(self, idx):
|
||||
strength_tensor = getattr(self, f"_lora_strength_{idx}")
|
||||
if self._lora_strength_is_scheduled[idx]:
|
||||
return strength_tensor.index_select(0, self._step).squeeze(0)
|
||||
return strength_tensor[0]
|
||||
|
||||
def _get_weight_with_lora(self, weight):
|
||||
"""Apply LoRA outside compiled region"""
|
||||
"""Apply LoRA using custom ops to avoid graph breaks"""
|
||||
if not hasattr(self, "lora_diff_0_0"):
|
||||
return weight
|
||||
|
||||
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
|
||||
if isinstance(lora_strength, list):
|
||||
lora_strength = lora_strength[self.step]
|
||||
if lora_strength == 0.0:
|
||||
continue
|
||||
elif lora_strength == 0.0:
|
||||
continue
|
||||
|
||||
for idx, lora_diff_names in enumerate(self.lora_diffs):
|
||||
lora_strength = self._get_lora_strength(idx)
|
||||
|
||||
if isinstance(lora_diff_names, tuple):
|
||||
lora_diff_0 = getattr(self, lora_diff_names[0])
|
||||
lora_diff_1 = getattr(self, lora_diff_names[1])
|
||||
lora_diff_2 = getattr(self, lora_diff_names[2])
|
||||
patch_diff = torch.mm(
|
||||
lora_diff_0.flatten(start_dim=1),
|
||||
lora_diff_1.flatten(start_dim=1)
|
||||
).reshape(weight.shape) + 0
|
||||
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0
|
||||
scale = lora_strength * alpha
|
||||
weight = weight.add(patch_diff, alpha=scale)
|
||||
|
||||
weight = self._apply_lora_impl(
|
||||
weight, lora_diff_0, lora_diff_1,
|
||||
float(lora_diff_2) if lora_diff_2 is not None else 0.0, lora_strength
|
||||
)
|
||||
else:
|
||||
lora_diff = getattr(self, lora_diff_names)
|
||||
weight = weight.add(lora_diff, alpha=lora_strength)
|
||||
weight = self._apply_single_lora_impl(weight, lora_diff, lora_strength)
|
||||
return weight
|
||||
|
||||
def _prepare_weight(self, input):
|
||||
"""Prepare weight tensor - handles both regular and GGUF weights"""
|
||||
if self.is_gguf:
|
||||
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
|
||||
else:
|
||||
weight = self.weight.to(input)
|
||||
return weight
|
||||
|
||||
def forward(self, input):
|
||||
weight = self._prepare_weight(input)
|
||||
|
||||
if self.bias is not None:
|
||||
bias = self.bias.to(input)
|
||||
bias = self.bias.to(input if not self.is_gguf else self.compute_dtype)
|
||||
else:
|
||||
bias = None
|
||||
weight = self.weight.to(input)
|
||||
|
||||
if self.scale_weight is not None:
|
||||
# Only apply scale_weight for non-GGUF models
|
||||
if not self.is_gguf and self.scale_weight is not None:
|
||||
if weight.numel() < input.numel():
|
||||
weight = weight * self.scale_weight
|
||||
else:
|
||||
input = input * self.scale_weight
|
||||
|
||||
weight = self._get_weight_with_lora(weight)
|
||||
out = self._linear_forward_impl(input, weight, bias)
|
||||
del weight, input, bias
|
||||
return out
|
||||
|
||||
def update_lora_step(module, step):
|
||||
for name, submodule in module.named_modules():
|
||||
if isinstance(submodule, CustomLinear) and hasattr(submodule, "_step"):
|
||||
submodule._step.fill_(step)
|
||||
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
def remove_lora_from_module(module):
|
||||
for name, submodule in module.named_modules():
|
||||
if hasattr(submodule, "lora_diffs"):
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
+1437
-1869
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1102
-984
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -1,780 +0,0 @@
|
||||
{
|
||||
"id": "8b7a9a57-2303-4ef5-9fc2-bf41713bd1fc",
|
||||
"revision": 0,
|
||||
"last_node_id": 46,
|
||||
"last_link_id": 58,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 33,
|
||||
"type": "Note",
|
||||
"pos": [
|
||||
227.3764190673828,
|
||||
-205.28524780273438
|
||||
],
|
||||
"size": [
|
||||
351.70458984375,
|
||||
88
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"Models:\nhttps://huggingface.co/Kijai/WanVideo_comfy/tree/main"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "LoadWanVideoT5TextEncoder",
|
||||
"pos": [
|
||||
224.15325927734375,
|
||||
-34.481563568115234
|
||||
],
|
||||
"size": [
|
||||
377.1661376953125,
|
||||
130
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "wan_t5_model",
|
||||
"type": "WANTEXTENCODER",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
15
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadWanVideoT5TextEncoder"
|
||||
},
|
||||
"widgets_values": [
|
||||
"umt5-xxl-enc-bf16.safetensors",
|
||||
"bf16",
|
||||
"offload_device",
|
||||
"disabled"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 28,
|
||||
"type": "WanVideoDecode",
|
||||
"pos": [
|
||||
1692.973876953125,
|
||||
-404.8614501953125
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
174
|
||||
],
|
||||
"flags": {},
|
||||
"order": 12,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "WANVAE",
|
||||
"link": 43
|
||||
},
|
||||
{
|
||||
"name": "samples",
|
||||
"type": "LATENT",
|
||||
"link": 33
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
48
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoDecode"
|
||||
},
|
||||
"widgets_values": [
|
||||
true,
|
||||
272,
|
||||
272,
|
||||
144,
|
||||
128
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 38,
|
||||
"type": "WanVideoVAELoader",
|
||||
"pos": [
|
||||
1687.4093017578125,
|
||||
-582.2750854492188
|
||||
],
|
||||
"size": [
|
||||
416.25482177734375,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "WANVAE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
43
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoVAELoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"wanvideo\\Wan2_1_VAE_bf16.safetensors",
|
||||
"bf16"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 42,
|
||||
"type": "GetImageSizeAndCount",
|
||||
"pos": [
|
||||
1708.7301025390625,
|
||||
-140.99705505371094
|
||||
],
|
||||
"size": [
|
||||
277.20001220703125,
|
||||
86
|
||||
],
|
||||
"flags": {},
|
||||
"order": 13,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 48
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
56
|
||||
]
|
||||
},
|
||||
{
|
||||
"label": "832 width",
|
||||
"name": "width",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"label": "480 height",
|
||||
"name": "height",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"label": "257 count",
|
||||
"name": "count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "GetImageSizeAndCount"
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 16,
|
||||
"type": "WanVideoTextEncode",
|
||||
"pos": [
|
||||
675.8850708007812,
|
||||
-36.032100677490234
|
||||
],
|
||||
"size": [
|
||||
420.30511474609375,
|
||||
261.5306701660156
|
||||
],
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "t5",
|
||||
"type": "WANTEXTENCODER",
|
||||
"link": 15
|
||||
},
|
||||
{
|
||||
"name": "model_to_offload",
|
||||
"shape": 7,
|
||||
"type": "WANVIDEOMODEL",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "text_embeds",
|
||||
"type": "WANVIDEOTEXTEMBEDS",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
30
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"high quality nature video featuring a red panda balancing on a bamboo stem while a bird lands on it's head, on the background there is a waterfall",
|
||||
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 30,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
2127.120849609375,
|
||||
-511.9014587402344
|
||||
],
|
||||
"size": [
|
||||
873.2135620117188,
|
||||
840.2385864257812
|
||||
],
|
||||
"flags": {},
|
||||
"order": 14,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 56
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 16,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "WanVideo2_1_T2V",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"trim_to_audio": false,
|
||||
"pingpong": false,
|
||||
"save_output": true,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "WanVideo2_1_T2V_00412.mp4",
|
||||
"subfolder": "",
|
||||
"type": "output",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 16,
|
||||
"workflow": "WanVideo2_1_T2V_00412.png",
|
||||
"fullpath": "N:\\AI\\ComfyUI\\output\\WanVideo2_1_T2V_00412.mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 37,
|
||||
"type": "WanVideoEmptyEmbeds",
|
||||
"pos": [
|
||||
1305.26708984375,
|
||||
-571.7843627929688
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "image_embeds",
|
||||
"type": "WANVIDIMAGE_EMBEDS",
|
||||
"links": [
|
||||
42
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoEmptyEmbeds"
|
||||
},
|
||||
"widgets_values": [
|
||||
832,
|
||||
480,
|
||||
257
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 35,
|
||||
"type": "WanVideoTorchCompileSettings",
|
||||
"pos": [
|
||||
193.47103881835938,
|
||||
-614.6900024414062
|
||||
],
|
||||
"size": [
|
||||
390.5999755859375,
|
||||
178
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "torch_compile_args",
|
||||
"type": "WANCOMPILEARGS",
|
||||
"slot_index": 0,
|
||||
"links": []
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoTorchCompileSettings"
|
||||
},
|
||||
"widgets_values": [
|
||||
"inductor",
|
||||
false,
|
||||
"default",
|
||||
false,
|
||||
64,
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 45,
|
||||
"type": "WanVideoTeaCache",
|
||||
"pos": [
|
||||
931.4036865234375,
|
||||
-792.5159912109375
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
154
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cache_args",
|
||||
"type": "CACHEARGS",
|
||||
"links": [
|
||||
58
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoTeaCache"
|
||||
},
|
||||
"widgets_values": [
|
||||
0.10000000000000002,
|
||||
1,
|
||||
-1,
|
||||
"offload_device",
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 36,
|
||||
"type": "Note",
|
||||
"pos": [
|
||||
796.0189208984375,
|
||||
-521.5020751953125
|
||||
],
|
||||
"size": [
|
||||
298.2554016113281,
|
||||
108.62744140625
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"sdpa should work too, haven't tested flaash\n\nfp8_fast seems to cause huge quality degradation"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 46,
|
||||
"type": "Note",
|
||||
"pos": [
|
||||
937.9556274414062,
|
||||
-940.750244140625
|
||||
],
|
||||
"size": [
|
||||
297.4364013671875,
|
||||
88
|
||||
],
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"TeaCache with context windows is VERY experimental and lower values than normal should be used."
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 27,
|
||||
"type": "WanVideoSampler",
|
||||
"pos": [
|
||||
1315.2401123046875,
|
||||
-401.48028564453125
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
574.1923217773438
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "WANVIDEOMODEL",
|
||||
"link": 29
|
||||
},
|
||||
{
|
||||
"name": "text_embeds",
|
||||
"type": "WANVIDEOTEXTEMBEDS",
|
||||
"link": 30
|
||||
},
|
||||
{
|
||||
"name": "image_embeds",
|
||||
"type": "WANVIDIMAGE_EMBEDS",
|
||||
"link": 42
|
||||
},
|
||||
{
|
||||
"name": "samples",
|
||||
"shape": 7,
|
||||
"type": "LATENT",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "feta_args",
|
||||
"shape": 7,
|
||||
"type": "FETAARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "context_options",
|
||||
"shape": 7,
|
||||
"type": "WANVIDCONTEXT",
|
||||
"link": 57
|
||||
},
|
||||
{
|
||||
"name": "cache_args",
|
||||
"shape": 7,
|
||||
"type": "CACHEARGS",
|
||||
"link": 58
|
||||
},
|
||||
{
|
||||
"name": "flowedit_args",
|
||||
"shape": 7,
|
||||
"type": "FLOWEDITARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "slg_args",
|
||||
"shape": 7,
|
||||
"type": "SLGARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "loop_args",
|
||||
"shape": 7,
|
||||
"type": "LOOPARGS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "samples",
|
||||
"type": "LATENT",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
33
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoSampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
30,
|
||||
6,
|
||||
5,
|
||||
1057359483639288,
|
||||
"fixed",
|
||||
true,
|
||||
"unipc",
|
||||
0,
|
||||
1,
|
||||
"",
|
||||
"comfy"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 43,
|
||||
"type": "WanVideoContextOptions",
|
||||
"pos": [
|
||||
1307.9542236328125,
|
||||
-855.8865356445312
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
226
|
||||
],
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "WANVAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "context_options",
|
||||
"type": "WANVIDCONTEXT",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
57
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoContextOptions"
|
||||
},
|
||||
"widgets_values": [
|
||||
"uniform_standard",
|
||||
81,
|
||||
4,
|
||||
16,
|
||||
true,
|
||||
false,
|
||||
6,
|
||||
2
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 22,
|
||||
"type": "WanVideoModelLoader",
|
||||
"pos": [
|
||||
620.3950805664062,
|
||||
-357.8426818847656
|
||||
],
|
||||
"size": [
|
||||
477.4410095214844,
|
||||
226.43276977539062
|
||||
],
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "compile_args",
|
||||
"shape": 7,
|
||||
"type": "WANCOMPILEARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "block_swap_args",
|
||||
"shape": 7,
|
||||
"type": "BLOCKSWAPARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "lora",
|
||||
"shape": 7,
|
||||
"type": "WANVIDLORA",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vram_management_args",
|
||||
"shape": 7,
|
||||
"type": "VRAM_MANAGEMENTARGS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "WANVIDEOMODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
29
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoModelLoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"WanVideo\\wan2.1_t2v_1.3B_fp16.safetensors",
|
||||
"fp16",
|
||||
"disabled",
|
||||
"offload_device",
|
||||
"sdpa"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
15,
|
||||
11,
|
||||
0,
|
||||
16,
|
||||
0,
|
||||
"WANTEXTENCODER"
|
||||
],
|
||||
[
|
||||
29,
|
||||
22,
|
||||
0,
|
||||
27,
|
||||
0,
|
||||
"WANVIDEOMODEL"
|
||||
],
|
||||
[
|
||||
30,
|
||||
16,
|
||||
0,
|
||||
27,
|
||||
1,
|
||||
"WANVIDEOTEXTEMBEDS"
|
||||
],
|
||||
[
|
||||
33,
|
||||
27,
|
||||
0,
|
||||
28,
|
||||
1,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
42,
|
||||
37,
|
||||
0,
|
||||
27,
|
||||
2,
|
||||
"WANVIDIMAGE_EMBEDS"
|
||||
],
|
||||
[
|
||||
43,
|
||||
38,
|
||||
0,
|
||||
28,
|
||||
0,
|
||||
"VAE"
|
||||
],
|
||||
[
|
||||
48,
|
||||
28,
|
||||
0,
|
||||
42,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
56,
|
||||
42,
|
||||
0,
|
||||
30,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
57,
|
||||
43,
|
||||
0,
|
||||
27,
|
||||
5,
|
||||
"WANVIDCONTEXT"
|
||||
],
|
||||
[
|
||||
58,
|
||||
45,
|
||||
0,
|
||||
27,
|
||||
6,
|
||||
"TEACACHEARGS"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8140274938684471,
|
||||
"offset": [
|
||||
-122.25834160503663,
|
||||
993.5739491626379
|
||||
]
|
||||
},
|
||||
"node_versions": {
|
||||
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
|
||||
"ComfyUI-KJNodes": "a5bd3c86c8ed6b83c55c2d0e7a59515b15a0137f",
|
||||
"ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8"
|
||||
},
|
||||
"VHS_latentpreview": true,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -220,6 +220,7 @@ class WanVideoAddFantasyPortrait:
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the portrait embedding"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||
"portrait_cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 20.0, "step": 0.01, "tooltip": "CFG scale for the portrait embedding"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -228,12 +229,13 @@ class WanVideoAddFantasyPortrait:
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, portrait_embeds, strength, start_percent=0.0, end_percent=1.0):
|
||||
def add(self, embeds, portrait_embeds, strength, start_percent=0.0, end_percent=1.0, portrait_cfg=1.0):
|
||||
new_entry = {
|
||||
"adapter_proj": portrait_embeds,
|
||||
"strength": strength,
|
||||
"start_percent": start_percent,
|
||||
"end_percent": end_percent,
|
||||
"cfg_scale": portrait_cfg,
|
||||
}
|
||||
|
||||
updated = dict(embeds)
|
||||
|
||||
+7
-143
@@ -1,15 +1,11 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import numpy as np
|
||||
import gguf
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
from .gguf_utils import GGUFParameter, dequantize_gguf_tensor
|
||||
from ..utils import log
|
||||
from .gguf_utils import GGUFParameter
|
||||
|
||||
def load_gguf(model_path):
|
||||
from gguf import GGUFReader
|
||||
reader = GGUFReader(model_path)
|
||||
reader = gguf.GGUFReader(model_path)
|
||||
parsed_parameters = {}
|
||||
for tensor in reader.tensors:
|
||||
# if the tensor is a torch supported dtype do not use GGUFParameter
|
||||
@@ -18,144 +14,12 @@ def load_gguf(model_path):
|
||||
parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor
|
||||
return parsed_parameters, reader
|
||||
|
||||
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
|
||||
from ..custom_linear import _replace_linear, set_lora_params, CustomLinear
|
||||
|
||||
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None):
|
||||
def _should_convert_to_gguf(state_dict, prefix):
|
||||
weight_key = prefix + "weight"
|
||||
return weight_key in state_dict and isinstance(state_dict[weight_key], GGUFParameter)
|
||||
|
||||
has_children = list(model.children())
|
||||
if not has_children:
|
||||
return
|
||||
|
||||
allow_compile = False
|
||||
|
||||
for name, module in model.named_children():
|
||||
if compile_args is not None:
|
||||
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
|
||||
module_prefix = prefix + name + "."
|
||||
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches, compile_args)
|
||||
|
||||
if (
|
||||
isinstance(module, nn.Linear)
|
||||
and not isinstance(module, GGUFLinear)
|
||||
and _should_convert_to_gguf(state_dict, module_prefix)
|
||||
and name not in modules_to_not_convert
|
||||
):
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
|
||||
with init_empty_weights():
|
||||
model._modules[name] = GGUFLinear(
|
||||
in_features,
|
||||
out_features,
|
||||
module.bias is not None,
|
||||
compute_dtype=compute_dtype,
|
||||
allow_compile=allow_compile
|
||||
)
|
||||
|
||||
model._modules[name].source_cls = type(module)
|
||||
model._modules[name].requires_grad_(False)
|
||||
return model
|
||||
return _replace_linear(model, compute_dtype, state_dict, prefix, patches, None, compile_args, modules_to_not_convert)
|
||||
|
||||
def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")):
|
||||
# Recursively set lora_diffs and lora_strengths for all GGUFLinear layers
|
||||
for name, child in module.named_children():
|
||||
params = list(child.parameters())
|
||||
if params:
|
||||
device = params[0].device
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
child_prefix = (f"{module_prefix}{name}.")
|
||||
set_lora_params_gguf(child, patches, child_prefix, device)
|
||||
if isinstance(module, GGUFLinear):
|
||||
key = f"diffusion_model.{module_prefix}weight"
|
||||
patch = patches.get(key, [])
|
||||
#print(f"Processing LoRA patches for {key}: {len(patch)} patches found")
|
||||
if len(patch) == 0:
|
||||
key = key.replace("_orig_mod.", "")
|
||||
patch = patches.get(key, [])
|
||||
if len(patch) != 0:
|
||||
lora_diffs = []
|
||||
for p in patch:
|
||||
lora_obj = p[1]
|
||||
if "head" in key:
|
||||
continue # For now skip LoRA for head layers
|
||||
elif hasattr(lora_obj, "weights"):
|
||||
lora_diffs.append(lora_obj.weights)
|
||||
elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff":
|
||||
lora_diffs.append(lora_obj[1])
|
||||
else:
|
||||
continue
|
||||
module.lora_strengths = [p[0] for p in patch]
|
||||
module.set_lora_diffs(lora_diffs, device=device)
|
||||
module.step = 0 # Initialize step for LoRA scheduling
|
||||
return set_lora_params(module, patches, module_prefix, device)
|
||||
|
||||
|
||||
class GGUFLinear(nn.Linear):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
out_features,
|
||||
bias=False,
|
||||
compute_dtype=None,
|
||||
device=None,
|
||||
allow_compile=False
|
||||
) -> None:
|
||||
super().__init__(in_features, out_features, bias, device)
|
||||
self.compute_dtype = compute_dtype
|
||||
self.lora_diffs = []
|
||||
self.lora_strengths = []
|
||||
self.step = 0
|
||||
self.allow_compile = allow_compile
|
||||
|
||||
if not allow_compile:
|
||||
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
|
||||
|
||||
def forward(self, inputs):
|
||||
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
|
||||
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
|
||||
|
||||
weight = self._get_weight_with_lora(weight)#.to(self.compute_dtype)
|
||||
|
||||
return torch.nn.functional.linear(inputs, weight, bias)
|
||||
|
||||
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
|
||||
self.lora_diffs = []
|
||||
for i, diff in enumerate(lora_diffs):
|
||||
if len(diff) > 1:
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
|
||||
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
|
||||
setattr(self, f"lora_diff_{i}_2", diff[2])
|
||||
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
|
||||
else:
|
||||
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
|
||||
self.lora_diffs.append(f"lora_diff_{i}_0")
|
||||
|
||||
def _get_weight_with_lora(self, weight):
|
||||
"""Apply LoRA outside compiled region"""
|
||||
if not hasattr(self, "lora_diff_0_0"):
|
||||
return weight
|
||||
|
||||
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
|
||||
if isinstance(lora_strength, list):
|
||||
lora_strength = lora_strength[self.step]
|
||||
if lora_strength == 0.0:
|
||||
continue
|
||||
elif lora_strength == 0.0:
|
||||
continue
|
||||
if isinstance(lora_diff_names, tuple):
|
||||
lora_diff_0 = getattr(self, lora_diff_names[0])
|
||||
lora_diff_1 = getattr(self, lora_diff_names[1])
|
||||
lora_diff_2 = getattr(self, lora_diff_names[2])
|
||||
patch_diff = torch.mm(
|
||||
lora_diff_0.flatten(start_dim=1),
|
||||
lora_diff_1.flatten(start_dim=1)
|
||||
).reshape(weight.shape) + 0
|
||||
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0
|
||||
scale = lora_strength * alpha
|
||||
weight = weight.add(patch_diff, alpha=scale)
|
||||
else:
|
||||
lora_diff = getattr(self, lora_diff_names)
|
||||
weight = weight.add(lora_diff, alpha=lora_strength)
|
||||
return weight
|
||||
GGUFLinear = CustomLinear
|
||||
@@ -0,0 +1,483 @@
|
||||
import torch
|
||||
import os
|
||||
import gc
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from ..latent_preview import prepare_callback
|
||||
from ..wanvideo.schedulers import get_scheduler
|
||||
from .multitalk import timestep_transform, add_noise
|
||||
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap
|
||||
from comfy.utils import load_torch_file
|
||||
from ..nodes_model_loading import load_weights
|
||||
from ..HuMo.nodes import get_audio_emb_window
|
||||
import comfy.model_management as mm
|
||||
from tqdm import tqdm
|
||||
import copy
|
||||
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
PATCH_SIZE = (1, 2, 2)
|
||||
vae_upscale_factor = 8
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
def multitalk_loop(self, **kwargs):
|
||||
# Unpack kwargs into local variables
|
||||
(latent, total_steps, steps, start_step, end_step, shift, cfg, denoise_strength,
|
||||
sigmas, weight_dtype, transformer, patcher, block_swap_args, model, vae, dtype,
|
||||
scheduler, scheduler_step_args, text_embeds, image_embeds, multitalk_embeds,
|
||||
multitalk_audio_embeds, unianim_data, dwpose_data, unianimate_poses, uni3c_embeds,
|
||||
humo_image_cond, humo_image_cond_neg, humo_audio, humo_reference_count,
|
||||
add_noise_to_samples, audio_stride, use_tsr, tsr_k, tsr_sigma, fantasy_portrait_input,
|
||||
noise, timesteps, force_offload, add_cond, control_latents, audio_proj,
|
||||
control_camera_latents, samples, masks, seed_g, gguf_reader, predict_func
|
||||
) = (kwargs.get(k) for k in (
|
||||
'latent', 'total_steps', 'steps', 'start_step', 'end_step', 'shift', 'cfg',
|
||||
'denoise_strength', 'sigmas', 'weight_dtype', 'transformer', 'patcher',
|
||||
'block_swap_args', 'model', 'vae', 'dtype', 'scheduler', 'scheduler_step_args',
|
||||
'text_embeds', 'image_embeds', 'multitalk_embeds', 'multitalk_audio_embeds',
|
||||
'unianim_data', 'dwpose_data', 'unianimate_poses', 'uni3c_embeds',
|
||||
'humo_image_cond', 'humo_image_cond_neg', 'humo_audio', 'humo_reference_count',
|
||||
'add_noise_to_samples', 'audio_stride', 'use_tsr', 'tsr_k', 'tsr_sigma',
|
||||
'fantasy_portrait_input', 'noise', 'timesteps', 'force_offload', 'add_cond',
|
||||
'control_latents', 'audio_proj', 'control_camera_latents', 'samples', 'masks',
|
||||
'seed_g', 'gguf_reader', 'predict_with_cfg'
|
||||
))
|
||||
|
||||
mode = image_embeds.get("multitalk_mode", "multitalk")
|
||||
if mode == "auto":
|
||||
mode = transformer.multitalk_model_type.lower()
|
||||
log.info(f"Multitalk mode: {mode}")
|
||||
cond_frame = None
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
offloaded = False
|
||||
tiled_vae = image_embeds.get("tiled_vae", False)
|
||||
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
|
||||
|
||||
clip_embeds = image_embeds.get("clip_context", None)
|
||||
if clip_embeds is not None:
|
||||
clip_embeds = clip_embeds.to(dtype)
|
||||
colormatch = image_embeds.get("colormatch", "disabled")
|
||||
motion_frame = image_embeds.get("motion_frame", 25)
|
||||
target_w = image_embeds.get("target_w", None)
|
||||
target_h = image_embeds.get("target_h", None)
|
||||
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
|
||||
if original_images is None:
|
||||
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
|
||||
|
||||
output_path = image_embeds.get("output_path", "")
|
||||
img_counter = 0
|
||||
|
||||
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
|
||||
face_scale = 0.1
|
||||
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
|
||||
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
|
||||
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
|
||||
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
|
||||
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
|
||||
human_mask2[x_min:x_max, righty_min:righty_max] = 1
|
||||
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
|
||||
human_masks = [human_mask1, human_mask2, background_mask]
|
||||
ref_target_masks = torch.stack(human_masks, dim=0)
|
||||
multitalk_embeds['ref_target_masks'] = ref_target_masks
|
||||
|
||||
gen_video_list = []
|
||||
is_first_clip = True
|
||||
arrive_last_frame = False
|
||||
cur_motion_frames_num = 1
|
||||
audio_start_idx = iteration_count = step_iteration_count = 0
|
||||
audio_end_idx = (audio_start_idx + clip_length) * audio_stride
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
current_condframe_index = 0
|
||||
|
||||
audio_embedding = multitalk_audio_embeds
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
cond_frame = None
|
||||
|
||||
uni3c_data = None
|
||||
if uni3c_embeds is not None:
|
||||
transformer.controlnet = uni3c_embeds["controlnet"]
|
||||
uni3c_data = uni3c_embeds.copy()
|
||||
|
||||
encoded_silence = None
|
||||
|
||||
try:
|
||||
silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
|
||||
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
|
||||
except:
|
||||
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
|
||||
|
||||
total_frames = len(audio_embedding[0])
|
||||
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
|
||||
callback = prepare_callback(patcher, estimated_iterations)
|
||||
|
||||
if frame_num >= total_frames:
|
||||
arrive_last_frame = True
|
||||
estimated_iterations = 1
|
||||
|
||||
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
|
||||
while True: # start video generation iteratively
|
||||
self.cache_state = [None, None]
|
||||
|
||||
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
|
||||
if mode == "infinitetalk":
|
||||
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
|
||||
if multitalk_embeds is not None:
|
||||
audio_embs = []
|
||||
# split audio with window size
|
||||
for human_idx in range(human_num):
|
||||
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
|
||||
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1)
|
||||
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
|
||||
audio_embs.append(audio_emb)
|
||||
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
|
||||
|
||||
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
|
||||
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
|
||||
latent_frame_num = (frame_num - 1) // 4 + 1
|
||||
|
||||
noise = torch.randn(
|
||||
16, latent_frame_num,
|
||||
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
|
||||
|
||||
# Calculate the correct latent slice based on current iteration
|
||||
if is_first_clip:
|
||||
latent_start_idx = 0
|
||||
latent_end_idx = noise.shape[1]
|
||||
else:
|
||||
new_frames_per_iteration = frame_num - motion_frame
|
||||
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
|
||||
latent_start_idx = iteration_count * new_latent_frames_per_iteration
|
||||
latent_end_idx = latent_start_idx + noise.shape[1]
|
||||
|
||||
if samples is not None:
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
input_samples = samples["samples"]
|
||||
if input_samples is not None:
|
||||
input_samples = input_samples.squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if latent_end_idx > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = latent_end_idx - input_samples.shape[1]
|
||||
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]}"
|
||||
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
# diff diff prep
|
||||
if noise_mask is not None:
|
||||
if len(noise_mask.shape) == 4:
|
||||
noise_mask = noise_mask.squeeze(1)
|
||||
if audio_end_idx > noise_mask.shape[0]:
|
||||
noise_mask = noise_mask.repeat(audio_end_idx // noise_mask.shape[0], 1, 1)
|
||||
noise_mask = noise_mask[audio_start_idx:audio_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.reshape(-1, 1, 1, 1, 1).to(device)
|
||||
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
|
||||
|
||||
# zero padding and vae encode for img cond
|
||||
if cond_image is not None or cond_frame is not None:
|
||||
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
|
||||
cond_frame_num = cond_.shape[2]
|
||||
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
|
||||
padding_frames_pixels_values = torch.concat([cond_.to(device, vae.dtype), video_frames], dim=2)
|
||||
|
||||
# encode
|
||||
vae.to(device)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
|
||||
|
||||
if mode == "multitalk":
|
||||
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
|
||||
else:
|
||||
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)
|
||||
|
||||
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
|
||||
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
|
||||
msk[:, :1] = 1
|
||||
y = torch.cat([msk, y]) # 4+C T H W
|
||||
mm.soft_empty_cache()
|
||||
else:
|
||||
y = None
|
||||
latent_motion_frames = noise[:, :1]
|
||||
|
||||
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
|
||||
if humo_image_cond is not None:
|
||||
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
|
||||
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
|
||||
if y is not None:
|
||||
partial_humo_cond_input[:, :1] = y[:, :1]
|
||||
if humo_reference_count > 0:
|
||||
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
|
||||
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
|
||||
|
||||
if humo_audio is not None:
|
||||
if is_first_clip:
|
||||
audio_embs = None
|
||||
|
||||
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
|
||||
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
|
||||
partial_humo_audio[-humo_reference_count:] = 0
|
||||
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
|
||||
|
||||
if scheduler == "multitalk":
|
||||
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
|
||||
timesteps.append(0.)
|
||||
timesteps = [torch.tensor([t], device=device) for t in timesteps]
|
||||
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
|
||||
else:
|
||||
if isinstance(scheduler, dict):
|
||||
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
|
||||
timesteps = scheduler["timesteps"]
|
||||
else:
|
||||
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas)
|
||||
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
|
||||
|
||||
# sample videos
|
||||
latent = noise
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip and mode == "multitalk":
|
||||
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])
|
||||
latent[:, :add_latent.shape[1]] = add_latent
|
||||
|
||||
if offloaded:
|
||||
# Load weights
|
||||
if transformer.patched_linear and gguf_reader is None:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
|
||||
elif gguf_reader is not None: #handle GGUF
|
||||
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
|
||||
#blockswap init
|
||||
init_blockswap(transformer, block_swap_args, model)
|
||||
|
||||
# Use the appropriate prompt for this section
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
|
||||
positive = [text_embeds["prompt_embeds"][prompt_index]]
|
||||
log.info(f"Using prompt index: {prompt_index}")
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
# uni3c slices
|
||||
if uni3c_embeds is not None:
|
||||
vae.to(device)
|
||||
# Pad original_images if needed
|
||||
num_frames = original_images.shape[2]
|
||||
if audio_end_idx > num_frames:
|
||||
pad_len = audio_end_idx - num_frames
|
||||
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
|
||||
padded_images = torch.cat([original_images, last_frame], dim=2)
|
||||
else:
|
||||
padded_images = original_images
|
||||
render_latent = vae.encode(
|
||||
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
|
||||
device=device, tiled=tiled_vae
|
||||
).to(dtype)
|
||||
|
||||
vae.to(offload_device)
|
||||
uni3c_data['render_latent'] = render_latent
|
||||
|
||||
# unianimate slices
|
||||
partial_unianim_data = None
|
||||
if unianim_data is not None:
|
||||
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
|
||||
partial_unianim_data = {
|
||||
"dwpose": partial_dwpose,
|
||||
"random_ref": unianim_data["random_ref"],
|
||||
"strength": unianimate_poses["strength"],
|
||||
"start_percent": unianimate_poses["start_percent"],
|
||||
"end_percent": unianimate_poses["end_percent"]
|
||||
}
|
||||
|
||||
# fantasy portrait slices
|
||||
partial_fantasy_portrait_input = None
|
||||
if fantasy_portrait_input is not None:
|
||||
adapter_proj = fantasy_portrait_input["adapter_proj"]
|
||||
if latent_end_idx > adapter_proj.shape[1]:
|
||||
pad_len = latent_end_idx - adapter_proj.shape[1]
|
||||
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
|
||||
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
|
||||
else:
|
||||
padded_proj = adapter_proj
|
||||
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
|
||||
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
# sampling loop
|
||||
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]
|
||||
latent_model_input = latent.to(device)
|
||||
if mode == "infinitetalk":
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
noise_pred, _, self.cache_state = predict_func(
|
||||
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, i, y, clip_embeds, control_latents, None, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input,
|
||||
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg,
|
||||
uni3c_data = uni3c_data)
|
||||
|
||||
if callback is not None:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * timestep.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 use_tsr:
|
||||
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
|
||||
if scheduler == "multitalk":
|
||||
noise_pred = -noise_pred
|
||||
dt = (timesteps[i] - timesteps[i + 1]) / 1000
|
||||
latent = latent + noise_pred * dt[:, None, None, None]
|
||||
else:
|
||||
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 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
|
||||
if not is_first_clip and mode == "multitalk":
|
||||
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])
|
||||
latent[:, :add_latent.shape[1]] = add_latent
|
||||
else:
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
del noise, latent_motion_frames
|
||||
if offload:
|
||||
offload_transformer(transformer, remove_lora=False)
|
||||
offloaded = True
|
||||
if humo_image_cond is not None and humo_reference_count > 0:
|
||||
latent = latent[:,:-humo_reference_count]
|
||||
vae.to(device)
|
||||
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()
|
||||
|
||||
# optional color correction (less relevant for InfiniteTalk)
|
||||
if colormatch != "disabled":
|
||||
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().float().numpy(), method=colormatch)
|
||||
else:
|
||||
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).permute(3, 0, 1, 2)
|
||||
|
||||
# optionally save generated samples to disk
|
||||
if output_path:
|
||||
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
|
||||
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
|
||||
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
|
||||
start_idx = 0 if is_first_clip else cur_motion_frames_num
|
||||
for i in range(start_idx, video_np.shape[0]):
|
||||
im = Image.fromarray(video_np[i])
|
||||
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
|
||||
img_counter += 1
|
||||
else:
|
||||
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
|
||||
|
||||
current_condframe_index += 1
|
||||
iteration_count += 1
|
||||
|
||||
# decide whether is done
|
||||
if arrive_last_frame:
|
||||
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 = cond_
|
||||
else:
|
||||
cond_image = cond_
|
||||
|
||||
del videos, latent
|
||||
|
||||
# Repeat audio emb
|
||||
if multitalk_embeds is not None:
|
||||
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
if audio_end_idx >= len(audio_embedding[0]):
|
||||
arrive_last_frame = True
|
||||
miss_lengths = []
|
||||
source_frames = []
|
||||
for human_inx in range(human_num):
|
||||
source_frame = len(audio_embedding[human_inx])
|
||||
source_frames.append(source_frame)
|
||||
if audio_end_idx >= len(audio_embedding[human_inx]):
|
||||
log.warning(f"Audio embedding for subject {human_inx} not long enough: {len(audio_embedding[human_inx])}, need {audio_end_idx}, padding...")
|
||||
miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3
|
||||
log.warning(f"Padding length: {miss_length}")
|
||||
if encoded_silence is not None:
|
||||
add_audio_emb = encoded_silence[-1*miss_length:]
|
||||
else:
|
||||
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
|
||||
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb.to(device, dtype)], dim=0)
|
||||
miss_lengths.append(miss_length)
|
||||
else:
|
||||
miss_lengths.append(0)
|
||||
if mode == "infinitetalk" and current_condframe_index >= original_images.shape[2]:
|
||||
last_frame = original_images[:, :, -1:, :, :]
|
||||
miss_length = 1
|
||||
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
|
||||
|
||||
if not output_path:
|
||||
gen_video_samples = torch.cat(gen_video_list, dim=1)
|
||||
else:
|
||||
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
|
||||
|
||||
if force_offload:
|
||||
if not model["auto_cpu_offload"]:
|
||||
offload_transformer(transformer)
|
||||
try:
|
||||
print_memory(device)
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
|
||||
+19
-1
@@ -7,6 +7,8 @@ from ..utils import log, set_module_tensor_to_device
|
||||
import os
|
||||
import json
|
||||
import datetime
|
||||
import scipy.signal as ss
|
||||
import numpy as np
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
folder_paths.add_model_folder_path("wav2vec2", os.path.join(folder_paths.models_dir, "wav2vec2"))
|
||||
@@ -134,6 +136,15 @@ def loudness_norm(audio_array, sr=16000, lufs=-23):
|
||||
return audio_array
|
||||
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
|
||||
return normalized_audio
|
||||
|
||||
def _add_noise_floor(audio, noise_db=-45):
|
||||
noise_amp = 10 ** (noise_db / 20)
|
||||
noise = np.random.randn(len(audio)) * noise_amp
|
||||
return audio + noise
|
||||
|
||||
def _smooth_transients(audio, sr=16000):
|
||||
b, a = ss.butter(3, 3000 / (sr/2))
|
||||
return ss.lfilter(b, a, audio)
|
||||
|
||||
class MultiTalkWav2VecEmbeds:
|
||||
@classmethod
|
||||
@@ -153,6 +164,8 @@ class MultiTalkWav2VecEmbeds:
|
||||
"audio_3": ("AUDIO",),
|
||||
"audio_4": ("AUDIO",),
|
||||
"ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space. Supply one mask per speaker (plus optional background) to guide mouth assignment"}),
|
||||
"add_noise_floor": ("BOOLEAN", {"default": False, "tooltip": "Add a low-level noise floor to the audio to reduce silent gaps"}),
|
||||
"smooth_transients": ("BOOLEAN", {"default": False, "tooltip": "Apply a low-pass filter to the audio to smooth out transients"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -161,7 +174,8 @@ class MultiTalkWav2VecEmbeds:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None, ref_target_masks=None):
|
||||
def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio_1, audio_scale, audio_cfg_scale, multi_audio_type, audio_2=None, audio_3=None, audio_4=None,
|
||||
ref_target_masks=None, add_noise_floor=False, smooth_transients=False):
|
||||
model_type = wav2vec_model["model_type"]
|
||||
if not "tencent" in model_type.lower():
|
||||
raise ValueError("Only tencent wav2vec2 models supported by MultiTalk")
|
||||
@@ -207,6 +221,10 @@ class MultiTalkWav2VecEmbeds:
|
||||
|
||||
if normalize_loudness:
|
||||
audio_segment = loudness_norm(audio_segment, sr=sr)
|
||||
if add_noise_floor:
|
||||
audio_segment = _add_noise_floor(audio_segment, noise_db=-45)
|
||||
if smooth_transients:
|
||||
audio_segment = _smooth_transients(audio_segment, sr=sr)
|
||||
|
||||
audio_feature = np.squeeze(
|
||||
wav2vec2_feature_extractor(audio_segment, sampling_rate=sr).input_values
|
||||
|
||||
@@ -1,11 +1,8 @@
|
||||
import os, gc, math
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
import hashlib
|
||||
|
||||
from .wanvideo.schedulers import get_scheduler, scheduler_list
|
||||
|
||||
from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device)
|
||||
from .taehv import TAEHV
|
||||
|
||||
@@ -22,30 +19,6 @@ offload_device = mm.unet_offload_device()
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
PATCH_SIZE = (1, 2, 2)
|
||||
|
||||
class WanVideoAddVideoPromptEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"video_prompt_embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"video_prompt_latents": ("LATENT", ),
|
||||
"text_embeds": ("WANVIDEOTEXTEMBEDS", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def add(self, image_embeds, video_prompt_embeds, video_prompt_latents, text_embeds):
|
||||
updated = dict(image_embeds)
|
||||
updated["video_prompt_embeds"] = video_prompt_embeds
|
||||
updated["video_prompt_embeds"]["video_prompt_latents"] = video_prompt_latents["samples"][0]
|
||||
updated["video_prompt_embeds"]["text_embeds"] = text_embeds
|
||||
return (updated,)
|
||||
|
||||
|
||||
class WanVideoEnhanceAVideo:
|
||||
@classmethod
|
||||
@@ -789,6 +762,98 @@ class WanVideoAddStandInLatent:
|
||||
updated = dict(embeds)
|
||||
updated["standin_input"] = new_entry
|
||||
return (updated,)
|
||||
|
||||
class WanVideoAddBindweaveEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"reference_latents": ("LATENT", {"tooltip": "Reference image to encode"}),
|
||||
},
|
||||
"optional": {
|
||||
"ref_masks": ("MASK", {"tooltip": "Reference mask to encode"}),
|
||||
"qwenvl_embeds_pos": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
|
||||
"qwenvl_embeds_neg": ("QWENVL_EMBEDS", {"tooltip": "Qwen-VL image embeddings for the reference image"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "LATENT", "MASK",)
|
||||
RETURN_NAMES = ("image_embeds", "image_embed_preview", "mask_preview",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, reference_latents, ref_masks=None, qwenvl_embeds_pos=None, qwenvl_embeds_neg=None):
|
||||
updated = dict(embeds)
|
||||
image_embeds = embeds["image_embeds"]
|
||||
max_refs = 4
|
||||
num_refs = reference_latents["samples"].shape[0]
|
||||
pad = torch.zeros(image_embeds.shape[0], max_refs-num_refs, image_embeds.shape[2], image_embeds.shape[3], device=image_embeds.device, dtype=image_embeds.dtype)
|
||||
if num_refs < max_refs:
|
||||
image_embeds = torch.cat([pad, image_embeds], dim=1)
|
||||
ref_latents = [ref_latent for ref_latent in reference_latents["samples"]]
|
||||
image_embeds = torch.cat([*ref_latents, image_embeds], dim=1)
|
||||
|
||||
mask = embeds.get("mask", None)
|
||||
if mask is not None:
|
||||
mask_pad = torch.zeros(mask.shape[0], max_refs-num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype)
|
||||
if num_refs < max_refs:
|
||||
mask = torch.cat([mask_pad, mask], dim=1)
|
||||
if ref_masks is not None:
|
||||
ref_mask_ = common_upscale(ref_masks.unsqueeze(1), mask.shape[3], mask.shape[2], "nearest", "disabled").movedim(0,1)
|
||||
ref_mask_ = torch.cat([ref_mask_, torch.zeros(3, ref_mask_.shape[1], ref_mask_.shape[2], ref_mask_.shape[3], device=ref_mask_.device, dtype=ref_mask_.dtype)])
|
||||
mask = torch.cat([ref_mask_, mask], dim=1)
|
||||
else:
|
||||
mask = torch.cat([torch.ones(mask.shape[0], num_refs, mask.shape[2], mask.shape[3], device=mask.device, dtype=mask.dtype), mask], dim=1)
|
||||
|
||||
updated["mask"] = mask
|
||||
|
||||
clip_embeds = updated.get("clip_context", None)
|
||||
if clip_embeds is not None:
|
||||
B, T, C = clip_embeds.shape
|
||||
target_len = max_refs * 257 # 4 * 257 = 1028
|
||||
if T < target_len:
|
||||
pad = torch.zeros(B, target_len - T, C, device=clip_embeds.device, dtype=clip_embeds.dtype)
|
||||
padded_embeds = torch.cat([clip_embeds, pad], dim=1)
|
||||
log.info(f"Padded clip embeds from {clip_embeds.shape} to {padded_embeds.shape} for Bindweave")
|
||||
updated["clip_context"] = padded_embeds
|
||||
else:
|
||||
updated["clip_context"] = clip_embeds
|
||||
|
||||
updated["image_embeds"] = image_embeds
|
||||
updated["qwenvl_embeds_pos"] = qwenvl_embeds_pos
|
||||
updated["qwenvl_embeds_neg"] = qwenvl_embeds_neg
|
||||
return (updated, {"samples": image_embeds.unsqueeze(0)}, mask[0].float())
|
||||
|
||||
class TextImageEncodeQwenVL():
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"clip": ("CLIP",),
|
||||
"prompt": ("STRING", {"default": "", "multiline": True}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("QWENVL_EMBEDS",)
|
||||
RETURN_NAMES = ("qwenvl_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(cls, clip, prompt, image=None):
|
||||
if image is None:
|
||||
input_images = []
|
||||
llama_template = None
|
||||
else:
|
||||
input_images = [image[:, :, :, :3]]
|
||||
|
||||
llama_template = "<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>{}<|im_end|>\n<|im_start|>assistant\n"
|
||||
|
||||
tokens = clip.tokenize(prompt, images=input_images, llama_template=llama_template)
|
||||
conditioning = clip.encode_from_tokens_scheduled(tokens)
|
||||
print("Qwen-VL embeds shape:", conditioning[0][0].shape)
|
||||
return (conditioning[0][0],)
|
||||
|
||||
class WanVideoAddMTVMotion:
|
||||
@classmethod
|
||||
@@ -823,6 +888,83 @@ class WanVideoAddMTVMotion:
|
||||
updated["mtv_crafter_motion"] = new_entry
|
||||
return (updated,)
|
||||
|
||||
class WanVideoAddStoryMemLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"vae": ("WANVAE",),
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"memory_images": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, vae, embeds, memory_images):
|
||||
updated = dict(embeds)
|
||||
story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images)
|
||||
updated["story_mem_latents"] = story_mem_latents["samples"].squeeze(2).permute(1, 0, 2, 3) # [C, T, H, W]
|
||||
return (updated,)
|
||||
|
||||
|
||||
class WanVideoSVIProEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"anchor_samples": ("LATENT", {"tooltip": "Initial start image encoded"}),
|
||||
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
|
||||
},
|
||||
"optional": {
|
||||
"prev_samples": ("LATENT", {"tooltip": "Last latent from previous generation"}),
|
||||
"motion_latent_count": ("INT", {"default": 1, "min": 0, "max": 100, "step": 1, "tooltip": "Number of latents used to continue"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, anchor_samples, num_frames, prev_samples=None, motion_latent_count=1):
|
||||
|
||||
anchor_latent = anchor_samples["samples"][0].clone()
|
||||
|
||||
C, T, H, W = anchor_latent.shape
|
||||
|
||||
total_latents = (num_frames - 1) // 4 + 1
|
||||
device = anchor_latent.device
|
||||
dtype = anchor_latent.dtype
|
||||
|
||||
if prev_samples is None or motion_latent_count == 0:
|
||||
padding_size = total_latents - anchor_latent.shape[1]
|
||||
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
|
||||
y = torch.concat([anchor_latent, padding], dim=1)
|
||||
else:
|
||||
prev_latent = prev_samples["samples"][0].clone()
|
||||
motion_latent = prev_latent[:, -motion_latent_count:]
|
||||
padding_size = total_latents - anchor_latent.shape[1] - motion_latent.shape[1]
|
||||
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
|
||||
y = torch.concat([anchor_latent, motion_latent, padding], dim=1)
|
||||
|
||||
msk = torch.ones(1, num_frames, H, W, device=device, dtype=dtype)
|
||||
msk[:, 1:] = 0
|
||||
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
|
||||
msk = msk.view(1, msk.shape[1] // 4, 4, H, W)
|
||||
msk = msk.transpose(1, 2)[0]
|
||||
|
||||
image_embeds = {
|
||||
"image_embeds": y,
|
||||
"num_frames": num_frames,
|
||||
"lat_h": H,
|
||||
"lat_w": W,
|
||||
"mask": msk
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
#region I2V encode
|
||||
class WanVideoImageToVideoEncode:
|
||||
@classmethod
|
||||
@@ -847,6 +989,8 @@ class WanVideoImageToVideoEncode:
|
||||
"extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}),
|
||||
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
|
||||
"add_cond_latents": ("ADD_COND_LATENTS", {"advanced": True, "tooltip": "Additional cond latents WIP"}),
|
||||
"augment_empty_frames": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "EXPERIMENTAL: Augment empty frames with the difference to the start image to force more motion"}),
|
||||
"empty_frame_pad_image": ("IMAGE", {"tooltip": "Use this image to pad empty frames instead of gray, used with SVI-shot and SVI 2.0 LoRAs"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -856,18 +1000,14 @@ class WanVideoImageToVideoEncode:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, width, height, num_frames, force_offload, noise_aug_strength,
|
||||
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
|
||||
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None):
|
||||
|
||||
if start_image is None and end_image is None and add_cond_latents is None:
|
||||
return WanVideoEmptyEmbeds().process(
|
||||
num_frames, width, height, control_embeds=control_embeds, extra_latents=extra_latents,
|
||||
)
|
||||
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
|
||||
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, augment_empty_frames=0.0, empty_frame_pad_image=None):
|
||||
|
||||
if vae is None:
|
||||
raise ValueError("VAE is required for image encoding.")
|
||||
H = height
|
||||
W = width
|
||||
|
||||
|
||||
lat_h = H // vae.upsampling_factor
|
||||
lat_w = W // vae.upsampling_factor
|
||||
|
||||
@@ -892,6 +1032,8 @@ class WanVideoImageToVideoEncode:
|
||||
mask = torch.cat([mask, torch.zeros(base_frames - mask.shape[0], lat_h, lat_w, device=device)])
|
||||
mask = mask.unsqueeze(0).to(device, vae.dtype)
|
||||
|
||||
pixel_mask = mask.clone()
|
||||
|
||||
# Repeat first frame and optionally end frame
|
||||
start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W
|
||||
if end_image is not None and not fun_or_fl2v_model:
|
||||
@@ -914,7 +1056,7 @@ class WanVideoImageToVideoEncode:
|
||||
resized_start_image = resized_start_image * 2 - 1
|
||||
if noise_aug_strength > 0.0:
|
||||
resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength)
|
||||
|
||||
|
||||
if end_image is not None:
|
||||
end_image = end_image[..., :3]
|
||||
if end_image.shape[1] != H or end_image.shape[2] != W:
|
||||
@@ -924,30 +1066,46 @@ class WanVideoImageToVideoEncode:
|
||||
resized_end_image = resized_end_image * 2 - 1
|
||||
if noise_aug_strength > 0.0:
|
||||
resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength)
|
||||
|
||||
|
||||
# Concatenate image with zero frames and encode
|
||||
if temporal_mask is None:
|
||||
if start_image is not None and end_image is None:
|
||||
zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype)
|
||||
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1)
|
||||
del resized_start_image, zero_frames
|
||||
elif start_image is None and end_image is not None:
|
||||
zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype)
|
||||
concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
|
||||
del zero_frames
|
||||
elif start_image is None and end_image is None:
|
||||
concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
|
||||
else:
|
||||
if fun_or_fl2v_model:
|
||||
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
|
||||
else:
|
||||
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype)
|
||||
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
|
||||
del resized_start_image, zero_frames
|
||||
if start_image is not None and end_image is None:
|
||||
zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype)
|
||||
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1)
|
||||
del resized_start_image, zero_frames
|
||||
elif start_image is None and end_image is not None:
|
||||
zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype)
|
||||
concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
|
||||
del zero_frames
|
||||
elif start_image is None and end_image is None:
|
||||
concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
|
||||
else:
|
||||
temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1)
|
||||
concatenated = resized_start_image[:,:num_frames].to(vae.dtype)# * temporal_mask[:num_frames].unsqueeze(0).to(vae.dtype)
|
||||
del resized_start_image, temporal_mask
|
||||
if fun_or_fl2v_model:
|
||||
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
|
||||
else:
|
||||
zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype)
|
||||
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
|
||||
del resized_start_image, zero_frames
|
||||
|
||||
if empty_frame_pad_image is not None:
|
||||
pad_img = empty_frame_pad_image.clone()[..., :3]
|
||||
if pad_img.shape[1] != H or pad_img.shape[2] != W:
|
||||
pad_img = common_upscale(pad_img.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
|
||||
pad_img = (pad_img.movedim(-1, 0) * 2 - 1).to(device, dtype=vae.dtype)
|
||||
|
||||
num_pad_frames = pad_img.shape[1]
|
||||
num_target_frames = concatenated.shape[1]
|
||||
if num_pad_frames < num_target_frames:
|
||||
pad_img = torch.cat([pad_img, pad_img[:, -1:].expand(-1, num_target_frames - num_pad_frames, -1, -1)], dim=1)
|
||||
else:
|
||||
pad_img = pad_img[:, :num_target_frames]
|
||||
|
||||
frame_is_empty = (pixel_mask[0].mean(dim=(-2, -1)) < 0.5)[:concatenated.shape[1]].clone()
|
||||
if start_image is not None:
|
||||
frame_is_empty[:start_image.shape[0]] = False
|
||||
if end_image is not None:
|
||||
frame_is_empty[-end_image.shape[0]:] = False
|
||||
|
||||
concatenated[:, frame_is_empty] = pad_img[:, frame_is_empty]
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
@@ -965,6 +1123,9 @@ class WanVideoImageToVideoEncode:
|
||||
has_ref = True
|
||||
y[:, :1] *= start_latent_strength
|
||||
y[:, -1:] *= end_latent_strength
|
||||
if augment_empty_frames > 0.0:
|
||||
frame_is_empty = (mask[0].mean(dim=(-2, -1)) < 0.5).view(1, -1, 1, 1)
|
||||
y = y[:, :1] + (y - y[:, :1]) * ((augment_empty_frames+1) * frame_is_empty + ~frame_is_empty)
|
||||
|
||||
# Calculate maximum sequence length
|
||||
patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
|
||||
@@ -973,14 +1134,14 @@ class WanVideoImageToVideoEncode:
|
||||
|
||||
if add_cond_latents is not None:
|
||||
add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device)
|
||||
|
||||
|
||||
if force_offload:
|
||||
vae.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
image_embeds = {
|
||||
"image_embeds": y,
|
||||
"image_embeds": y.cpu(),
|
||||
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
|
||||
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None,
|
||||
"max_seq_len": max_seq_len,
|
||||
@@ -992,11 +1153,11 @@ class WanVideoImageToVideoEncode:
|
||||
"fun_or_fl2v_model": fun_or_fl2v_model,
|
||||
"has_ref": has_ref,
|
||||
"add_cond_latents": add_cond_latents,
|
||||
"mask": mask
|
||||
"mask": mask.cpu()
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
|
||||
# region WanAnimate
|
||||
class WanVideoAnimateEmbeds:
|
||||
@classmethod
|
||||
@@ -1007,15 +1168,15 @@ class WanVideoAnimateEmbeds:
|
||||
"height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}),
|
||||
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
"frame_window_size": ("INT", {"default": 77, "min": 1, "max": 1000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}),
|
||||
"frame_window_size": ("INT", {"default": 77, "min": 1, "max": 10000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}),
|
||||
"colormatch": (
|
||||
[
|
||||
[
|
||||
'disabled',
|
||||
'mkl',
|
||||
'hm',
|
||||
'reinhard',
|
||||
'mvgd',
|
||||
'hm-mvgd-hm',
|
||||
'hm',
|
||||
'reinhard',
|
||||
'mvgd',
|
||||
'hm-mvgd-hm',
|
||||
'hm-mkl-hm',
|
||||
], {
|
||||
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
|
||||
@@ -1182,6 +1343,46 @@ class WanVideoAnimateEmbeds:
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
# region UniLumos
|
||||
class WanVideoUniLumosEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}),
|
||||
"height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}),
|
||||
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
|
||||
},
|
||||
"optional": {
|
||||
"foreground_latents": ("LATENT", {"tooltip": "Video foreground latents"}),
|
||||
"background_latents": ("LATENT", {"tooltip": "Video background latents"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, num_frames, width, height, foreground_latents=None, background_latents=None):
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
height // VAE_STRIDE[1],
|
||||
width // VAE_STRIDE[2])
|
||||
|
||||
embeds = {
|
||||
"target_shape": target_shape,
|
||||
"num_frames": num_frames,
|
||||
}
|
||||
if foreground_latents is not None:
|
||||
embeds["foreground_latents"] = foreground_latents["samples"][0]
|
||||
else:
|
||||
embeds["foreground_latents"] = torch.zeros(target_shape[0], target_shape[1], target_shape[2], target_shape[3], device=torch.device("cpu"), dtype=torch.float32)
|
||||
if background_latents is not None:
|
||||
embeds["background_latents"] = background_latents["samples"][0]
|
||||
else:
|
||||
embeds["background_latents"] = torch.zeros(target_shape[0], target_shape[1], target_shape[2], target_shape[3], device=torch.device("cpu"), dtype=torch.float32)
|
||||
|
||||
return (embeds,)
|
||||
|
||||
class WanVideoEmptyEmbeds:
|
||||
@classmethod
|
||||
@@ -1702,33 +1903,7 @@ class WanVideoContextOptions:
|
||||
}
|
||||
|
||||
return (context_options,)
|
||||
|
||||
|
||||
class WanVideoFlowEdit:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"source_embeds": ("WANVIDEOTEXTEMBEDS", ),
|
||||
"skip_steps": ("INT", {"default": 4, "min": 0}),
|
||||
"drift_steps": ("INT", {"default": 0, "min": 0}),
|
||||
"drift_flow_shift": ("FLOAT", {"default": 3.0, "min": 1.0, "max": 30.0, "step": 0.01}),
|
||||
"source_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
|
||||
"drift_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
|
||||
},
|
||||
"optional": {
|
||||
"source_image_embeds": ("WANVIDIMAGE_EMBEDS", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOWEDITARGS", )
|
||||
RETURN_NAMES = ("flowedit_args",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Flowedit options for WanVideo"
|
||||
|
||||
def process(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
class WanVideoLoopArgs:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -1800,155 +1975,6 @@ class WanVideoFreeInitArgs:
|
||||
|
||||
def process(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
class WanVideoScheduler: #WIP
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"scheduler": (scheduler_list, {"default": "unipc"}),
|
||||
"steps": ("INT", {"default": 30, "min": 1, "tooltip": "Number of steps for the scheduler"}),
|
||||
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||
"start_step": ("INT", {"default": 0, "min": 0, "tooltip": "Starting step for the scheduler"}),
|
||||
"end_step": ("INT", {"default": -1, "min": -1, "tooltip": "Ending step for the scheduler"})
|
||||
},
|
||||
"optional": {
|
||||
"sigmas": ("SIGMAS", ),
|
||||
},
|
||||
"hidden": {
|
||||
"unique_id": "UNIQUE_ID",
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SIGMAS", "INT", "FLOAT", scheduler_list, "INT", "INT",)
|
||||
RETURN_NAMES = ("sigmas", "steps", "shift", "scheduler", "start_step", "end_step")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None):
|
||||
sample_scheduler, timesteps, start_idx, end_idx = get_scheduler(
|
||||
scheduler,
|
||||
steps,
|
||||
start_step, end_step, shift,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
log_timesteps=True)
|
||||
|
||||
scheduler_dict = {
|
||||
"sample_scheduler": sample_scheduler,
|
||||
"timesteps": timesteps,
|
||||
}
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
import io
|
||||
import base64
|
||||
import matplotlib.pyplot as plt
|
||||
except:
|
||||
PromptServer = None
|
||||
if unique_id and PromptServer is not None:
|
||||
try:
|
||||
# Plot sigmas and save to a buffer
|
||||
sigmas_np = sample_scheduler.full_sigmas.cpu().numpy()
|
||||
if not np.isclose(sigmas_np[-1], 0.0, atol=1e-6):
|
||||
sigmas_np = np.append(sigmas_np, 0.0)
|
||||
buf = io.BytesIO()
|
||||
fig = plt.figure(facecolor='#353535')
|
||||
ax = fig.add_subplot(111)
|
||||
ax.set_facecolor('#353535') # Set axes background color
|
||||
x_values = range(0, len(sigmas_np))
|
||||
ax.plot(x_values, sigmas_np)
|
||||
# Annotate each sigma value
|
||||
ax.scatter(x_values, sigmas_np, color='white', s=20, zorder=3) # Small dots at each sigma
|
||||
for x, y in zip(x_values, sigmas_np):
|
||||
# Show all annotations if few steps, or just show split step annotations
|
||||
show_annotation = len(sigmas_np) <= 10
|
||||
is_split_step = (start_idx > 0 and x == start_idx) or (end_idx != -1 and x == end_idx + 1)
|
||||
|
||||
if show_annotation or is_split_step:
|
||||
color = 'orange'
|
||||
if is_split_step:
|
||||
color = 'yellow'
|
||||
ax.annotate(f"{y:.3f}", (x, y), textcoords="offset points", xytext=(10, 1), ha='center', color=color, fontsize=12)
|
||||
ax.set_xticks(x_values)
|
||||
ax.set_title("Sigmas", color='white') # Title font color
|
||||
ax.set_xlabel("Step", color='white') # X label font color
|
||||
ax.set_ylabel("Sigma Value", color='white') # Y label font color
|
||||
ax.tick_params(axis='x', colors='white', labelsize=10) # X tick color
|
||||
ax.tick_params(axis='y', colors='white', labelsize=10) # Y tick color
|
||||
# Add split point if end_step is defined
|
||||
end_idx += 1
|
||||
if end_idx != -1 and 0 <= end_idx < len(sigmas_np) - 1:
|
||||
ax.axvline(end_idx, color='red', linestyle='--', linewidth=2, label='end_step split')
|
||||
# Add split point if start_step is defined
|
||||
if start_idx > 0 and 0 <= start_idx < len(sigmas_np):
|
||||
ax.axvline(start_idx, color='green', linestyle='--', linewidth=2, label='start_step split')
|
||||
if (end_idx != -1 and 0 <= end_idx < len(sigmas_np)) or (start_idx > 0 and 0 <= start_idx < len(sigmas_np)):
|
||||
handles, labels = ax.get_legend_handles_labels()
|
||||
if labels:
|
||||
ax.legend()
|
||||
if start_idx < end_idx and 0 <= start_idx < len(sigmas_np) and 0 < end_idx < len(sigmas_np):
|
||||
ax.axvspan(start_idx, end_idx, color='lightblue', alpha=0.1, label='Sampled Range')
|
||||
plt.tight_layout()
|
||||
plt.savefig(buf, format='png')
|
||||
plt.close(fig)
|
||||
buf.seek(0)
|
||||
img_base64 = base64.b64encode(buf.read()).decode('utf-8')
|
||||
buf.close()
|
||||
|
||||
# Send as HTML img tag with base64 data
|
||||
html_img = f"<img src='data:image/png;base64,{img_base64}' alt='Sigmas Plot' style='max-width:100%; height:100%; overflow:hidden; display:block;'>"
|
||||
PromptServer.instance.send_progress_text(html_img, unique_id)
|
||||
except Exception as e:
|
||||
print("Failed to send sigmas plot:", e)
|
||||
pass
|
||||
|
||||
return (sigmas, steps, shift, scheduler_dict, start_step, end_step)
|
||||
|
||||
class WanVideoSchedulerSA_ODE:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"use_adaptive_order": ("BOOLEAN", {"default": False, "tooltip": "Use adaptive order"}),
|
||||
"use_velocity_smoothing": ("BOOLEAN", {"default": True, "tooltip": "Use velocity smoothing"}),
|
||||
"convergence_threshold": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Convergence threshold for velocity smoothing"}),
|
||||
"smoothing_factor": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Smoothing factor for velocity smoothing"}),
|
||||
"steps": ("INT", {"default": 30, "min": 1, "tooltip": "Number of steps for the scheduler"}),
|
||||
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||
"start_step": ("INT", {"default": 0, "min": 0, "tooltip": "Starting step for the scheduler"}),
|
||||
"end_step": ("INT", {"default": -1, "min": -1, "tooltip": "Ending step for the scheduler"})
|
||||
},
|
||||
"optional": {
|
||||
"sigmas": ("SIGMAS", ),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SIGMAS", "INT", "FLOAT", scheduler_list, "INT", "INT",)
|
||||
RETURN_NAMES = ("sigmas", "steps", "shift", "scheduler", "start_step", "end_step")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def process(self, steps, start_step, end_step, shift, use_adaptive_order, use_velocity_smoothing, convergence_threshold, smoothing_factor, sigmas=None):
|
||||
sample_scheduler, timesteps, _, _ = get_scheduler(
|
||||
scheduler="sa_ode_stable/lowstep",
|
||||
steps=steps,
|
||||
start_step=start_step, end_step=end_step, shift=shift,
|
||||
device=device,
|
||||
sigmas=sigmas,
|
||||
log_timesteps=True,
|
||||
use_adaptive_order=use_adaptive_order,
|
||||
use_velocity_smoothing=use_velocity_smoothing,
|
||||
convergence_threshold=convergence_threshold,
|
||||
smoothing_factor=smoothing_factor
|
||||
)
|
||||
|
||||
scheduler_dict = {
|
||||
"sample_scheduler": sample_scheduler,
|
||||
"timesteps": timesteps,
|
||||
}
|
||||
|
||||
return (sigmas, steps, shift, scheduler_dict, start_step, end_step)
|
||||
|
||||
rope_functions = ["default", "comfy", "comfy_chunked"]
|
||||
class WanVideoRoPEFunction:
|
||||
@@ -1979,6 +2005,53 @@ class WanVideoRoPEFunction:
|
||||
return (rope_func_dict,)
|
||||
return (rope_function,)
|
||||
|
||||
#region TTM
|
||||
class WanVideoAddTTMLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"reference_latents": ("LATENT", {"tooltip": "Latents used as reference for TTM"}),
|
||||
"mask": ("MASK", {"tooltip": "Mask used for TTM"}),
|
||||
"start_step": ("INT", {"default": 0, "min": -1, "max": 1000, "step": 1, "tooltip": "Start step for whole denoising process"}),
|
||||
"end_step": ("INT", {"default": 1, "min": 1, "max": 1000, "step": 1, "tooltip": "The step to stop applying TTM"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
|
||||
RETURN_NAMES = ("image_embeds", )
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "https://github.com/time-to-move/TTM"
|
||||
|
||||
def add(self, embeds, reference_latents, mask, start_step, end_step):
|
||||
|
||||
if end_step < max(0, start_step):
|
||||
raise ValueError(f"`end_step` ({end_step}) must be >= `start_step` ({start_step}).")
|
||||
|
||||
mask_sampled = mask[::4]
|
||||
mask_sampled = mask_sampled.unsqueeze(1).unsqueeze(0) # [1, T, 1, H, W]
|
||||
|
||||
vae_upscale_factor = 8
|
||||
if reference_latents["samples"].shape[1] == 48:
|
||||
vae_upscale_factor = 16
|
||||
|
||||
# Upsample spatially to latent resolution
|
||||
H_latent = mask_sampled.shape[-2] // vae_upscale_factor
|
||||
W_latent = mask_sampled.shape[-1] // vae_upscale_factor
|
||||
mask_latent = F.interpolate(
|
||||
mask_sampled.float(),
|
||||
size=(mask_sampled.shape[2], H_latent, W_latent),
|
||||
mode="nearest"
|
||||
)
|
||||
|
||||
updated = dict(embeds)
|
||||
updated["ttm_reference_latents"] = reference_latents["samples"].squeeze(0)
|
||||
updated["ttm_mask"] = mask_latent.squeeze(0).movedim(1, 0) # [T, 1, H, W]
|
||||
updated["ttm_start_step"] = start_step
|
||||
updated["ttm_end_step"] = end_step
|
||||
|
||||
return (updated,)
|
||||
|
||||
#region VideoDecode
|
||||
class WanVideoDecode:
|
||||
@@ -2000,7 +2073,7 @@ class WanVideoDecode:
|
||||
"tile_stride_y": ("INT", {"default": 128, "min": 32, "max": 2040, "step": 8, "tooltip": "Tile stride height in pixels. Smaller values use less VRAM but will introduce more seams."}),
|
||||
},
|
||||
"optional": {
|
||||
"normalization": (["default", "minmax"], {"advanced": True}),
|
||||
"normalization": (["default", "minmax", "none"], {"advanced": True}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2024,7 +2097,7 @@ class WanVideoDecode:
|
||||
video.clamp_(-1.0, 1.0)
|
||||
video.add_(1.0).div_(2.0)
|
||||
return video.cpu().float(),
|
||||
latents = samples["samples"]
|
||||
latents = samples["samples"].clone()
|
||||
end_image = samples.get("end_image", None)
|
||||
has_ref = samples.get("has_ref", False)
|
||||
drop_last = samples.get("drop_last", False)
|
||||
@@ -2043,25 +2116,24 @@ class WanVideoDecode:
|
||||
if drop_last:
|
||||
latents = latents[:, :, :-1]
|
||||
|
||||
if type(vae).__name__ == "TAEHV":
|
||||
if type(vae).__name__ == "TAEHV":
|
||||
images = vae.decode_video(latents.permute(0, 2, 1, 3, 4), cond=flashvsr_LQ_images.to(vae.dtype) if flashvsr_LQ_images is not None else None)[0].permute(1, 0, 2, 3)
|
||||
images = torch.clamp(images, 0.0, 1.0)
|
||||
images = images.permute(1, 2, 3, 0).cpu().float()
|
||||
return (images,)
|
||||
else:
|
||||
if end_image is not None:
|
||||
enable_vae_tiling = False
|
||||
images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0]
|
||||
|
||||
|
||||
|
||||
|
||||
images = images.cpu().float()
|
||||
|
||||
if normalization == "minmax":
|
||||
images.sub_(images.min()).div_(images.max() - images.min())
|
||||
else:
|
||||
images.clamp_(-1.0, 1.0)
|
||||
images.add_(1.0).div_(2.0)
|
||||
|
||||
if normalization != "none":
|
||||
if normalization == "minmax":
|
||||
images.sub_(images.min()).div_(images.max() - images.min())
|
||||
else:
|
||||
images.clamp_(-1.0, 1.0)
|
||||
images.add_(1.0).div_(2.0)
|
||||
|
||||
if is_looped:
|
||||
temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2)
|
||||
temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0]
|
||||
@@ -2072,7 +2144,7 @@ class WanVideoDecode:
|
||||
if end_image is not None:
|
||||
images = images[:, 0:-1]
|
||||
|
||||
|
||||
|
||||
vae.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
|
||||
@@ -2101,7 +2173,7 @@ class WanVideoEncodeLatentBatch:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Encodes a batch of images individually to create a latent video batch where each video is a single frame, useful for I2V init purposes, for example as multiple context window inits"
|
||||
|
||||
def encode(self, vae, images, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, latent_strength=1.0):
|
||||
def encode(self, vae, images, enable_vae_tiling=False, tile_x=272, tile_y=272, tile_stride_x=144, tile_stride_y=128, latent_strength=1.0):
|
||||
vae.to(device)
|
||||
|
||||
images = images.clone()
|
||||
@@ -2123,7 +2195,7 @@ class WanVideoEncodeLatentBatch:
|
||||
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
|
||||
else:
|
||||
latent = vae.encode(img.unsqueeze(0).unsqueeze(0).permute(0, 4, 1, 2, 3), device=device, tiled=enable_vae_tiling)
|
||||
|
||||
|
||||
if latent_strength != 1.0:
|
||||
latent *= latent_strength
|
||||
latent_list.append(latent.squeeze(0).cpu())
|
||||
@@ -2183,14 +2255,16 @@ class WanVideoEncode:
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
latents = vae.encode(image * 2.0 - 1.0, device=device, tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))
|
||||
|
||||
|
||||
vae.to(offload_device)
|
||||
if latent_strength != 1.0:
|
||||
latents *= latent_strength
|
||||
|
||||
latents = latents.cpu()
|
||||
|
||||
log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}")
|
||||
mm.soft_empty_cache()
|
||||
|
||||
|
||||
return ({"samples": latents, "noise_mask": mask},)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -2205,7 +2279,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo,
|
||||
"WanVideoContextOptions": WanVideoContextOptions,
|
||||
"WanVideoTextEmbedBridge": WanVideoTextEmbedBridge,
|
||||
"WanVideoFlowEdit": WanVideoFlowEdit,
|
||||
"WanVideoControlEmbeds": WanVideoControlEmbeds,
|
||||
"WanVideoSLG": WanVideoSLG,
|
||||
"WanVideoLoopArgs": WanVideoLoopArgs,
|
||||
@@ -2221,7 +2294,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoBlockList": WanVideoBlockList,
|
||||
"WanVideoTextEncodeCached": WanVideoTextEncodeCached,
|
||||
"WanVideoAddExtraLatent": WanVideoAddExtraLatent,
|
||||
"WanVideoScheduler": WanVideoScheduler,
|
||||
"WanVideoAddStandInLatent": WanVideoAddStandInLatent,
|
||||
"WanVideoAddControlEmbeds": WanVideoAddControlEmbeds,
|
||||
"WanVideoAddMTVMotion": WanVideoAddMTVMotion,
|
||||
@@ -2229,8 +2301,12 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddPusaNoise": WanVideoAddPusaNoise,
|
||||
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
||||
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
|
||||
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
|
||||
"WanVideoAddVideoPromptEmbeds": WanVideoAddVideoPromptEmbeds,
|
||||
"WanVideoAddBindweaveEmbeds": WanVideoAddBindweaveEmbeds,
|
||||
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
|
||||
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
|
||||
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
|
||||
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
|
||||
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -2246,7 +2322,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video",
|
||||
"WanVideoContextOptions": "WanVideo Context Options",
|
||||
"WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge",
|
||||
"WanVideoFlowEdit": "WanVideo FlowEdit",
|
||||
"WanVideoControlEmbeds": "WanVideo Control Embeds",
|
||||
"WanVideoSLG": "WanVideo SLG",
|
||||
"WanVideoLoopArgs": "WanVideo Loop Args",
|
||||
@@ -2269,5 +2344,9 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddPusaNoise": "WanVideo Add Pusa Noise",
|
||||
"WanVideoAnimateEmbeds": "WanVideo Animate Embeds",
|
||||
"WanVideoAddLucyEditLatents": "WanVideo Add LucyEdit Latents",
|
||||
"WanVideoSchedulerSA_ODE": "WanVideo Scheduler SA-ODE",
|
||||
"WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds",
|
||||
"WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds",
|
||||
"WanVideoAddTTMLatents": "WanVideo Add TTMLatents",
|
||||
"WanVideoAddStoryMemLatents": "WanVideo Add StoryMem Latents",
|
||||
"WanVideoSVIProEmbeds": "WanVideo SVIPro Embeds",
|
||||
}
|
||||
|
||||
+353
-135
File diff suppressed because it is too large
Load Diff
+723
-1004
File diff suppressed because it is too large
Load Diff
+190
-69
@@ -1,6 +1,9 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
from comfy.utils import common_upscale
|
||||
from comfy import model_management
|
||||
from tqdm import tqdm
|
||||
from .utils import log
|
||||
from einops import rearrange
|
||||
|
||||
@@ -12,6 +15,9 @@ except:
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
PATCH_SIZE = (1, 2, 2)
|
||||
|
||||
main_device = model_management.get_torch_device()
|
||||
offload_device = model_management.unet_offload_device()
|
||||
|
||||
class WanVideoImageResizeToClosest:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -30,7 +36,7 @@ class WanVideoImageResizeToClosest:
|
||||
DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code"
|
||||
|
||||
def process(self, image, generation_width, generation_height, aspect_ratio_preservation ):
|
||||
|
||||
|
||||
H, W = image.shape[1], image.shape[2]
|
||||
max_area = generation_width * generation_height
|
||||
|
||||
@@ -42,7 +48,7 @@ class WanVideoImageResizeToClosest:
|
||||
aspect_ratio = generation_height / generation_width
|
||||
if aspect_ratio_preservation == "crop_to_new":
|
||||
crop = "center"
|
||||
|
||||
|
||||
lat_h = round(
|
||||
np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] //
|
||||
PATCH_SIZE[1] * PATCH_SIZE[1])
|
||||
@@ -130,27 +136,27 @@ class WanVideoVACEStartToEndFrame:
|
||||
# Convert negative end_index to positive
|
||||
if end_index < 0:
|
||||
end_index = num_frames + end_index
|
||||
|
||||
|
||||
# Create output batch with empty frames
|
||||
out_batch = torch.ones((num_frames, H, W, 3), device=device) * empty_frame_level
|
||||
|
||||
|
||||
# Create mask tensor with proper dimensions
|
||||
masks = torch.ones((num_frames, H, W), device=device)
|
||||
|
||||
|
||||
# Pre-process all images at once to avoid redundant work
|
||||
if end_image is not None and (end_image.shape[1] != H or end_image.shape[2] != W):
|
||||
end_image = common_upscale(end_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
|
||||
|
||||
|
||||
if control_images is not None and (control_images.shape[1] != H or control_images.shape[2] != W):
|
||||
control_images = common_upscale(control_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1)
|
||||
|
||||
|
||||
# Place start image at start_index
|
||||
if start_image is not None:
|
||||
frames_to_copy = min(start_image.shape[0], num_frames - start_index)
|
||||
if frames_to_copy > 0:
|
||||
out_batch[start_index:start_index + frames_to_copy] = start_image[:frames_to_copy]
|
||||
masks[start_index:start_index + frames_to_copy] = 0
|
||||
|
||||
|
||||
# Place end image at end_index
|
||||
if end_image is not None:
|
||||
# Calculate where to start placing end images
|
||||
@@ -158,28 +164,28 @@ class WanVideoVACEStartToEndFrame:
|
||||
if end_start < 0: # Handle case where end images won't all fit
|
||||
end_image = end_image[abs(end_start):]
|
||||
end_start = 0
|
||||
|
||||
|
||||
frames_to_copy = min(end_image.shape[0], num_frames - end_start)
|
||||
if frames_to_copy > 0:
|
||||
out_batch[end_start:end_start + frames_to_copy] = end_image[:frames_to_copy]
|
||||
masks[end_start:end_start + frames_to_copy] = 0
|
||||
|
||||
|
||||
# Apply control images to remaining frames that don't have start or end images
|
||||
if control_images is not None:
|
||||
# Create a mask of frames that are still empty (mask == 1)
|
||||
empty_frames = masks.sum(dim=(1, 2)) > 0.5 * H * W
|
||||
|
||||
|
||||
if empty_frames.any():
|
||||
# Only apply control images where they exist
|
||||
control_length = control_images.shape[0]
|
||||
for frame_idx in range(num_frames):
|
||||
if empty_frames[frame_idx] and frame_idx < control_length:
|
||||
out_batch[frame_idx] = control_images[frame_idx]
|
||||
|
||||
|
||||
# Apply inpaint mask if provided
|
||||
if inpaint_mask is not None:
|
||||
inpaint_mask = common_upscale(inpaint_mask.unsqueeze(1), W, H, "nearest-exact", "disabled").squeeze(1).to(device)
|
||||
|
||||
|
||||
# Handle different mask lengths efficiently
|
||||
if inpaint_mask.shape[0] > num_frames:
|
||||
inpaint_mask = inpaint_mask[:num_frames]
|
||||
@@ -221,31 +227,31 @@ class CreateCFGScheduleFloatList:
|
||||
cfg_list = [1.0] * 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(cfg_scale_start + factor * (cfg_scale_end - cfg_scale_start), 2)
|
||||
|
||||
|
||||
# If start_percent > 0, always include the first step
|
||||
if start_percent > 0:
|
||||
cfg_list[0] = 1.0
|
||||
|
||||
if unique_id and PromptServer is not None:
|
||||
try:
|
||||
try:
|
||||
PromptServer.instance.send_progress_text(
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
@@ -254,7 +260,7 @@ class CreateCFGScheduleFloatList:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
|
||||
|
||||
class CreateScheduleFloatList:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -284,16 +290,16 @@ class CreateScheduleFloatList:
|
||||
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":
|
||||
@@ -308,7 +314,7 @@ class CreateScheduleFloatList:
|
||||
cfg_list[0] = default_value
|
||||
|
||||
if unique_id and PromptServer is not None:
|
||||
try:
|
||||
try:
|
||||
PromptServer.instance.send_progress_text(
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
@@ -317,7 +323,7 @@ class CreateScheduleFloatList:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
|
||||
|
||||
|
||||
class DummyComfyWanModelObject:
|
||||
@classmethod
|
||||
@@ -343,7 +349,7 @@ class DummyComfyWanModelObject:
|
||||
return model_sampling
|
||||
return None
|
||||
return (DummyModel(),)
|
||||
|
||||
|
||||
class WanVideoLatentReScale:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -400,7 +406,7 @@ class WanVideoLatentReScale:
|
||||
samples["samples"] = latents
|
||||
|
||||
return (samples,)
|
||||
|
||||
|
||||
class WanVideoSigmaToStep:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -417,7 +423,7 @@ class WanVideoSigmaToStep:
|
||||
|
||||
def convert(self, sigma):
|
||||
return (sigma,)
|
||||
|
||||
|
||||
class NormalizeAudioLoudness:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -432,11 +438,11 @@ class NormalizeAudioLoudness:
|
||||
FUNCTION = "normalize"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def normalize(self, audio, lufs):
|
||||
def normalize(self, audio, lufs):
|
||||
audio_input = audio["waveform"]
|
||||
sample_rate = audio["sample_rate"]
|
||||
if audio_input.dim() == 3:
|
||||
audio_input = audio_input.squeeze(0)
|
||||
audio_input = audio_input.squeeze(0)
|
||||
audio_input_np = audio_input.detach().transpose(0, 1).numpy().astype(np.float32)
|
||||
audio_input_np = np.ascontiguousarray(audio_input_np)
|
||||
normalized_audio = self.loudness_norm(audio_input_np, sr=sample_rate, lufs=lufs)
|
||||
@@ -444,7 +450,7 @@ class NormalizeAudioLoudness:
|
||||
out_audio = {"waveform": torch.from_numpy(normalized_audio).transpose(0, 1).unsqueeze(0).float(), "sample_rate": sample_rate}
|
||||
|
||||
return (out_audio, )
|
||||
|
||||
|
||||
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
|
||||
try:
|
||||
import pyloudnorm
|
||||
@@ -456,7 +462,7 @@ class NormalizeAudioLoudness:
|
||||
return audio_array
|
||||
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
|
||||
return normalized_audio
|
||||
|
||||
|
||||
class WanVideoPassImagesFromSamples:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -500,15 +506,15 @@ class FaceMaskFromPoseKeypoints:
|
||||
for i, pose_frame in enumerate(pose_frames):
|
||||
selected_idx, prev_center = self.select_closest_person(pose_frame, person_index if i == 0 else prev_center)
|
||||
np_frames.append(self.draw_kps(pose_frame, selected_idx))
|
||||
|
||||
|
||||
if not np_frames:
|
||||
# Handle case where no frames were processed
|
||||
log.warning("No valid pose frames found, returning empty mask")
|
||||
return (torch.zeros((1, 64, 64), dtype=torch.float32),)
|
||||
|
||||
|
||||
np_frames = np.stack(np_frames, axis=0)
|
||||
tensor = torch.from_numpy(np_frames).float() / 255.
|
||||
print("tensor.shape:", tensor.shape)
|
||||
log.info(f"tensor.shape: {tensor.shape}")
|
||||
tensor = tensor[:, :, :, 0]
|
||||
return (tensor,)
|
||||
|
||||
@@ -516,41 +522,41 @@ class FaceMaskFromPoseKeypoints:
|
||||
people = pose_frame["people"]
|
||||
if not people:
|
||||
return -1, None
|
||||
|
||||
|
||||
centers = []
|
||||
valid_people_indices = []
|
||||
|
||||
|
||||
for idx, person in enumerate(people):
|
||||
# Check if face keypoints exist and are valid
|
||||
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
|
||||
continue
|
||||
|
||||
|
||||
kps = np.array(person["face_keypoints_2d"])
|
||||
if len(kps) == 0:
|
||||
continue
|
||||
|
||||
|
||||
n = len(kps) // 3
|
||||
if n == 0:
|
||||
continue
|
||||
|
||||
|
||||
facial_kps = rearrange(kps, "(n c) -> n c", n=n, c=3)[:, :2]
|
||||
|
||||
|
||||
# Check if we have valid coordinates (not all zeros)
|
||||
if np.all(facial_kps == 0):
|
||||
continue
|
||||
|
||||
|
||||
center = facial_kps.mean(axis=0)
|
||||
|
||||
|
||||
# Check if center is valid (not NaN or infinite)
|
||||
if np.isnan(center).any() or np.isinf(center).any():
|
||||
continue
|
||||
|
||||
|
||||
centers.append(center)
|
||||
valid_people_indices.append(idx)
|
||||
|
||||
|
||||
if not centers:
|
||||
return -1, None
|
||||
|
||||
|
||||
if isinstance(prev_center_or_index, (int, np.integer)):
|
||||
# First frame: use person_index, but map to valid people
|
||||
if 0 <= prev_center_or_index < len(valid_people_indices):
|
||||
@@ -582,58 +588,58 @@ class FaceMaskFromPoseKeypoints:
|
||||
width, height = pose_frame["canvas_width"], pose_frame["canvas_height"]
|
||||
canvas = np.zeros((height, width, 3), dtype=np.uint8)
|
||||
people = pose_frame["people"]
|
||||
|
||||
|
||||
if person_index < 0 or person_index >= len(people):
|
||||
return canvas # Out of bounds, return blank
|
||||
|
||||
|
||||
person = people[person_index]
|
||||
|
||||
|
||||
# Check if face keypoints exist and are valid
|
||||
if "face_keypoints_2d" not in person or not person["face_keypoints_2d"]:
|
||||
return canvas # No face keypoints, return blank
|
||||
|
||||
|
||||
face_kps_data = person["face_keypoints_2d"]
|
||||
if len(face_kps_data) == 0:
|
||||
return canvas # Empty keypoints, return blank
|
||||
|
||||
|
||||
n = len(face_kps_data) // 3
|
||||
if n < 17: # Need at least 17 points for outer contour
|
||||
return canvas # Not enough keypoints, return blank
|
||||
|
||||
|
||||
facial_kps = rearrange(np.array(face_kps_data), "(n c) -> n c", n=n, c=3)[:, :2]
|
||||
|
||||
|
||||
# Check if we have valid coordinates (not all zeros)
|
||||
if np.all(facial_kps == 0):
|
||||
return canvas # All keypoints are zero, return blank
|
||||
|
||||
|
||||
# Check for NaN or infinite values
|
||||
if np.isnan(facial_kps).any() or np.isinf(facial_kps).any():
|
||||
return canvas # Invalid coordinates, return blank
|
||||
|
||||
|
||||
# Check for negative coordinates or coordinates that would create streaks
|
||||
if np.any(facial_kps < 0):
|
||||
return canvas # Negative coordinates, likely bad detection
|
||||
|
||||
|
||||
# Check if coordinates are reasonable (not too close to edges which might indicate bad detection)
|
||||
min_margin = 5 # Minimum distance from edges
|
||||
if (np.any(facial_kps[:, 0] < min_margin) or
|
||||
np.any(facial_kps[:, 1] < min_margin) or
|
||||
np.any(facial_kps[:, 0] > width - min_margin) or
|
||||
if (np.any(facial_kps[:, 0] < min_margin) or
|
||||
np.any(facial_kps[:, 1] < min_margin) or
|
||||
np.any(facial_kps[:, 0] > width - min_margin) or
|
||||
np.any(facial_kps[:, 1] > height - min_margin)):
|
||||
# Check if this looks like a streak to corner (many points near 0,0)
|
||||
corner_points = np.sum((facial_kps[:, 0] < min_margin) & (facial_kps[:, 1] < min_margin))
|
||||
if corner_points > 3: # Too many points near corner, likely bad detection
|
||||
return canvas
|
||||
|
||||
|
||||
facial_kps = facial_kps.astype(np.int32)
|
||||
|
||||
|
||||
# Ensure coordinates are within canvas bounds
|
||||
facial_kps[:, 0] = np.clip(facial_kps[:, 0], 0, width - 1)
|
||||
facial_kps[:, 1] = np.clip(facial_kps[:, 1], 0, height - 1)
|
||||
|
||||
|
||||
part_color = (255, 255, 255)
|
||||
outer_contour = facial_kps[:17]
|
||||
|
||||
|
||||
# Additional validation for the contour before drawing
|
||||
# Check if contour points are too spread out (indicating bad detection)
|
||||
if len(outer_contour) >= 3:
|
||||
@@ -642,11 +648,11 @@ class FaceMaskFromPoseKeypoints:
|
||||
max_x, max_y = np.max(outer_contour, axis=0)
|
||||
contour_width = max_x - min_x
|
||||
contour_height = max_y - min_y
|
||||
|
||||
|
||||
# If contour spans more than 80% of canvas, likely bad detection
|
||||
if (contour_width > 0.8 * width or contour_height > 0.8 * height):
|
||||
return canvas
|
||||
|
||||
|
||||
# Check if we have a valid contour (at least 3 unique points)
|
||||
unique_points = np.unique(outer_contour, axis=0)
|
||||
if len(unique_points) >= 3:
|
||||
@@ -654,13 +660,124 @@ class FaceMaskFromPoseKeypoints:
|
||||
# Calculate area to see if it's too large or too small
|
||||
contour_area = cv2.contourArea(outer_contour)
|
||||
canvas_area = width * height
|
||||
|
||||
|
||||
# If contour is less than 0.1% or more than 50% of canvas, skip
|
||||
if 0.001 * canvas_area <= contour_area <= 0.5 * canvas_area:
|
||||
cv2.fillPoly(canvas, pts=[outer_contour], color=part_color)
|
||||
|
||||
|
||||
return canvas
|
||||
|
||||
|
||||
|
||||
class DrawGaussianNoiseOnImage:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"image": ("IMAGE", ),
|
||||
"mask": ("MASK", ),
|
||||
},
|
||||
"optional": {
|
||||
"device": (["cpu", "gpu"], {"default": "cpu", "tooltip": "Device to use for processing"}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
RETURN_NAMES = ("images",)
|
||||
FUNCTION = "apply"
|
||||
CATEGORY = "KJNodes/masking"
|
||||
DESCRIPTION = "Fills the background (masked area) with Gaussian noise sampled using the mean and variance of the subject (unmasked) region."
|
||||
|
||||
def apply(self, image, mask, device="cpu", seed=0):
|
||||
B, H, W, C = image.shape
|
||||
BM, HM, WM = mask.shape
|
||||
|
||||
processing_device = main_device if device == "gpu" else torch.device("cpu")
|
||||
|
||||
in_masks = mask.clone().to(processing_device)
|
||||
in_images = image.clone().to(processing_device)
|
||||
|
||||
# Resize mask to match image dimensions
|
||||
if HM != H or WM != W:
|
||||
in_masks = F.interpolate(mask.unsqueeze(1), size=(H, W), mode='nearest-exact').squeeze(1)
|
||||
|
||||
# Match batch sizes
|
||||
if B > BM:
|
||||
in_masks = in_masks.repeat((B + BM - 1) // BM, 1, 1)[:B]
|
||||
elif BM > B:
|
||||
in_masks = in_masks[:B]
|
||||
|
||||
output_images = []
|
||||
|
||||
# Set random seed for reproducibility
|
||||
generator = torch.Generator(device=processing_device).manual_seed(seed)
|
||||
|
||||
for i in tqdm(range(B), desc="DrawGaussianNoiseOnImage batch"):
|
||||
curr_mask = in_masks[i]
|
||||
img_idx = min(i, B - 1)
|
||||
curr_image = in_images[img_idx]
|
||||
|
||||
# Expand mask to 3 channels
|
||||
mask_expanded = curr_mask.unsqueeze(-1).expand(-1, -1, 3)
|
||||
|
||||
# Calculate mean and std per channel from the subject region (where mask is 1)
|
||||
subject_mask = mask_expanded > 0.5
|
||||
|
||||
# Initialize noise tensor
|
||||
noise = torch.zeros_like(curr_image)
|
||||
|
||||
for c in range(C):
|
||||
channel = curr_image[:, :, c]
|
||||
channel_mask = subject_mask[:, :, c]
|
||||
|
||||
if channel_mask.sum() > 0:
|
||||
# Get subject pixels
|
||||
subject_pixels = channel[channel_mask]
|
||||
|
||||
# Calculate statistics
|
||||
mean = subject_pixels.mean()
|
||||
std = subject_pixels.std()
|
||||
|
||||
# Generate Gaussian noise for this channel
|
||||
noise[:, :, c] = torch.normal(mean=mean.item(), std=std.item(),
|
||||
size=(H, W), generator=generator,
|
||||
device=processing_device)
|
||||
|
||||
# Clamp noise to valid range
|
||||
noise = torch.clamp(noise, 0.0, 1.0)
|
||||
|
||||
# Apply: keep subject, fill background with noise
|
||||
masked_image = curr_image * mask_expanded + noise * (1 - mask_expanded)
|
||||
output_images.append(masked_image)
|
||||
|
||||
# If no masks were processed, return empty tensor
|
||||
if not output_images:
|
||||
return (torch.zeros((0, H, W, 3), dtype=image.dtype),)
|
||||
|
||||
out_rgb = torch.stack(output_images, dim=0).cpu()
|
||||
|
||||
return (out_rgb, )
|
||||
|
||||
|
||||
class WanVideoPreviewEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT", "MASK")
|
||||
RETURN_NAMES = ("image_embeds", "mask",)
|
||||
FUNCTION = "get"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def get(self, embeds):
|
||||
latents = embeds.get("image_embeds", None)
|
||||
mask = embeds.get("mask", None)
|
||||
if mask is not None:
|
||||
mask = mask[0].float().cpu()
|
||||
return ({"samples": latents.unsqueeze(0)}, mask)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
|
||||
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
|
||||
@@ -673,6 +790,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"NormalizeAudioLoudness": NormalizeAudioLoudness,
|
||||
"WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples,
|
||||
"FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints,
|
||||
"DrawGaussianNoiseOnImage": DrawGaussianNoiseOnImage,
|
||||
"WanVideoPreviewEmbeds": WanVideoPreviewEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
|
||||
@@ -686,4 +805,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"NormalizeAudioLoudness": "Normalize Audio Loudness",
|
||||
"WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples",
|
||||
"FaceMaskFromPoseKeypoints": "Face Mask From Pose Keypoints",
|
||||
}
|
||||
"DrawGaussianNoiseOnImage": "Draw Gaussian Noise On Image",
|
||||
"WanVideoPreviewEmbeds": "WanVideo Preview Embeds",
|
||||
}
|
||||
|
||||
@@ -0,0 +1,440 @@
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
from einops import rearrange
|
||||
import numpy as np
|
||||
from typing import Tuple
|
||||
|
||||
from .unet_causal_3d_blocks import get_down_block3d, CausalConv3d
|
||||
|
||||
class ControlNetCausalConditioningEmbedding(nn.Module):
|
||||
def __init__(self, conditioning_embedding_channels: int, conditioning_channels: int = 3, block_out_channels: Tuple[int, ...] = (16, 32, 96, 256)):
|
||||
super().__init__()
|
||||
self.conv_in = CausalConv3d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
|
||||
self.blocks = nn.ModuleList([])
|
||||
|
||||
for i in range(len(block_out_channels) - 1):
|
||||
channel_in = block_out_channels[i]
|
||||
channel_out = block_out_channels[i + 1]
|
||||
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
|
||||
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
|
||||
|
||||
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
|
||||
|
||||
def forward(self, conditioning):
|
||||
embedding = self.conv_in(conditioning)
|
||||
embedding = F.silu(embedding)
|
||||
|
||||
for block in self.blocks:
|
||||
embedding = block(embedding)
|
||||
embedding = F.silu(embedding)
|
||||
|
||||
embedding = self.conv_out(embedding)
|
||||
|
||||
return embedding
|
||||
|
||||
class MiniHunyuanEncoder(nn.Module):
|
||||
'''
|
||||
a direct copy of hunyuan encoder
|
||||
'''
|
||||
def __init__(
|
||||
self,
|
||||
in_channels = 3,
|
||||
out_channels = 3,
|
||||
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
|
||||
block_out_channels = [128, 256, 512, 512],
|
||||
layers_per_block = 2,
|
||||
norm_num_groups = 32,
|
||||
act_fn: str = "silu",
|
||||
time_compression_ratio: int = 4,
|
||||
spatial_compression_ratio: int = 8,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
self.conv_in = CausalConv3d(
|
||||
in_channels, block_out_channels[0], kernel_size=3, stride=1)
|
||||
self.mid_block = None
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
# down
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
num_spatial_downsample_layers = int(
|
||||
np.log2(spatial_compression_ratio))
|
||||
num_time_downsample_layers = int(np.log2(time_compression_ratio))
|
||||
|
||||
if time_compression_ratio == 4:
|
||||
add_spatial_downsample = bool(
|
||||
i < num_spatial_downsample_layers)
|
||||
add_time_downsample = bool(i >= (
|
||||
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
|
||||
elif time_compression_ratio == 8:
|
||||
add_spatial_downsample = bool(
|
||||
i < num_spatial_downsample_layers)
|
||||
add_time_downsample = bool(i < num_time_downsample_layers)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported time_compression_ratio: {time_compression_ratio}")
|
||||
|
||||
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
|
||||
downsample_stride_T = (2, ) if add_time_downsample else (1, )
|
||||
downsample_stride = tuple(
|
||||
downsample_stride_T + downsample_stride_HW)
|
||||
down_block = get_down_block3d(
|
||||
down_block_type,
|
||||
num_layers=self.layers_per_block,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
add_downsample=bool(
|
||||
add_spatial_downsample or add_time_downsample),
|
||||
downsample_stride=downsample_stride,
|
||||
resnet_eps=1e-6,
|
||||
downsample_padding=0,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
|
||||
|
||||
def forward(self, sample):
|
||||
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
|
||||
sample = self.conv_in(sample)
|
||||
# down
|
||||
for down_block in self.down_blocks:
|
||||
sample = down_block(sample)
|
||||
sample = self.conv_out(sample)
|
||||
return sample
|
||||
|
||||
|
||||
class ControlNetConditioningEmbedding(nn.Module):
|
||||
"""
|
||||
Quoting from https://arxiv.org/abs/2302.05543: "Stable Diffusion uses a pre-processing method similar to VQ-GAN
|
||||
[11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized
|
||||
training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the
|
||||
convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides
|
||||
(activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full
|
||||
model) to encode image-space conditions ... into feature maps ..."
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
conditioning_embedding_channels: int,
|
||||
conditioning_channels: int = 3,
|
||||
block_out_channels: Tuple[int, ...] = (16, 32, 96, 256),
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
|
||||
|
||||
self.blocks = nn.ModuleList([])
|
||||
|
||||
for i in range(len(block_out_channels) - 1):
|
||||
channel_in = block_out_channels[i]
|
||||
channel_out = block_out_channels[i + 1]
|
||||
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
|
||||
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
|
||||
|
||||
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
|
||||
|
||||
def forward(self, conditioning):
|
||||
embedding = self.conv_in(conditioning)
|
||||
embedding = F.silu(embedding)
|
||||
|
||||
for block in self.blocks:
|
||||
embedding = block(embedding)
|
||||
embedding = F.silu(embedding)
|
||||
|
||||
embedding = self.conv_out(embedding)
|
||||
|
||||
return embedding
|
||||
|
||||
|
||||
class InflatedGroupNorm(nn.GroupNorm):
|
||||
def forward(self, x):
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
class InflatedConv3d(nn.Conv2d):
|
||||
def forward(self, x):
|
||||
video_length = x.shape[2]
|
||||
|
||||
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||
x = super().forward(x)
|
||||
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class ResnetBlockInflated(nn.Module):
|
||||
def __init__(self, *, in_channels, out_channels=None, dropout=0.0, groups=32, groups_out=None, pre_norm=True, eps=1e-6, non_linearity="swish", output_scale_factor=1.0):
|
||||
super().__init__()
|
||||
self.pre_norm = pre_norm
|
||||
self.pre_norm = True
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.output_scale_factor = output_scale_factor
|
||||
|
||||
if groups_out is None:
|
||||
groups_out = groups
|
||||
|
||||
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
if non_linearity == "swish":
|
||||
self.nonlinearity = lambda x: F.silu(x)
|
||||
elif non_linearity == "silu":
|
||||
self.nonlinearity = nn.SiLU()
|
||||
|
||||
def forward(self, input_tensor, temb):
|
||||
if temb is not None:
|
||||
print("Warning: temb is None in ResnetBlockInflated")
|
||||
hidden_states = input_tensor
|
||||
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
|
||||
hidden_states = self.conv1(hidden_states)
|
||||
|
||||
if temb is not None:
|
||||
hidden_states = hidden_states + temb
|
||||
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.dropout(hidden_states)
|
||||
hidden_states = self.conv2(hidden_states)
|
||||
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
return output_tensor
|
||||
|
||||
class DownEncoderBlockInflated(nn.Module):
|
||||
def __init__(self, *, num_layers: int, in_channels: int, out_channels: int, add_downsample: bool, downsample_stride: tuple = (1, 2, 2),
|
||||
resnet_eps: float = 1e-6, resnet_act_fn: str = "silu", resnet_groups: int = 32):
|
||||
super().__init__()
|
||||
|
||||
self.resnets = nn.ModuleList([ResnetBlockInflated(
|
||||
in_channels=in_channels if i == 0 else out_channels,
|
||||
out_channels=out_channels,
|
||||
eps=resnet_eps,
|
||||
non_linearity=resnet_act_fn,
|
||||
groups=resnet_groups,
|
||||
) for i in range(num_layers)])
|
||||
|
||||
self.downsamplers = nn.ModuleList()
|
||||
if add_downsample:
|
||||
self.downsamplers.append(
|
||||
InflatedConv3d(
|
||||
out_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=2,
|
||||
padding=1,
|
||||
)
|
||||
)
|
||||
self.down_stride = downsample_stride
|
||||
else:
|
||||
self.down_stride = (1, 1, 1)
|
||||
|
||||
def forward(self, x, temb=None):
|
||||
for resnet in self.resnets:
|
||||
x = resnet(x, temb)
|
||||
|
||||
for down in self.downsamplers:
|
||||
x = down(x)
|
||||
return x
|
||||
|
||||
|
||||
class SFT(nn.Module): # 2D SFT
|
||||
def __init__(
|
||||
self, in_channels, out_channels, intermediate_channels=128, groups=32, eps=1e-6):
|
||||
super().__init__()
|
||||
self.out_channels = out_channels
|
||||
self.norm = InflatedGroupNorm(groups, out_channels, eps, affine=True)
|
||||
self.mlp_shared = nn.Sequential(InflatedConv3d(in_channels, intermediate_channels, kernel_size=3, stride=1, padding=1), nn.SiLU())
|
||||
self.mlp_gamma = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
self.mlp_beta = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
def forward(self, hidden_state, condition):
|
||||
"""
|
||||
hidden_state : (B, Cout, T, H, W)
|
||||
condition : (B, Cin, 1, H, W)
|
||||
"""
|
||||
hidden_state = self.norm(hidden_state) #2D SFT 2D Norm
|
||||
|
||||
actv = self.mlp_shared(condition)
|
||||
gamma = self.mlp_gamma(actv)
|
||||
beta = self.mlp_beta(actv)
|
||||
|
||||
return torch.addcmul(beta, hidden_state, 1 + gamma)
|
||||
|
||||
|
||||
class MiniEncoder2D(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
down_block_types: list = (
|
||||
"DownEncoderBlockInflated",
|
||||
"DownEncoderBlockInflated",
|
||||
"DownEncoderBlockInflated",
|
||||
"DownEncoderBlockInflated",
|
||||
),
|
||||
block_out_channels: list = (128, 256, 512, 512),
|
||||
layers_per_block: int = 2,
|
||||
norm_num_groups: int = 32,
|
||||
act_fn: str = "silu",
|
||||
spatial_compression_ratio: int = 8,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# -------------------------------------------------------------------
|
||||
# conv in
|
||||
# -------------------------------------------------------------------
|
||||
self.conv_in = InflatedConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
|
||||
|
||||
self.down_blocks = nn.ModuleList()
|
||||
output_channel = block_out_channels[0]
|
||||
num_spatial_down_layers = int(np.log2(spatial_compression_ratio))
|
||||
|
||||
for i, block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
# is_final_block = i == len(block_out_channels) - 1
|
||||
|
||||
add_spatial_downsample = bool(i < num_spatial_down_layers)
|
||||
|
||||
downsample_stride = (1, 2, 2) if add_spatial_downsample else (1, 1, 1)
|
||||
|
||||
down_block = DownEncoderBlockInflated(
|
||||
num_layers=layers_per_block,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
add_downsample=add_spatial_downsample,
|
||||
downsample_stride=downsample_stride,
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
self.conv_out = InflatedConv3d(output_channel, out_channels, kernel_size=3, stride=1, padding=1)
|
||||
|
||||
def forward(self, x):
|
||||
# (B,C,1,H,W)
|
||||
x = self.conv_in(x)
|
||||
|
||||
for block in self.down_blocks:
|
||||
x = block(x)
|
||||
|
||||
return self.conv_out(x)
|
||||
|
||||
|
||||
class Driven_Ref_PoseEncoder(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels = 3, out_channels = 3,
|
||||
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
|
||||
block_out_channels = [128, 256, 512, 512], layers_per_block = 2, norm_num_groups = 32,
|
||||
act_fn: str = "silu", time_compression_ratio: int = 4, spatial_compression_ratio: int = 8,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
|
||||
self.mid_block = None
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
# down
|
||||
output_channel = block_out_channels[0]
|
||||
for i, down_block_type in enumerate(down_block_types):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
num_spatial_downsample_layers = int(
|
||||
np.log2(spatial_compression_ratio))
|
||||
num_time_downsample_layers = int(np.log2(time_compression_ratio))
|
||||
|
||||
if time_compression_ratio == 4:
|
||||
add_spatial_downsample = bool(
|
||||
i < num_spatial_downsample_layers)
|
||||
add_time_downsample = bool(i >= (
|
||||
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
|
||||
elif time_compression_ratio == 8:
|
||||
add_spatial_downsample = bool(
|
||||
i < num_spatial_downsample_layers)
|
||||
add_time_downsample = bool(i < num_time_downsample_layers)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported time_compression_ratio: {time_compression_ratio}")
|
||||
|
||||
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
|
||||
downsample_stride_T = (2, ) if add_time_downsample else (1, )
|
||||
downsample_stride = tuple(
|
||||
downsample_stride_T + downsample_stride_HW)
|
||||
down_block = get_down_block3d(
|
||||
down_block_type,
|
||||
num_layers=self.layers_per_block,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
add_downsample=bool(
|
||||
add_spatial_downsample or add_time_downsample),
|
||||
downsample_stride=downsample_stride,
|
||||
resnet_eps=1e-6,
|
||||
downsample_padding=0,
|
||||
resnet_act_fn=act_fn,
|
||||
resnet_groups=norm_num_groups,
|
||||
attention_head_dim=output_channel,
|
||||
)
|
||||
self.down_blocks.append(down_block)
|
||||
|
||||
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
|
||||
|
||||
self.ref_pose_encoder = MiniEncoder2D(
|
||||
in_channels = in_channels,
|
||||
out_channels = out_channels,
|
||||
block_out_channels = block_out_channels,
|
||||
norm_num_groups = norm_num_groups,
|
||||
layers_per_block = layers_per_block,
|
||||
spatial_compression_ratio = spatial_compression_ratio,
|
||||
)
|
||||
self.sft_layers = nn.ModuleList()
|
||||
for i, ch in enumerate(block_out_channels):
|
||||
if i == 0: # 0 层 (H/2,W/2) 不做 SFT
|
||||
self.sft_layers.append(None)
|
||||
else: # H/4、H/8、H/16 做 SFT
|
||||
self.sft_layers.append(
|
||||
SFT(
|
||||
in_channels=ch,
|
||||
out_channels=ch,
|
||||
intermediate_channels=max(8, ch // 2),
|
||||
groups=norm_num_groups,
|
||||
)
|
||||
)
|
||||
|
||||
def forward(self, driven_pose, ref_pose):
|
||||
# driven_pose b c t h w
|
||||
# ref_pose b c 1 h w
|
||||
ref_pose_cond, ref_feats = self.ref_pose_encoder(ref_pose)
|
||||
|
||||
x = self.conv_in(driven_pose)
|
||||
for i, down_block in enumerate(self.down_blocks):
|
||||
x = down_block(x)
|
||||
|
||||
if self.sft_layers[i] is not None:
|
||||
cond_feat = ref_feats[i]
|
||||
x = self.sft_layers[i](x, cond_feat)
|
||||
|
||||
driven_pose_cond = self.conv_out(x)
|
||||
return driven_pose_cond, ref_pose_cond
|
||||
@@ -0,0 +1,152 @@
|
||||
import torch
|
||||
from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
class WanVideoAddOneToAllReferenceEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||
"ref_image": ("IMAGE",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||
},
|
||||
"optional": {
|
||||
"ref_mask": ("MASK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, ref_mask=None):
|
||||
updated = dict(embeds)
|
||||
|
||||
ref_latent = ref_latent_empty = None
|
||||
vae.to(device)
|
||||
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
|
||||
ref_latent = vae.encode([ref_image_in], device, tiled=False)
|
||||
ref_mask_in = None
|
||||
if ref_mask is not None:
|
||||
ref_mask_in = (ref_mask.unsqueeze(0).repeat(3, 1, 1, 1) * 2 - 1.).to(device, vae.dtype)
|
||||
else:
|
||||
ref_mask_in = torch.zeros_like(ref_image_in)-1
|
||||
ref_mask_latent = vae.encode([ref_mask_in], device, tiled=False)
|
||||
|
||||
if ref_mask is not None and not torch.all(ref_mask == 0):
|
||||
ref_latent_empty = vae.encode([torch.zeros_like(ref_image_in)-1], device, tiled=False)
|
||||
else:
|
||||
ref_latent_empty = ref_mask_latent
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
updated.setdefault("one_to_all_embeds", {})
|
||||
updated["one_to_all_embeds"]["ref_latent_pos"] = torch.cat([ref_latent, ref_latent_empty], dim=1)
|
||||
updated["one_to_all_embeds"]["ref_latent_neg"] = torch.cat([ref_latent_empty, ref_latent_empty], dim=1)
|
||||
updated["one_to_all_embeds"]["ref_strength"] = strength
|
||||
updated["one_to_all_embeds"]["ref_start_percent"] = start_percent
|
||||
updated["one_to_all_embeds"]["ref_end_percent"] = end_percent
|
||||
|
||||
return (updated,)
|
||||
|
||||
class WanVideoAddOneToAllPoseEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
|
||||
},
|
||||
"optional": {
|
||||
"pose_prefix_image": ("IMAGE",),
|
||||
"pose_cfg_scale": ("FLOAT", {"default": 1.5, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "CFG scale for the pose control, has no effect if main cfg scale is 1.0"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, pose_images, strength, pose_prefix_image=None, start_percent=0.0, end_percent=1.0, pose_cfg_scale=1.5):
|
||||
updated = dict(embeds)
|
||||
updated.setdefault("one_to_all_embeds", {})
|
||||
pose_images_in = pose_images[..., :3].unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
|
||||
updated["one_to_all_embeds"]["pose_images"] = pose_images_in
|
||||
if pose_prefix_image is not None:
|
||||
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_prefix_image.unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
|
||||
else:
|
||||
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_images_in[:, :, :1]
|
||||
|
||||
updated["one_to_all_embeds"]["controlnet_strength"] = strength
|
||||
updated["one_to_all_embeds"]["controlnet_start_percent"] = start_percent
|
||||
updated["one_to_all_embeds"]["controlnet_end_percent"] = end_percent
|
||||
updated["one_to_all_embeds"]["pose_cfg_scale"] = pose_cfg_scale
|
||||
|
||||
return (updated,)
|
||||
|
||||
class WanVideoAddOneToAllExtendEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"prev_latents": ("LATENT", {"tooltip": "Previous latents to be used to continue generation"}),
|
||||
"window_size": ("INT", {"default": 81, "min": 1, "max": 256, "step": 1, "tooltip": "Number of new frames to generate" }),
|
||||
"overlap": ("INT", {"default": 5, "min": 0, "max": 64, "step": 1, "tooltip": "Number of overlapping frames between previous and new frames" }),
|
||||
"frames_processed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Number of frames already processed in the video" }),
|
||||
"if_not_enough_frames": (["pad_with_last", "error"], {"default": "pad_with_last", "tooltip": "What to do if there are not enough frames in pose_images for the window"}),
|
||||
},
|
||||
"optional": {
|
||||
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "IMAGE",)
|
||||
RETURN_NAMES = ("image_embeds", "pose_slice",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, prev_latents, if_not_enough_frames, window_size=81, overlap=5, frames_processed=0, pose_images=None):
|
||||
updated = dict(embeds)
|
||||
updated.setdefault("one_to_all_embeds", {})
|
||||
updated["one_to_all_embeds"]["prev_latents"] = prev_latents["samples"][0]
|
||||
if pose_images is not None:
|
||||
pose_images_in = pose_images.clone()[..., :3]
|
||||
start = max(0, frames_processed - overlap)
|
||||
end = start + window_size
|
||||
log.info(f"Extracting pose images from {start} to {end}")
|
||||
if start >= pose_images_in.shape[0]:
|
||||
raise ValueError(f"start index {start} exceeds pose images length {pose_images_in.shape[0]}")
|
||||
if end > pose_images_in.shape[0]:
|
||||
if if_not_enough_frames == "pad_with_last":
|
||||
padding_needed = end - pose_images_in.shape[0]
|
||||
pose_images_in = torch.cat([pose_images_in, pose_images_in[-1:].repeat(padding_needed, 1, 1, 1)], dim=0)
|
||||
log.info(f"Not enough frames, padding with {padding_needed} frames to reach {end} total frames")
|
||||
else:
|
||||
raise ValueError(f"end index {end} exceeds pose images length {pose_images.shape[0]}")
|
||||
pose_slice = pose_images_in[start:end]
|
||||
else:
|
||||
pose_slice = torch.zeros((1, 64, 64, 3))
|
||||
|
||||
return (updated, pose_slice)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddOneToAllReferenceEmbeds": WanVideoAddOneToAllReferenceEmbeds,
|
||||
"WanVideoAddOneToAllPoseEmbeds": WanVideoAddOneToAllPoseEmbeds,
|
||||
"WanVideoAddOneToAllExtendEmbeds": WanVideoAddOneToAllExtendEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddOneToAllReferenceEmbeds": "WanVideo Add OneToAll Reference Embeds",
|
||||
"WanVideoAddOneToAllPoseEmbeds": "WanVideo Add OneToAll Pose Embeds",
|
||||
"WanVideoAddOneToAllExtendEmbeds": "WanVideo Add OneToAll Extend Embeds",
|
||||
}
|
||||
@@ -0,0 +1,220 @@
|
||||
from typing import Dict, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from ..wanvideo.modules.model import WanLayerNorm, WanSelfAttention, EmbedND_RifleX, sinusoidal_embedding_1d, apply_rotary_emb_split, apply_rope_comfy1
|
||||
|
||||
class WanAttentionBlock(nn.Module):
|
||||
def __init__(self, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default"):
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = out_features // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
self.attention_mode = attention_mode
|
||||
self.rope_func = rope_func
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(self.dim, eps)
|
||||
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function, head_norm=False)
|
||||
self.norm2 = WanLayerNorm(self.dim, eps)
|
||||
self.ffn = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
|
||||
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
|
||||
|
||||
|
||||
def get_mod(self, e, modulation):
|
||||
if e.dim() == 3:
|
||||
if e.shape[-1] == 512:
|
||||
e = self.modulation(e)
|
||||
return e.unsqueeze(2).chunk(6, dim=-1)
|
||||
return (modulation + e).chunk(6, dim=1) # 1, 6, dim
|
||||
elif e.dim() == 4:
|
||||
e_mod = modulation.unsqueeze(2) + e
|
||||
return [ei.squeeze(1) for ei in e_mod.unbind(dim=1)]
|
||||
|
||||
|
||||
def modulate(self, norm_x, shift_msa, scale_msa):
|
||||
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
|
||||
|
||||
def ffn_chunked(self, mod_x, num_chunks=4):
|
||||
seq_len = mod_x.shape[1]
|
||||
if seq_len <= 8192 or num_chunks <= 1:
|
||||
return self.ffn(mod_x)
|
||||
return torch.cat([self.ffn(chunk.contiguous()) for chunk in mod_x.chunk(num_chunks, dim=1)], dim=1)
|
||||
|
||||
#region attention forward
|
||||
def forward(self, x, e, seq_lens, freqs, split_rope=True, e_tr=None, tr_start=0, tr_num=0):
|
||||
|
||||
use_token_replace = False
|
||||
if e_tr is not None and tr_num > 0:
|
||||
tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = self.get_mod(e_tr.to(x.device), self.modulation)
|
||||
use_token_replace = True
|
||||
tr_start = tr_start or 0
|
||||
tr_end = tr_start + (tr_num or 0)
|
||||
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
|
||||
del e
|
||||
input_dtype = x.dtype
|
||||
|
||||
if use_token_replace:
|
||||
norm_x = self.norm1(x.to(shift_msa.dtype))
|
||||
input_x = torch.cat([
|
||||
torch.addcmul(shift_msa, norm_x[:, :tr_start], 1 + scale_msa), # before replace → T
|
||||
torch.addcmul(tr_shift_msa, norm_x[:, tr_start:tr_end], 1 + tr_scale_msa), # replace segment → t=0
|
||||
torch.addcmul(shift_msa, norm_x[:, tr_end:], 1 + scale_msa) # after replace → T
|
||||
], dim=1).to(input_dtype)
|
||||
else:
|
||||
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa).to(input_dtype)
|
||||
del shift_msa, scale_msa
|
||||
|
||||
b, s, n, d = *x.shape[:2], self.self_attn.num_heads, self.self_attn.head_dim
|
||||
h_dim = w_dim = 2 * (self.head_dim // 6)
|
||||
t_dim = self.head_dim - h_dim - w_dim
|
||||
|
||||
q = self.self_attn.norm_q(self.self_attn.q(input_x)).to(self.self_attn.norm_q.weight.dtype).view(b, s, n, d)
|
||||
if split_rope:
|
||||
q = apply_rotary_emb_split(q, freqs, t_dim) # Apply split rotary embedding (only to H/W dimensions, leaving T unchanged)
|
||||
else:
|
||||
q = apply_rope_comfy1(q, freqs)
|
||||
|
||||
k = self.self_attn.norm_k(self.self_attn.k(input_x).to(self.self_attn.norm_k.weight.dtype)).to(input_x.dtype).view(b, s, n, d)
|
||||
if split_rope:
|
||||
k = apply_rotary_emb_split(k, freqs, t_dim)
|
||||
else:
|
||||
k = apply_rope_comfy1(k, freqs)
|
||||
|
||||
v = self.self_attn.v(input_x).view(b, s, n, d)
|
||||
del input_x
|
||||
|
||||
y = self.self_attn.forward(q, k, v, seq_lens)
|
||||
del q, k, v
|
||||
if use_token_replace:
|
||||
x = x + torch.cat([
|
||||
y[:, :tr_start] * gate_msa,
|
||||
y[:, tr_start:tr_end] * tr_gate_msa,
|
||||
y[:, tr_end:] * gate_msa
|
||||
], dim=1).to(input_dtype)
|
||||
else:
|
||||
x = x.addcmul(y, gate_msa)
|
||||
del y, gate_msa
|
||||
|
||||
# ffn
|
||||
if use_token_replace:
|
||||
norm2_x = self.norm2(x.to(shift_mlp.dtype))
|
||||
mod_x = torch.cat([
|
||||
torch.addcmul(shift_mlp, norm2_x[:, :tr_start], 1 + scale_mlp),
|
||||
torch.addcmul(tr_shift_mlp, norm2_x[:, tr_start:tr_end], 1 + tr_scale_mlp),
|
||||
torch.addcmul(shift_mlp, norm2_x[:, tr_end:], 1 + scale_mlp)
|
||||
], dim=1)
|
||||
else:
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
|
||||
del shift_mlp, scale_mlp
|
||||
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
|
||||
del mod_x
|
||||
|
||||
# gate_mlp
|
||||
if use_token_replace:
|
||||
x = x + torch.cat([
|
||||
x_ffn[:, :tr_start] * gate_mlp,
|
||||
x_ffn[:, tr_start:tr_end] * tr_gate_mlp,
|
||||
x_ffn[:, tr_end:] * gate_mlp
|
||||
], dim=1).to(input_dtype)
|
||||
else:
|
||||
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
|
||||
del gate_mlp
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class WanRefextractor(nn.Module):
|
||||
def __init__(self, patch_size=(1, 2, 2), in_dim=16, dim=5120, in_features=5120, out_features=5120, ffn_dim=8192, ffn2_dim=8192,
|
||||
freq_dim=256, num_heads=16, num_layers=32, eps=1e-6,
|
||||
qk_norm=True, cross_attn_norm=True,
|
||||
attention_mode='sdpa', rope_func='comfy', rms_norm_function='default',
|
||||
main_device=torch.device('cuda'), offload_device=torch.device('cpu'), dtype=torch.float16):
|
||||
super().__init__()
|
||||
self.patch_size = patch_size
|
||||
self.freq_dim = freq_dim
|
||||
self.dim = dim
|
||||
self.main_device = main_device
|
||||
self.base_dtype = dtype
|
||||
self.attention_mode = attention_mode
|
||||
self.patch_embedding = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
self.time_embedding = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
||||
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
|
||||
self.blocks = nn.ModuleList([
|
||||
WanAttentionBlock(in_features, out_features, ffn_dim, ffn2_dim, num_heads,
|
||||
qk_norm, cross_attn_norm, eps, attention_mode="sdpa", rope_func=rope_func, rms_norm_function=rms_norm_function)
|
||||
for i in range(num_layers)
|
||||
])
|
||||
|
||||
self.ref_blocks = nn.ModuleList([])
|
||||
for _ in range(len(self.blocks)+1):
|
||||
self.ref_blocks.append(nn.Linear(in_features, out_features))
|
||||
|
||||
d = dim // num_heads
|
||||
self.rope_embedder = EmbedND_RifleX(d,10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=1, k=0)
|
||||
|
||||
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
|
||||
patch_size = self.patch_size
|
||||
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
||||
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
|
||||
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
|
||||
|
||||
if steps_t is None:
|
||||
steps_t = t_len
|
||||
if steps_h is None:
|
||||
steps_h = h_len
|
||||
if steps_w is None:
|
||||
steps_w = w_len
|
||||
|
||||
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
|
||||
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
|
||||
|
||||
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
|
||||
return freqs
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
B, C, F, H, W = x.shape
|
||||
|
||||
freqs = self.rope_encode_comfy(F, H, W, device=x.device, dtype=x.dtype)
|
||||
|
||||
self.patch_embedding.to(self.main_device)
|
||||
x = self.patch_embedding(x.float()).to(x.dtype).flatten(2).transpose(1, 2).to(self.base_dtype)
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
|
||||
|
||||
time_embed_dtype = self.time_embedding[0].weight.dtype
|
||||
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
||||
time_embed_dtype = self.base_dtype
|
||||
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep.flatten()).to(time_embed_dtype)) # b, dim
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim)).to(self.base_dtype) # b, 6, dim
|
||||
del e
|
||||
|
||||
# 4. Transformer blocks
|
||||
block_samples = ()
|
||||
for block in self.blocks:
|
||||
block_samples = block_samples + (x, )
|
||||
x = block(x, e0, seq_lens, freqs)
|
||||
|
||||
block_samples = block_samples + (x, )
|
||||
|
||||
ref_block_samples = ()
|
||||
for block_sample, ref_block in zip(block_samples, self.ref_blocks):
|
||||
block_sample = ref_block(block_sample)
|
||||
ref_block_samples = ref_block_samples + (block_sample, )
|
||||
|
||||
return ref_block_samples, freqs
|
||||
@@ -0,0 +1,144 @@
|
||||
# Copyright 2024 The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
#
|
||||
# Modified from diffusers==0.29.2
|
||||
#
|
||||
# ==============================================================================
|
||||
|
||||
from typing import Optional
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
import comfy.ops
|
||||
ops = comfy.ops.disable_weight_init
|
||||
|
||||
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
|
||||
seq_len = n_frame * n_hw
|
||||
mask = torch.full((seq_len, seq_len), float(
|
||||
"-inf"), dtype=dtype, device=device)
|
||||
for i in range(seq_len):
|
||||
i_frame = i // n_hw
|
||||
mask[i, : (i_frame + 1) * n_hw] = 0
|
||||
if batch_size is not None:
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
return mask
|
||||
|
||||
|
||||
class CausalConv3d(nn.Module):
|
||||
def __init__(self, chan_in, chan_out, kernel_size, stride = 1, dilation = 1, pad_mode='replicate', **kwargs):
|
||||
super().__init__()
|
||||
self.pad_mode = pad_mode
|
||||
padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0) # W, H, T
|
||||
self.time_causal_padding = padding
|
||||
self.conv = ops.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class DownsampleCausal3D(nn.Module):
|
||||
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv", kernel_size=3, bias=True, stride=2):
|
||||
super().__init__()
|
||||
self.channels, self.out_channels, self.use_conv, self.padding, self.name = channels, out_channels or channels, use_conv, padding, name
|
||||
self.conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, bias=bias)
|
||||
|
||||
def forward(self, x, scale=1.0):
|
||||
return self.conv(x)
|
||||
|
||||
|
||||
class ResnetBlockCausal3D(nn.Module):
|
||||
def __init__(self, *, in_channels: int, out_channels: Optional[int] = None, groups: int = 32, eps: float = 1e-6, conv_3d_out_channels: Optional[int] = None):
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
out_channels = in_channels if out_channels is None else out_channels
|
||||
self.out_channels = out_channels
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
self.norm2 = torch.nn.GroupNorm(num_groups=groups, num_channels=out_channels, eps=eps, affine=True)
|
||||
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
|
||||
conv_3d_out_channels = conv_3d_out_channels or out_channels
|
||||
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
|
||||
|
||||
def forward(self, input_tensor: torch.FloatTensor, temb: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
|
||||
hidden_states = input_tensor
|
||||
hidden_states = self.conv1(nn.SiLU()(self.norm1(hidden_states)))
|
||||
if temb is not None:
|
||||
hidden_states = hidden_states + temb
|
||||
hidden_states = self.conv2(nn.SiLU()(self.norm2(hidden_states)))
|
||||
return input_tensor + hidden_states
|
||||
|
||||
|
||||
def get_down_block3d(down_block_type: str, num_layers: int, in_channels: int, out_channels: int,
|
||||
add_downsample: bool, downsample_stride: int, resnet_eps: float, resnet_act_fn: str, resnet_groups: Optional[int] = None,
|
||||
downsample_padding: Optional[int] = None, **kwargs):
|
||||
|
||||
down_block_type = down_block_type[7:] if down_block_type.startswith(
|
||||
"UNetRes") else down_block_type
|
||||
if down_block_type == "DownEncoderBlockCausal3D":
|
||||
return DownEncoderBlockCausal3D(
|
||||
num_layers=num_layers,
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
add_downsample=add_downsample,
|
||||
downsample_stride=downsample_stride,
|
||||
resnet_eps=resnet_eps,
|
||||
resnet_act_fn=resnet_act_fn,
|
||||
resnet_groups=resnet_groups,
|
||||
downsample_padding=downsample_padding,
|
||||
)
|
||||
raise ValueError(f"{down_block_type} does not exist.")
|
||||
|
||||
|
||||
class DownEncoderBlockCausal3D(nn.Module):
|
||||
def __init__(self, in_channels: int, out_channels: int, num_layers: int = 1, resnet_eps: float = 1e-6,
|
||||
resnet_groups: int = 32, add_downsample: bool = True, downsample_stride: int = 2, downsample_padding: int = 1, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
resnets = []
|
||||
for i in range(num_layers):
|
||||
in_channels = in_channels if i == 0 else out_channels
|
||||
resnets.append(
|
||||
ResnetBlockCausal3D(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
eps=resnet_eps,
|
||||
groups=resnet_groups,
|
||||
)
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if add_downsample:
|
||||
self.downsamplers = nn.ModuleList([DownsampleCausal3D(
|
||||
out_channels,
|
||||
use_conv=True,
|
||||
out_channels=out_channels,
|
||||
padding=downsample_padding,
|
||||
name="op",
|
||||
stride=downsample_stride,
|
||||
)])
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states, temb=None, scale=scale)
|
||||
|
||||
if self.downsamplers is not None:
|
||||
for downsampler in self.downsamplers:
|
||||
hidden_states = downsampler(hidden_states, scale)
|
||||
|
||||
return hidden_states
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "ComfyUI-WanVideoWrapper"
|
||||
description = "ComfyUI wrapper nodes for WanVideo"
|
||||
version = "1.3.8"
|
||||
version = "1.4.5"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
|
||||
|
||||
|
||||
@@ -1,7 +1,35 @@
|
||||
## Note: Due to the stupid amount of bots or people thinking this is some of video generation service, I've blocked new accounts from posting issues for now.
|
||||
|
||||
# ComfyUI wrapper nodes for [WanVideo](https://github.com/Wan-Video/Wan2.1) and related models.
|
||||
|
||||
|
||||
## Memory use update (again)
|
||||
|
||||
I've made everythign less reliant on torch.compile for VRAM efficiency, so things should work better even without it. Also figured workaround for some issues when using compile that made first run use drastically more VRAM, issue I battled with myself a lot.
|
||||
|
||||
|
||||
## Update notification that can affect memory use in old workflows
|
||||
|
||||
In a recent update I changed how unmerged LoRA weights are handled:
|
||||
|
||||
Previously mostly due to my laziness they were always loaded from RAM when used, this was of course inefficient and also made using torch.compile for LoRA applying difficult, thus forcing a graph break when using unmerged LoRAs.
|
||||
|
||||
Now the LoRA weights are assigned as buffers to the corresponding modules, so they are part of the blocks and obey the block swapping unifying the offloading and allowing LoRA weights to benefit from the prefetch feature for async offoading. Downside is that this means if you did not use block swap, you will see increased memory use as the LoRAs are part of the model and all on VRAM.
|
||||
|
||||
If you use block swap, the LoRAs are swapped along the rest of the block, but the block size is now larger, this means you may have to compensate with couple of more blocks swapped.
|
||||
|
||||
Example situation: you use 1GB LoRA unmerged and swap 20 blocks on 14B model, we can divide the LoRA size by block count, single block grows by 25MB, 20 blocks grow by 500MB, so your VRAM usage would be 500MB more than before, to compensate you swap 2 more blocks.
|
||||
|
||||
### Unrelated other VRAM issue with torch.compile
|
||||
|
||||
After any update that modifies the model code and when using torch.compile it's common to run into issues with VRAM, this can be caused by using older pytorch/triton version without latest compile fixes, and/or from old triton caches, mostly in Windows. This manifests in the issue that first run of new input size may have drastically increased memory use, which can clear from simply running it again, and once cached, not manifest again. Again I've only seen this happen in Windows.
|
||||
|
||||
To clear your Triton cache you can delete the contents of following (default) folders:
|
||||
|
||||
`C:\Users\<username>\.triton`
|
||||
`C:\Users\<username>\AppData\Local\Temp\torchinductor_<username>`
|
||||
|
||||
|
||||
## Note: Due to the stupid amount of bots or people thinking this is some of video generation service, I've blocked new accounts from posting issues for now.
|
||||
|
||||
# WORK IN PROGRESS (perpetually)
|
||||
|
||||
# Why should I use custom nodes when WanVideo works natively?
|
||||
@@ -76,6 +104,22 @@ WanAnimate: https://github.com/Wan-Video/Wan2.2/tree/main/wan/modules/animate
|
||||
|
||||
Lynx: https://github.com/bytedance/lynx
|
||||
|
||||
MoCha: https://github.com/Orange-3DV-Team/MoCha
|
||||
|
||||
UniLumos: https://github.com/alibaba-damo-academy/Lumos-Custom
|
||||
|
||||
Bindweave: https://github.com/bytedance/BindWeave
|
||||
|
||||
Training free techniques:
|
||||
|
||||
TimeToMove: https://github.com/time-to-move/TTM
|
||||
|
||||
SteadyDancer: https://github.com/MCG-NJU/SteadyDancer
|
||||
|
||||
One-to-all-Animation: https://github.com/ssj9596/One-to-All-Animation
|
||||
|
||||
SCAIL: https://github.com/zai-org/SCAIL
|
||||
|
||||
|
||||
Not exactly Wan model, but close enough to work with the code base:
|
||||
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
# Modify from https://github.com/liyunsheng13/dcd/blob/main/models/imagenet/mobilenetv2_dcd.py
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class Hsigmoid(nn.Module):
|
||||
def __init__(self, inplace=True):
|
||||
super(Hsigmoid, self).__init__()
|
||||
self.inplace = inplace
|
||||
|
||||
def forward(self, x):
|
||||
return F.relu6(x + 3., inplace=self.inplace) / 3.
|
||||
|
||||
|
||||
class DYModule(nn.Module):
|
||||
def __init__(self, inp, oup, fc_squeeze=8):
|
||||
super(DYModule, self).__init__()
|
||||
self.conv = nn.Conv2d(inp, oup, 1, 1, 0, bias=False)
|
||||
if inp < oup:
|
||||
self.mul = 4
|
||||
reduction = 8
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(2)
|
||||
else:
|
||||
self.mul = 1
|
||||
reduction = 2
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
|
||||
self.dim = min((inp * self.mul) // reduction, oup // reduction)
|
||||
while self.dim ** 2 > inp * self.mul * 2:
|
||||
reduction *= 2
|
||||
self.dim = min((inp * self.mul) // reduction, oup // reduction)
|
||||
if self.dim < 4:
|
||||
self.dim = 4
|
||||
|
||||
squeeze = max(inp * self.mul, self.dim ** 2) // fc_squeeze
|
||||
if squeeze < 4:
|
||||
squeeze = 4
|
||||
self.conv_q = nn.Conv2d(inp, self.dim, 1, 1, 0, bias=False)
|
||||
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(inp * self.mul, squeeze, bias=False),
|
||||
SEModule_small(squeeze),
|
||||
)
|
||||
self.fc_phi = nn.Linear(squeeze, self.dim ** 2, bias=False)
|
||||
self.fc_scale = nn.Linear(squeeze, oup, bias=False)
|
||||
self.hs = Hsigmoid()
|
||||
self.conv_p = nn.Conv2d(self.dim, oup, 1, 1, 0, bias=False)
|
||||
# self.bn1 = nn.BatchNorm2d(self.dim)
|
||||
self.bn1 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
|
||||
# self.bn2 = nn.BatchNorm1d(self.dim)
|
||||
self.bn2 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
|
||||
|
||||
def forward(self, x):
|
||||
x_type = x.dtype
|
||||
r = self.conv(x.to(self.conv.weight.dtype)).to(x_type)
|
||||
b, c, h, w = x.size()
|
||||
|
||||
y = self.avg_pool(x).view(b, c * self.mul)
|
||||
y = self.fc(y)
|
||||
dy_phi = self.fc_phi(y).view(b, self.dim, self.dim)
|
||||
dy_scale = self.hs(self.fc_scale(y)).view(b, -1, 1, 1)
|
||||
r = dy_scale.expand_as(r) * r
|
||||
|
||||
x = self.conv_q(x.to(self.conv_q.weight.dtype)).to(self.bn1.weight.dtype)
|
||||
x = self.bn1(x)
|
||||
|
||||
x = x.view(b, -1, h * w)
|
||||
x = x + self.bn2(torch.matmul(dy_phi, x.to(dy_phi.dtype)).to(self.bn2.weight.dtype))
|
||||
x = x.view(b, -1, h, w)
|
||||
|
||||
x = self.conv_p(x.to(self.conv_p.weight.dtype)).to(x_type)
|
||||
|
||||
return x + r
|
||||
|
||||
|
||||
class SEModule_small(nn.Module):
|
||||
def __init__(self, channel):
|
||||
super(SEModule_small, self).__init__()
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(channel, channel, bias=False),
|
||||
Hsigmoid()
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
y = self.fc(x)
|
||||
return x * y
|
||||
|
||||
|
||||
class SEModule(nn.Module):
|
||||
def __init__(self, channel, reduction=4):
|
||||
super(SEModule, self).__init__()
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(channel, channel // reduction, bias=False),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(channel // reduction, channel, bias=False),
|
||||
Hsigmoid()
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
b, c, _, _ = x.size()
|
||||
y = self.avg_pool(x).view(b, c)
|
||||
y = self.fc(y).view(b, c, 1, 1)
|
||||
return x * y.expand_as(x)
|
||||
@@ -0,0 +1,62 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from ..utils import log
|
||||
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import load_torch_file, ProgressBar
|
||||
import folder_paths
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
|
||||
class WanVideoAddSteadyDancerEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"pose_latents_positive": ("LATENT",),
|
||||
"pose_strength_spatial": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the pose embedding"}),
|
||||
"pose_strength_temporal": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the pose embedding"}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||
},
|
||||
"optional": {
|
||||
"pose_latents_negative": ("LATENT",),
|
||||
"clip_vision_embeds": ("WANVIDIMAGE_CLIPEMBEDS",),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, pose_latents_positive, pose_strength_spatial, pose_strength_temporal, start_percent=0.0, end_percent=1.0, pose_latents_negative=None, clip_vision_embeds=None):
|
||||
sdancer_embeds = {
|
||||
"cond_pos": pose_latents_positive["samples"][0],
|
||||
"cond_neg": pose_latents_negative["samples"][0] if pose_latents_negative else None,
|
||||
"pose_strength_spatial": pose_strength_spatial,
|
||||
"pose_strength_temporal": pose_strength_temporal,
|
||||
"start_percent": start_percent,
|
||||
"end_percent": end_percent,
|
||||
"clip_fea": clip_vision_embeds,
|
||||
}
|
||||
|
||||
updated = dict(embeds)
|
||||
updated["sdancer_embeds"] = sdancer_embeds
|
||||
return (updated,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddSteadyDancerEmbeds": WanVideoAddSteadyDancerEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddSteadyDancerEmbeds": "WanVideo Add SteadyDancer Embeds",
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class FactorConv3d(nn.Module):
|
||||
"""
|
||||
(2+1)D decomposition of 3D convolution: 1xHxW spatial convolution → Swish → Tx1x1 temporal convolution
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
kernel_size,
|
||||
stride: int = 1,
|
||||
dilation: int = 1):
|
||||
super().__init__()
|
||||
|
||||
if isinstance(kernel_size, int):
|
||||
k_t, k_h, k_w = kernel_size, kernel_size, kernel_size
|
||||
else:
|
||||
k_t, k_h, k_w = kernel_size
|
||||
|
||||
pad_t = (k_t - 1) * dilation // 2
|
||||
pad_hw = (k_h - 1) * dilation // 2
|
||||
|
||||
self.spatial = nn.Conv3d(
|
||||
in_channels, in_channels,
|
||||
kernel_size=(1, k_h, k_w),
|
||||
stride=(1, stride, stride),
|
||||
padding=(0, pad_hw, pad_hw),
|
||||
dilation=(1, dilation, dilation),
|
||||
groups=in_channels,
|
||||
bias=False
|
||||
)
|
||||
|
||||
self.temporal = nn.Conv3d(
|
||||
in_channels, out_channels,
|
||||
kernel_size=(k_t, 1, 1),
|
||||
stride=(stride, 1, 1),
|
||||
padding=(pad_t, 0, 0),
|
||||
dilation=(dilation, 1, 1),
|
||||
bias=True
|
||||
)
|
||||
|
||||
self.act = nn.SiLU()
|
||||
|
||||
def forward(self, x):
|
||||
out_dtype = x.dtype
|
||||
x = self.spatial(x.to(self.spatial.weight.dtype)).to(out_dtype)
|
||||
x = self.act(x)
|
||||
return self.temporal(x.to(self.temporal.weight.dtype)).to(out_dtype)
|
||||
|
||||
class LayerNorm2D(nn.Module):
|
||||
"""
|
||||
LayerNorm over C for a 4-D tensor (B, C, H, W)
|
||||
"""
|
||||
def __init__(self, num_channels, eps=1e-5, affine=True):
|
||||
super().__init__()
|
||||
self.num_channels = num_channels
|
||||
self.eps = eps
|
||||
self.affine = affine
|
||||
if affine:
|
||||
self.weight = nn.Parameter(torch.ones(1, num_channels, 1, 1))
|
||||
self.bias = nn.Parameter(torch.zeros(1, num_channels, 1, 1))
|
||||
|
||||
def forward(self, x):
|
||||
# x: (B, C, H, W)
|
||||
mean = x.mean(dim=1, keepdim=True) # (B, 1, H, W)
|
||||
var = x.var (dim=1, keepdim=True, unbiased=False)
|
||||
x = (x - mean) / torch.sqrt(var + self.eps)
|
||||
if self.affine:
|
||||
x = x * self.weight + self.bias
|
||||
return x
|
||||
|
||||
|
||||
class PoseRefNetNoBNV3(nn.Module):
|
||||
def __init__(self,
|
||||
in_channels_c: int,
|
||||
in_channels_x: int,
|
||||
hidden_dim: int = 256,
|
||||
num_heads: int = 8,
|
||||
dropout: float = 0.1):
|
||||
super().__init__()
|
||||
self.d_model = hidden_dim
|
||||
self.nhead = num_heads
|
||||
|
||||
self.proj_p = nn.Conv2d(in_channels_c, hidden_dim, kernel_size=1)
|
||||
self.proj_r = nn.Conv2d(in_channels_x, hidden_dim, kernel_size=1)
|
||||
|
||||
self.proj_p_back = nn.Conv2d(hidden_dim, in_channels_c, kernel_size=1)
|
||||
|
||||
self.cross_attn = nn.MultiheadAttention(hidden_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout)
|
||||
|
||||
self.ffn_pose = nn.Sequential(
|
||||
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1)
|
||||
)
|
||||
|
||||
self.norm1 = LayerNorm2D(hidden_dim)
|
||||
self.norm2 = LayerNorm2D(hidden_dim)
|
||||
|
||||
def forward(self, pose, ref, mask=None):
|
||||
"""
|
||||
pose : (B, C1, T, H, W)
|
||||
ref : (B, C2, T, H, W)
|
||||
mask : (B, T*H*W) optional key_padding_mask
|
||||
return: (B, d_model, T, H, W)
|
||||
"""
|
||||
B, _, T, H, W = pose.shape
|
||||
|
||||
p_trans = pose.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
|
||||
r_trans = ref.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
|
||||
|
||||
p_trans = self.proj_p(p_trans.to(self.proj_p.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
|
||||
r_trans = self.proj_r(r_trans.to(self.proj_r.weight.dtype)).to(self.cross_attn.in_proj_weight.dtype).flatten(2).transpose(1, 2)
|
||||
|
||||
out = self.cross_attn(query=r_trans,
|
||||
key=p_trans,
|
||||
value=p_trans,
|
||||
key_padding_mask=mask)[0]
|
||||
|
||||
out = self.norm1(out.transpose(1, 2).contiguous().view(B*T, -1, H, W))
|
||||
|
||||
out_type = out.dtype
|
||||
|
||||
out = out + self.ffn_pose(out.to(self.ffn_pose[0].weight.dtype)).to(out_type)
|
||||
out = self.norm2(out)
|
||||
|
||||
out = self.proj_p_back(out.to(self.proj_p_back.weight.dtype)).to(out_type)
|
||||
|
||||
return out.view(B, T, -1, H, W).contiguous().transpose(1, 2)
|
||||
+9
-2
@@ -8,6 +8,7 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from tqdm.auto import tqdm
|
||||
from collections import namedtuple
|
||||
from ..wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
|
||||
|
||||
DecoderResult = namedtuple("DecoderResult", ("frame", "memory"))
|
||||
TWorkItem = namedtuple("TWorkItem", ("input_tensor", "block_index"))
|
||||
@@ -146,7 +147,7 @@ def apply_model_with_memblocks(model, x, parallel, show_progress_bar):
|
||||
return x
|
||||
|
||||
class TAEHV(nn.Module):
|
||||
def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16):
|
||||
def __init__(self, state_dict, parallel=False, decoder_time_upscale=(True, True), decoder_space_upscale=(True, True, True), dtype=torch.float16, model_name="taehv"):
|
||||
"""Initialize pretrained TAEHV from the given checkpoint.
|
||||
|
||||
Arg:
|
||||
@@ -161,6 +162,7 @@ class TAEHV(nn.Module):
|
||||
if self.latent_channels == 48:
|
||||
self.patch_size = 2
|
||||
self.dtype = dtype
|
||||
self.model_name = model_name
|
||||
|
||||
self.encoder = nn.Sequential(
|
||||
conv(self.image_channels*self.patch_size**2, 64), nn.ReLU(inplace=True),
|
||||
@@ -180,8 +182,11 @@ class TAEHV(nn.Module):
|
||||
)
|
||||
if state_dict is not None:
|
||||
self.load_state_dict(self.patch_tgrow_layers(state_dict))
|
||||
|
||||
|
||||
self.parallel = parallel
|
||||
orig_vae = WanVideoVAE38() if self.latent_channels == 48 else WanVideoVAE()
|
||||
self.mean = orig_vae.mean.to(dtype).movedim(1, 2)
|
||||
self.inv_std = orig_vae.inv_std.to(dtype).movedim(1, 2)
|
||||
|
||||
def patch_tgrow_layers(self, sd):
|
||||
"""Patch TGrow layers to use a smaller kernel if needed.
|
||||
@@ -221,6 +226,8 @@ class TAEHV(nn.Module):
|
||||
if False, frames will be processed sequentially.
|
||||
Returns NTCHW RGB tensor with ~[0, 1] values.
|
||||
"""
|
||||
if "light" in self.model_name.lower():
|
||||
x = x / self.inv_std.to(x) + self.mean.to(x)
|
||||
x = apply_model_with_memblocks(self.decoder, x, self.parallel, show_progress_bar)
|
||||
if self.patch_size > 1: x = F.pixel_shuffle(x, self.patch_size)
|
||||
return x[:, self.frames_to_trim:]
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/attn_qk_int8_per_block.py
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
|
||||
K_ptrs, K_scale_ptr, V_ptrs, stride_kn, stride_vn,
|
||||
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
start_m,
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr, offs_m: tl.constexpr, offs_n: tl.constexpr,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
frame_tokens: tl.constexpr = 1560,
|
||||
sigmoid_a: tl.constexpr = 1.0,
|
||||
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
|
||||
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
|
||||
sink_width: tl.constexpr = 4,
|
||||
window_width: tl.constexpr = 16,
|
||||
multi_factor: tl.constexpr = None,
|
||||
entropy_factor: tl.constexpr = None,
|
||||
):
|
||||
|
||||
|
||||
lo, hi = 0, kv_len
|
||||
for start_n in range(lo, hi, BLOCK_N):
|
||||
start_n = tl.multiple_of(start_n, BLOCK_N)
|
||||
k_mask = offs_n[None, :] < (kv_len - start_n)
|
||||
k = tl.load(K_ptrs, mask = k_mask)
|
||||
k_scale = tl.load(K_scale_ptr)
|
||||
|
||||
|
||||
m = offs_m[:, None]
|
||||
n = start_n + offs_n
|
||||
|
||||
qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale
|
||||
|
||||
window_th = frame_tokens * window_width / 2
|
||||
dist2 = tl.abs(m - n).to(tl.int32)
|
||||
dist_mask = dist2 <= window_th
|
||||
|
||||
negative_mask = (qk<0)
|
||||
|
||||
qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor)
|
||||
|
||||
window3 = (m <= frame_tokens) & (n > window_width*frame_tokens)
|
||||
qk = tl.where(window3, -1e4, qk)
|
||||
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1))
|
||||
qk = qk - m_ij[:, None]
|
||||
p = tl.math.exp2(qk)
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptrs, mask = offs_n[:, None] < (kv_len - start_n))
|
||||
p = p.to(tl.float16)
|
||||
|
||||
acc += tl.dot(p, v, out_dtype=tl.float16)
|
||||
m_i = m_ij
|
||||
K_ptrs += BLOCK_N * stride_kn
|
||||
K_scale_ptr += 1
|
||||
V_ptrs += BLOCK_N * stride_vn
|
||||
return acc, l_i
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd(Q, K, V, Q_scale, K_scale, Out,
|
||||
Block_bias, Decay_mask,
|
||||
flags, stride_f_b, stride_f_h,
|
||||
stride_qz, stride_qh, stride_qn,
|
||||
stride_kz, stride_kh, stride_kn,
|
||||
stride_vz, stride_vh, stride_vn,
|
||||
stride_oz, stride_oh, stride_on,
|
||||
stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
qo_len, kv_len, H: tl.constexpr, num_kv_groups: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
BLOCK_M: tl.constexpr,
|
||||
BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
frame_tokens: tl.constexpr = 1560,
|
||||
sigmoid_a: tl.constexpr = 1.0,
|
||||
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
|
||||
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
|
||||
sink_width: tl.constexpr = 4,
|
||||
window_width: tl.constexpr = 16,
|
||||
multi_factor: tl.constexpr = None,
|
||||
entropy_factor: tl.constexpr = None,
|
||||
):
|
||||
start_m = tl.program_id(0)
|
||||
|
||||
off_z = tl.program_id(2).to(tl.int64)
|
||||
off_h = tl.program_id(1).to(tl.int64)
|
||||
|
||||
q_scale_offset = (off_z * H + off_h) * tl.cdiv(qo_len, BLOCK_M)
|
||||
k_scale_offset = (off_z * (H // num_kv_groups) + off_h // num_kv_groups) * tl.cdiv(kv_len, BLOCK_N)
|
||||
|
||||
flag_ptr = flags + off_z * stride_f_b + off_h * stride_f_h
|
||||
current_flag = tl.load(flag_ptr)
|
||||
|
||||
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
Q_ptrs = Q + (off_z * stride_qz + off_h * stride_qh) + offs_m[:, None] * stride_qn + offs_k[None, :]
|
||||
Q_scale_ptr = Q_scale + q_scale_offset + start_m
|
||||
K_ptrs = K + (off_z * stride_kz + (off_h // num_kv_groups) * stride_kh) + offs_n[None, :] * stride_kn + offs_k[:, None]
|
||||
K_scale_ptr = K_scale + k_scale_offset
|
||||
V_ptrs = V + (off_z * stride_vz + (off_h // num_kv_groups) * stride_vh) + offs_n[:, None] * stride_vn + offs_k[None, :]
|
||||
O_block_ptr = Out + (off_z * stride_oz + off_h * stride_oh) + offs_m[:, None] * stride_on + offs_k[None, :]
|
||||
|
||||
# # 计算block_bias指针
|
||||
Block_bias_ptrs = Block_bias + off_z * stride_bbz + off_h * stride_bbh
|
||||
|
||||
# 计算decay_mask指针
|
||||
Decay_mask_ptrs = Decay_mask + off_z * stride_dmz + off_h * stride_dmh
|
||||
|
||||
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
q = tl.load(Q_ptrs, mask = offs_m[:, None] < qo_len)
|
||||
q_scale = tl.load(Q_scale_ptr)
|
||||
acc, l_i = _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag, K_ptrs, K_scale_ptr, V_ptrs,
|
||||
stride_kn, stride_vn,
|
||||
Block_bias_ptrs, stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
Decay_mask_ptrs, stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
start_m,
|
||||
BLOCK_M, HEAD_DIM, BLOCK_N,
|
||||
4 - STAGE, offs_m, offs_n,
|
||||
xpos_xi=xpos_xi,
|
||||
frame_tokens=frame_tokens,
|
||||
sigmoid_a=sigmoid_a,
|
||||
alpha_xpos_xi=alpha_xpos_xi,
|
||||
beta_xpos_xi=beta_xpos_xi,
|
||||
sink_width=sink_width,
|
||||
window_width=window_width,
|
||||
multi_factor=multi_factor,
|
||||
entropy_factor=entropy_factor,
|
||||
)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(O_block_ptr, acc.to(Out.type.element_ty), mask = (offs_m[:, None] < qo_len))
|
||||
|
||||
def forward(q, k, v, flags, block_bias, decay_mask, q_scale, k_scale, tensor_layout="HND", output_dtype=torch.float16,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
frame_tokens: tl.constexpr = 1560,
|
||||
sigmoid_a: tl.constexpr = 1.0,
|
||||
alpha_xpos_xi: tl.constexpr = 0.9999967941742395,
|
||||
beta_xpos_xi: tl.constexpr = 0.9999860536252945,
|
||||
BLOCK_M: tl.constexpr = 128,
|
||||
BLOCK_N: tl.constexpr = 128,
|
||||
sink_width: tl.constexpr = 4,
|
||||
window_width: tl.constexpr = 16,
|
||||
multi_factor: tl.constexpr = None,
|
||||
entropy_factor: tl.constexpr = None,
|
||||
):
|
||||
stage = 1
|
||||
|
||||
o = torch.empty(q.shape, dtype=output_dtype, device=q.device)
|
||||
|
||||
b, h_qo, qo_len, head_dim = q.shape
|
||||
if block_bias is None:
|
||||
block_bias = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.float16, device=q.device)
|
||||
|
||||
if decay_mask is None:
|
||||
decay_mask = torch.zeros((b, h_qo, (qo_len + BLOCK_M - 1) // BLOCK_M, (qo_len + BLOCK_N - 1) // BLOCK_N), dtype=torch.bool, device=q.device)
|
||||
|
||||
if tensor_layout == "HND":
|
||||
b, h_qo, qo_len, head_dim = q.shape
|
||||
_, h_kv, kv_len, _ = k.shape
|
||||
|
||||
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
|
||||
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
|
||||
stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(1), v.stride(2)
|
||||
stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(1), o.stride(2)
|
||||
stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride()
|
||||
stride_dmz, stride_dmh, stride_dm, stride_dn = decay_mask.stride()
|
||||
# elif tensor_layout == "NHD":
|
||||
# b, qo_len, h_qo, head_dim = q.shape
|
||||
# _, kv_len, h_kv, _ = k.shape
|
||||
|
||||
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
|
||||
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
|
||||
# stride_bz_v, stride_h_v, stride_seq_v = v.stride(0), v.stride(2), v.stride(1)
|
||||
# stride_bz_o, stride_h_o, stride_seq_o = o.stride(0), o.stride(2), o.stride(1)
|
||||
# stride_bbz, stride_bbh, stride_bm, stride_bn = block_bias.stride(0), block_bias.stride(2), block_bias.stride(1), block_bias.stride(3)
|
||||
else:
|
||||
raise ValueError(f"tensor_layout {tensor_layout} not supported")
|
||||
|
||||
stride_f_b, stride_f_h = flags.stride()
|
||||
|
||||
HEAD_DIM_K = head_dim
|
||||
num_kv_groups = h_qo // h_kv
|
||||
|
||||
grid = (triton.cdiv(qo_len, BLOCK_M), h_qo, b)
|
||||
_attn_fwd[grid](
|
||||
q, k, v, q_scale, k_scale, o,
|
||||
block_bias, decay_mask,
|
||||
flags,
|
||||
stride_f_b, stride_f_h,
|
||||
stride_bz_q, stride_h_q, stride_seq_q,
|
||||
stride_bz_k, stride_h_k, stride_seq_k,
|
||||
stride_bz_v, stride_h_v, stride_seq_v,
|
||||
stride_bz_o, stride_h_o, stride_seq_o,
|
||||
stride_bbz, stride_bbh, stride_bm, stride_bn,
|
||||
stride_dmz, stride_dmh, stride_dm, stride_dn,
|
||||
qo_len, kv_len,
|
||||
h_qo, num_kv_groups,
|
||||
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
|
||||
STAGE=stage,
|
||||
num_warps=4 if head_dim == 64 else 8,
|
||||
num_stages=3 if head_dim == 64 else 4,
|
||||
xpos_xi=xpos_xi,
|
||||
frame_tokens=frame_tokens,
|
||||
sigmoid_a=sigmoid_a,
|
||||
alpha_xpos_xi=alpha_xpos_xi,
|
||||
beta_xpos_xi=beta_xpos_xi,
|
||||
sink_width=sink_width,
|
||||
window_width=window_width,
|
||||
multi_factor=multi_factor,
|
||||
entropy_factor=entropy_factor,
|
||||
)
|
||||
return o
|
||||
@@ -0,0 +1,64 @@
|
||||
# source https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/core.py
|
||||
|
||||
import torch
|
||||
import triton.language as tl
|
||||
|
||||
from .quant_per_block import per_block_int8
|
||||
from .attn_qk_int8_per_block import forward as attn_false
|
||||
|
||||
from typing import Optional
|
||||
|
||||
def sage_attention(
|
||||
qkv: list[torch.Tensor],
|
||||
tensor_layout: str ="HND",
|
||||
is_causal=False,
|
||||
sm_scale: Optional[float] = None,
|
||||
smooth_k: bool =True,
|
||||
xpos_xi: tl.constexpr = 0.9999934149894527,
|
||||
flags = None,
|
||||
block_bias = None,
|
||||
sigmoid_a: float = 1.0,
|
||||
alpha_xpos_xi: float = 0.97,
|
||||
beta_xpos_xi: float = 0.8,
|
||||
decay_mask = None,
|
||||
sink_width: int = 4,
|
||||
window_width: int = 21,
|
||||
multi_factor: Optional[float] = None,
|
||||
entropy_factor: Optional[float] = None,
|
||||
block_size : int = 64,
|
||||
**kwargs
|
||||
) -> torch.Tensor:
|
||||
dtype = qkv[0].dtype
|
||||
q, k, v = qkv[0].transpose(1, 2), qkv[1].transpose(1, 2), qkv[2].transpose(1, 2) # to HND
|
||||
|
||||
if flags == None:
|
||||
flags = torch.zeros([q.shape[0],q.shape[1]], dtype=torch.int32, device=q.device)
|
||||
|
||||
seq_dim = 2
|
||||
|
||||
if smooth_k:
|
||||
km = k.mean(dim=seq_dim, keepdim=True)
|
||||
k -= km
|
||||
else:
|
||||
km = None
|
||||
|
||||
if dtype == torch.bfloat16 or dtype == torch.float32:
|
||||
v = v.to(torch.float16)
|
||||
|
||||
if q.dtype != k.dtype or q.dtype != v.dtype:
|
||||
k, v = k.to(q.dtype), v.to(q.dtype)
|
||||
|
||||
q_int8, q_scale, k_int8, k_scale = per_block_int8(q, k, sm_scale=sm_scale, tensor_layout=tensor_layout, BLKQ=block_size, BLKK=block_size)
|
||||
del q, k
|
||||
|
||||
o = attn_false(q_int8, k_int8, v, flags, block_bias, decay_mask, q_scale, k_scale,
|
||||
tensor_layout=tensor_layout, output_dtype=dtype, xpos_xi=xpos_xi, sigmoid_a=sigmoid_a,
|
||||
alpha_xpos_xi=alpha_xpos_xi, beta_xpos_xi=beta_xpos_xi,
|
||||
BLOCK_M=block_size, BLOCK_N=block_size,
|
||||
sink_width=sink_width,
|
||||
window_width=window_width,
|
||||
multi_factor=multi_factor,
|
||||
entropy_factor=entropy_factor,
|
||||
)
|
||||
|
||||
return o.transpose(1, 2).contiguous()
|
||||
@@ -0,0 +1,84 @@
|
||||
# https://github.com/thu-ml/DiT-Extrapolation/blob/ultra-wan/sageattn/quant_per_block.py
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def quant_per_block_int8_kernel(Input, Output, Scale, L,
|
||||
stride_iz, stride_ih, stride_in,
|
||||
stride_oz, stride_oh, stride_on,
|
||||
stride_sz, stride_sh,
|
||||
sm_scale,
|
||||
C: tl.constexpr, BLK: tl.constexpr):
|
||||
off_blk = tl.program_id(0)
|
||||
off_h = tl.program_id(1)
|
||||
off_b = tl.program_id(2)
|
||||
|
||||
offs_n = off_blk * BLK + tl.arange(0, BLK)
|
||||
offs_k = tl.arange(0, C)
|
||||
|
||||
input_ptrs = Input + off_b * stride_iz + off_h * stride_ih + offs_n[:, None] * stride_in + offs_k[None, :]
|
||||
output_ptrs = Output + off_b * stride_oz + off_h * stride_oh + offs_n[:, None] * stride_on + offs_k[None, :]
|
||||
scale_ptrs = Scale + off_b * stride_sz + off_h * stride_sh + off_blk
|
||||
|
||||
x = tl.load(input_ptrs, mask=offs_n[:, None] < L)
|
||||
x = x.to(tl.float32)
|
||||
x *= sm_scale
|
||||
scale = tl.max(tl.abs(x)) / 127.
|
||||
x_int8 = x / scale
|
||||
x_int8 += 0.5 * tl.where(x_int8 >= 0, 1, -1)
|
||||
x_int8 = x_int8.to(tl.int8)
|
||||
tl.store(output_ptrs, x_int8, mask=offs_n[:, None] < L)
|
||||
tl.store(scale_ptrs, scale)
|
||||
|
||||
def per_block_int8(q, k, BLKQ=128, BLKK=64, sm_scale=None, tensor_layout="HND"):
|
||||
q_int8 = torch.empty(q.shape, dtype=torch.int8, device=q.device)
|
||||
k_int8 = torch.empty(k.shape, dtype=torch.int8, device=k.device)
|
||||
|
||||
if tensor_layout == "HND":
|
||||
b, h_qo, qo_len, head_dim = q.shape
|
||||
_, h_kv, kv_len, _ = k.shape
|
||||
|
||||
stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(1), q.stride(2)
|
||||
stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(1), q_int8.stride(2)
|
||||
stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(1), k.stride(2)
|
||||
stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(1), k_int8.stride(2)
|
||||
# elif tensor_layout == "NHD":
|
||||
# b, qo_len, h_qo, head_dim = q.shape
|
||||
# _, kv_len, h_kv, _ = k.shape
|
||||
|
||||
# stride_bz_q, stride_h_q, stride_seq_q = q.stride(0), q.stride(2), q.stride(1)
|
||||
# stride_bz_qo, stride_h_qo, stride_seq_qo = q_int8.stride(0), q_int8.stride(2), q_int8.stride(1)
|
||||
# stride_bz_k, stride_h_k, stride_seq_k = k.stride(0), k.stride(2), k.stride(1)
|
||||
# stride_bz_ko, stride_h_ko, stride_seq_ko = k_int8.stride(0), k_int8.stride(2), k_int8.stride(1)
|
||||
else:
|
||||
raise ValueError(f"Unknown tensor layout: {tensor_layout}")
|
||||
|
||||
q_scale = torch.empty((b, h_qo, (qo_len + BLKQ - 1) // BLKQ, 1), device=q.device, dtype=torch.float32)
|
||||
k_scale = torch.empty((b, h_kv, (kv_len + BLKK - 1) // BLKK, 1), device=q.device, dtype=torch.float32)
|
||||
|
||||
if sm_scale is None:
|
||||
sm_scale = head_dim**-0.5
|
||||
|
||||
grid = ((qo_len + BLKQ - 1) // BLKQ, h_qo, b)
|
||||
quant_per_block_int8_kernel[grid](
|
||||
q, q_int8, q_scale, qo_len,
|
||||
stride_bz_q, stride_h_q, stride_seq_q,
|
||||
stride_bz_qo, stride_h_qo, stride_seq_qo,
|
||||
q_scale.stride(0), q_scale.stride(1),
|
||||
sm_scale=(sm_scale * 1.44269504),
|
||||
C=head_dim, BLK=BLKQ
|
||||
)
|
||||
|
||||
grid = ((kv_len + BLKK - 1) // BLKK, h_kv, b)
|
||||
quant_per_block_int8_kernel[grid](
|
||||
k, k_int8, k_scale, kv_len,
|
||||
stride_bz_k, stride_h_k, stride_seq_k,
|
||||
stride_bz_ko, stride_h_ko, stride_seq_ko,
|
||||
k_scale.stride(0), k_scale.stride(1),
|
||||
sm_scale=1.0,
|
||||
C=head_dim, BLK=BLKK
|
||||
)
|
||||
|
||||
return q_int8, q_scale, k_int8, k_scale
|
||||
+6
-13
@@ -115,16 +115,9 @@ class WanRotaryPosEmbed(nn.Module):
|
||||
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
|
||||
return freqs
|
||||
|
||||
|
||||
from ..wanvideo.modules.attention import sageattn_func
|
||||
|
||||
def zero_module(module):
|
||||
# Zero out the parameters of a module and return it.
|
||||
for p in module.parameters():
|
||||
p.detach().zero_()
|
||||
return module
|
||||
|
||||
|
||||
class SimpleAttnProcessor2_0:
|
||||
def __init__(self, attention_mode):
|
||||
self.attention_mode = attention_mode
|
||||
@@ -278,7 +271,7 @@ class MaskCamEmbed(nn.Module):
|
||||
mid_channels = controlnet_cfg.get("mid_channels", 64)
|
||||
self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)),
|
||||
nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU())
|
||||
self.mask_zero_proj = zero_module(nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2)))
|
||||
self.mask_zero_proj = nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))
|
||||
|
||||
def forward(self, add_inputs: torch.Tensor):
|
||||
# render_mask.shape [b,c,f,h,w]
|
||||
@@ -321,7 +314,7 @@ class WanControlNet(ModelMixin):
|
||||
)
|
||||
self.proj_out = nn.ModuleList(
|
||||
[
|
||||
zero_module(nn.Linear(self.dim, 5120))
|
||||
nn.Linear(self.dim, 5120)
|
||||
for _ in range(controlnet_cfg["num_layers"])
|
||||
]
|
||||
)
|
||||
@@ -341,7 +334,7 @@ class WanControlNet(ModelMixin):
|
||||
|
||||
self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg)
|
||||
|
||||
def forward(self, render_latent, render_mask, camera_embedding, temb, device):
|
||||
def forward(self, render_latent, render_mask, camera_embedding, temb, out_device):
|
||||
controlnet_rotary_emb = self.controlnet_rope(render_latent)
|
||||
controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32))
|
||||
if not self.quantized:
|
||||
@@ -361,7 +354,7 @@ class WanControlNet(ModelMixin):
|
||||
if add_inputs is not None:
|
||||
add_inputs = self.controlnet_mask_embedding(add_inputs)
|
||||
controlnet_inputs = controlnet_inputs + add_inputs
|
||||
|
||||
|
||||
hidden_states = self.proj_in(controlnet_inputs)
|
||||
|
||||
controlnet_states = []
|
||||
@@ -371,6 +364,6 @@ class WanControlNet(ModelMixin):
|
||||
temb=temb,
|
||||
rotary_emb=controlnet_rotary_emb
|
||||
)
|
||||
controlnet_states.append(self.proj_out[i](hidden_states).to(device))
|
||||
controlnet_states.append(self.proj_out[i](hidden_states).to(out_device))
|
||||
|
||||
return controlnet_states
|
||||
|
||||
+29
-33
@@ -10,8 +10,6 @@ from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
import folder_paths
|
||||
|
||||
import json
|
||||
import numpy as np
|
||||
|
||||
class WanVideoUni3C_ControlnetLoader:
|
||||
@classmethod
|
||||
@@ -22,7 +20,7 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
|
||||
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
|
||||
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
"attention_mode": ([
|
||||
"sdpa",
|
||||
"sageattn",
|
||||
@@ -45,17 +43,17 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
transformer_load_device = device if load_device == "main_device" else offload_device
|
||||
|
||||
|
||||
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
|
||||
|
||||
|
||||
|
||||
model_path = folder_paths.get_full_path_or_raise("controlnet", model)
|
||||
|
||||
|
||||
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
|
||||
|
||||
if not "controlnet_patch_embedding.weight" in sd:
|
||||
raise ValueError("Invalid ControlNet model")
|
||||
|
||||
|
||||
in_channels = sd["controlnet_patch_embedding.weight"].shape[1]
|
||||
ffn_dim = sd["controlnet_blocks.0.ffn.0.bias"].shape[0]
|
||||
|
||||
@@ -79,7 +77,7 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
with init_empty_weights():
|
||||
controlnet = WanControlNet(controlnet_cfg)
|
||||
controlnet.eval()
|
||||
|
||||
|
||||
if quantization == "disabled":
|
||||
for k, v in sd.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
@@ -97,18 +95,18 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "proj_in"}
|
||||
|
||||
|
||||
log.info("Using accelerate to load and assign controlnet model weights to device...")
|
||||
param_count = sum(1 for _ in controlnet.named_parameters())
|
||||
for name, param in tqdm(controlnet.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
for name, param in tqdm(controlnet.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
total=param_count,
|
||||
leave=True):
|
||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
||||
if "controlnet_patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
set_module_tensor_to_device(controlnet, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
|
||||
|
||||
del sd
|
||||
|
||||
if compile_args is not None:
|
||||
@@ -123,8 +121,8 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
for i, block in enumerate(controlnet.controlnet_blocks):
|
||||
controlnet.controlnet_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
else:
|
||||
controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
|
||||
if load_device == "offload_device" and controlnet.device != offload_device:
|
||||
log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}")
|
||||
@@ -146,6 +144,7 @@ class WanVideoUni3C_embeds:
|
||||
"optional": {
|
||||
"render_latent": ("LATENT",),
|
||||
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
|
||||
"offload": ("BOOLEAN", {"default": True, "tooltip": "If enabled, the controlnet model will be offloaded before main model block processing to save VRAM."}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -154,9 +153,7 @@ class WanVideoUni3C_embeds:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None, offload=True):
|
||||
|
||||
latent_mask = latents = None
|
||||
if render_latent is not None:
|
||||
@@ -164,17 +161,17 @@ class WanVideoUni3C_embeds:
|
||||
# nframe = latents.shape[2] * 4
|
||||
# height = latents.shape[3] * 8
|
||||
# width = latents.shape[4] * 8
|
||||
|
||||
|
||||
if render_mask is not None:
|
||||
raise NotImplementedError("render_mask is not implemented at this time")
|
||||
mask = torch.nn.functional.interpolate(
|
||||
render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
size=(nframe, height, width),
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
).squeeze(0)
|
||||
latent_mask = mask.unsqueeze(0).to(device)
|
||||
log.info(f"latent mask shape {latent_mask.shape}")
|
||||
# mask = torch.nn.functional.interpolate(
|
||||
# render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
# size=(nframe, height, width),
|
||||
# mode='trilinear',
|
||||
# align_corners=False
|
||||
# ).squeeze(0)
|
||||
# latent_mask = mask.unsqueeze(0).to(device)
|
||||
# log.info(f"latent mask shape {latent_mask.shape}")
|
||||
|
||||
# # load camera
|
||||
# cam_info = json.load(open(f"{render_path}/cam_info.json"))
|
||||
@@ -199,7 +196,7 @@ class WanVideoUni3C_embeds:
|
||||
# K_inv = K.inverse()
|
||||
# intrinsic = K[None].repeat(nframe, 1, 1)
|
||||
|
||||
|
||||
|
||||
# w2c_0, c2w_0 = set_initial_camera(start_elevation, depth_avg)
|
||||
# w2cs, c2ws, intrinsic = build_cameras(cam_traj=cam_traj,
|
||||
# w2c_0=w2c_0,
|
||||
@@ -215,7 +212,7 @@ class WanVideoUni3C_embeds:
|
||||
# y_offset=y_offset,
|
||||
# z_offset=z_offset)
|
||||
|
||||
|
||||
|
||||
# from .camera import get_camera_embedding
|
||||
# camera_embedding = get_camera_embedding(intrinsic, w2cs, nframe, height, width, normalize=True)
|
||||
#print("camera embedding shape", camera_embedding.shape)
|
||||
@@ -227,11 +224,12 @@ class WanVideoUni3C_embeds:
|
||||
"end": end_percent,
|
||||
"render_latent": latents,
|
||||
"render_mask": latent_mask,
|
||||
"camera_embedding": None
|
||||
"camera_embedding": None,
|
||||
"offload": offload,
|
||||
}
|
||||
|
||||
|
||||
return (uni3c_embeds,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader,
|
||||
"WanVideoUni3C_embeds": WanVideoUni3C_embeds,
|
||||
@@ -240,5 +238,3 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader",
|
||||
"WanVideoUni3C_embeds": "WanVideo Uni3C Embeds",
|
||||
}
|
||||
|
||||
|
||||
+42
-53
@@ -9,37 +9,28 @@ from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import ProgressBar
|
||||
|
||||
import comfy.ops
|
||||
ops = comfy.ops.disable_weight_init
|
||||
|
||||
def update_transformer(transformer, state_dict):
|
||||
|
||||
|
||||
concat_dim = 4
|
||||
transformer.dwpose_embedding = nn.Sequential(
|
||||
nn.Conv3d(3, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
|
||||
nn.SiLU(),
|
||||
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
|
||||
nn.SiLU(),
|
||||
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
|
||||
nn.SiLU(),
|
||||
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,2,2), padding=(1,1,1)),
|
||||
nn.SiLU(),
|
||||
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv3d(concat_dim * 4, 5120, (1,2,2), stride=(1,2,2), padding=0))
|
||||
ops.Conv3d(3, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
|
||||
ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
|
||||
ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)), nn.SiLU(),
|
||||
ops.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,2,2), padding=(1,1,1)), nn.SiLU(),
|
||||
ops.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1), nn.SiLU(),
|
||||
ops.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1), nn.SiLU(),
|
||||
ops.Conv3d(concat_dim * 4, 5120, (1,2,2), stride=(1,2,2), padding=0))
|
||||
|
||||
randomref_dim = 20
|
||||
transformer.randomref_embedding_pose = nn.Sequential(
|
||||
nn.Conv2d(3, concat_dim * 4, 3, stride=1, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(3, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1), nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1), nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1), nn.SiLU(),
|
||||
nn.Conv2d(concat_dim * 4, randomref_dim, 3, stride=2, padding=1),
|
||||
)
|
||||
unianimate_sd = {}
|
||||
@@ -123,7 +114,7 @@ class DWposeDetector:
|
||||
body = candidate[:,:18].copy()
|
||||
body = body.reshape(nums*18, locs)
|
||||
score = subset[:,:18].copy()
|
||||
|
||||
|
||||
for i in range(len(score)):
|
||||
for j in range(len(score[i])):
|
||||
if score[i][j] > score_threshold:
|
||||
@@ -142,17 +133,17 @@ class DWposeDetector:
|
||||
else:
|
||||
bodyfoot_score[i][j] = -1
|
||||
if -1 not in bodyfoot_score[:,18] and -1 not in bodyfoot_score[:,19]:
|
||||
bodyfoot_score[:,18] = np.array([18.])
|
||||
bodyfoot_score[:,18] = np.array([18.])
|
||||
else:
|
||||
bodyfoot_score[:,18] = np.array([-1.])
|
||||
if -1 not in bodyfoot_score[:,21] and -1 not in bodyfoot_score[:,22]:
|
||||
bodyfoot_score[:,19] = np.array([19.])
|
||||
bodyfoot_score[:,19] = np.array([19.])
|
||||
else:
|
||||
bodyfoot_score[:,19] = np.array([-1.])
|
||||
bodyfoot_score = bodyfoot_score[:, :20]
|
||||
|
||||
bodyfoot = candidate[:,:24].copy()
|
||||
|
||||
|
||||
for i in range(nums):
|
||||
if -1 not in bodyfoot[i][18] and -1 not in bodyfoot[i][19]:
|
||||
bodyfoot[i][18] = (bodyfoot[i][18]+bodyfoot[i][19])/2
|
||||
@@ -162,7 +153,7 @@ class DWposeDetector:
|
||||
bodyfoot[i][19] = (bodyfoot[i][21]+bodyfoot[i][22])/2
|
||||
else:
|
||||
bodyfoot[i][19] = np.array([-1., -1.])
|
||||
|
||||
|
||||
bodyfoot = bodyfoot[:,:20,:]
|
||||
bodyfoot = bodyfoot.reshape(nums*20, locs)
|
||||
|
||||
@@ -172,7 +163,7 @@ class DWposeDetector:
|
||||
|
||||
hands = candidate[:,92:113]
|
||||
hands = np.vstack([hands, candidate[:,113:]])
|
||||
|
||||
|
||||
# bodies = dict(candidate=body, subset=score)
|
||||
bodies = dict(candidate=bodyfoot, subset=bodyfoot_score, score=bodyfoot_score)
|
||||
pose = dict(bodies=bodies, hands=hands, faces=faces)
|
||||
@@ -180,7 +171,7 @@ class DWposeDetector:
|
||||
# return draw_pose(pose, H, W)
|
||||
return pose
|
||||
|
||||
def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_feet=True,
|
||||
def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_feet=True,
|
||||
body_keypoint_size=4, hand_keypoint_size=4, draw_head=True):
|
||||
from .dwpose.util import draw_body_and_foot, draw_handpose, draw_facepose
|
||||
bodies = pose['bodies']
|
||||
@@ -202,7 +193,7 @@ def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_fe
|
||||
def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_threshold, stick_width,
|
||||
draw_body=True, draw_hands=True, hand_keypoint_size=4, draw_feet=True,
|
||||
body_keypoint_size=4, handle_not_detected="repeat", draw_head=True):
|
||||
|
||||
|
||||
results_vis = []
|
||||
comfy_pbar = ProgressBar(len(pose_images))
|
||||
|
||||
@@ -224,7 +215,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
||||
pose = np.zeros_like(img)
|
||||
results_vis.append(pose)
|
||||
comfy_pbar.update(1)
|
||||
|
||||
|
||||
bodies = results_vis[0]['bodies']
|
||||
faces = results_vis[0]['faces']
|
||||
hands = results_vis[0]['hands']
|
||||
@@ -268,7 +259,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
||||
results_vis[0]['faces'][:,:,1] *= y_ratio
|
||||
results_vis[0]['hands'][:,:,0] *= x_ratio
|
||||
results_vis[0]['hands'][:,:,1] *= y_ratio
|
||||
|
||||
|
||||
########neck########
|
||||
l_neck_ref = ((ref_candidate[0][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[0][1] - ref_candidate[1][1]) ** 2) ** 0.5
|
||||
l_neck_0 = ((candidate[0][0] - candidate[1][0]) ** 2 + (candidate[0][1] - candidate[1][1]) ** 2) ** 0.5
|
||||
@@ -287,7 +278,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
||||
results_vis[0]['bodies']['candidate'][16,1] += y_offset_neck
|
||||
results_vis[0]['bodies']['candidate'][17,0] += x_offset_neck
|
||||
results_vis[0]['bodies']['candidate'][17,1] += y_offset_neck
|
||||
|
||||
|
||||
########shoulder2########
|
||||
l_shoulder2_ref = ((ref_candidate[2][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[2][1] - ref_candidate[1][1]) ** 2) ** 0.5
|
||||
l_shoulder2_0 = ((candidate[2][0] - candidate[1][0]) ** 2 + (candidate[2][1] - candidate[1][1]) ** 2) ** 0.5
|
||||
@@ -435,9 +426,9 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
||||
|
||||
results_vis[0]['bodies']['candidate'][17,0] += x_offset_head17
|
||||
results_vis[0]['bodies']['candidate'][17,1] += y_offset_head17
|
||||
|
||||
|
||||
########MovingAverage########
|
||||
|
||||
|
||||
########left leg########
|
||||
l_ll1_ref = ((ref_candidate[8][0] - ref_candidate[9][0]) ** 2 + (ref_candidate[8][1] - ref_candidate[9][1]) ** 2) ** 0.5
|
||||
l_ll1_0 = ((candidate[8][0] - candidate[9][0]) ** 2 + (candidate[8][1] - candidate[9][1]) ** 2) ** 0.5
|
||||
@@ -522,7 +513,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
||||
results_vis[i]['bodies']['candidate'][17,1] += y_offset_neck
|
||||
|
||||
########shoulder2########
|
||||
|
||||
|
||||
|
||||
x_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][2][0])*(1.-shoulder2_ratio)
|
||||
y_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][2][1])*(1.-shoulder2_ratio)
|
||||
@@ -676,7 +667,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
||||
results_vis[i]['bodies']['candidate'] += offset[np.newaxis, :]
|
||||
results_vis[i]['faces'] += offset[np.newaxis, np.newaxis, :]
|
||||
results_vis[i]['hands'] += offset[np.newaxis, np.newaxis, :]
|
||||
|
||||
|
||||
dwpose_woface_list = []
|
||||
for i in range(len(results_vis)):
|
||||
#try:
|
||||
@@ -724,11 +715,11 @@ class WanVideoUniAnimateDWPoseDetector:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, pose_images, score_threshold, stick_width, reference_pose_image=None, draw_body=True, body_keypoint_size=4,
|
||||
def process(self, pose_images, score_threshold, stick_width, reference_pose_image=None, draw_body=True, body_keypoint_size=4,
|
||||
draw_feet=True, draw_hands=True, hand_keypoint_size=4, colorspace="RGB", handle_not_detected="empty", draw_head=True):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
|
||||
|
||||
#model loading
|
||||
dw_pose_model = "dw-ll_ucoco_384_bs5.torchscript.pt"
|
||||
yolo_model = "yolox_l.torchscript.pt"
|
||||
@@ -742,27 +733,27 @@ class WanVideoUniAnimateDWPoseDetector:
|
||||
if not os.path.exists(model_det):
|
||||
log.info(f"Downloading yolo model to: {model_base_path}")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id="hr16/yolox-onnx",
|
||||
snapshot_download(repo_id="hr16/yolox-onnx",
|
||||
allow_patterns=[f"*{yolo_model}*"],
|
||||
local_dir=model_base_path,
|
||||
local_dir=model_base_path,
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
|
||||
if not os.path.exists(model_pose):
|
||||
log.info(f"Downloading dwpose model to: {model_base_path}")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5",
|
||||
snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5",
|
||||
allow_patterns=[f"*{dw_pose_model}*"],
|
||||
local_dir=model_base_path,
|
||||
local_dir=model_base_path,
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
if not hasattr(self, "det") or not hasattr(self, "pose"):
|
||||
self.det = torch.jit.load(model_det, map_location=device)
|
||||
self.pose = torch.jit.load(model_pose, map_location=device)
|
||||
self.dwpose_detector = DWposeDetector(self.det, self.pose)
|
||||
self.dwpose_detector = DWposeDetector(self.det, self.pose)
|
||||
|
||||
#model inference
|
||||
height, width = pose_images.shape[1:3]
|
||||
|
||||
|
||||
pose_np = pose_images.cpu().numpy() * 255
|
||||
ref_np = None
|
||||
if reference_pose_image is not None:
|
||||
@@ -772,11 +763,11 @@ class WanVideoUniAnimateDWPoseDetector:
|
||||
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_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:
|
||||
@@ -828,11 +819,9 @@ class WanVideoUniAnimatePoseInput:
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoUniAnimatePoseInput": WanVideoUniAnimatePoseInput,
|
||||
"WanVideoUniAnimateDWPoseDetector": WanVideoUniAnimateDWPoseDetector,
|
||||
|
||||
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoUniAnimatePoseInput": "WanVideo UniAnimate Pose Input",
|
||||
"WanVideoUniAnimateDWPoseDetector": "WanVideo UniAnimate DWPose Detector",
|
||||
}
|
||||
|
||||
|
||||
@@ -4,17 +4,116 @@ import logging
|
||||
import math
|
||||
from tqdm import tqdm
|
||||
from pathlib import Path
|
||||
import os
|
||||
import gc
|
||||
import types, collections
|
||||
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
|
||||
from comfy.model_patcher import get_key_weight, string_to_seed
|
||||
from comfy.lora import calculate_weight
|
||||
from comfy.model_management import cast_to_device
|
||||
|
||||
from comfy.float import stochastic_rounding
|
||||
from .custom_linear import remove_lora_from_module
|
||||
import folder_paths
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
import comfy.model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
try:
|
||||
from .gguf.gguf import GGUFParameter
|
||||
except:
|
||||
pass
|
||||
|
||||
COLOR_CODES = {
|
||||
"reset": "\033[0m",
|
||||
"red": "\033[31m",
|
||||
"green": "\033[32m",
|
||||
"yellow": "\033[33m",
|
||||
"blue": "\033[34m",
|
||||
"magenta": "\033[35m",
|
||||
"cyan": "\033[36m",
|
||||
"white": "\033[37m",
|
||||
}
|
||||
|
||||
def color_text(text, color):
|
||||
try:
|
||||
return f"{COLOR_CODES.get(color, COLOR_CODES['reset'])}{text}{COLOR_CODES['reset']}"
|
||||
except Exception:
|
||||
return text
|
||||
|
||||
class MetaParameter(torch.nn.Parameter):
|
||||
def __new__(cls, dtype, quant_type=None):
|
||||
data = torch.empty(0, dtype=dtype)
|
||||
self = torch.nn.Parameter(data, requires_grad=False)
|
||||
self.quant_type = quant_type
|
||||
return self
|
||||
|
||||
def offload_transformer(transformer, remove_lora=True):
|
||||
transformer.teacache_state.clear_all()
|
||||
transformer.magcache_state.clear_all()
|
||||
transformer.easycache_state.clear_all()
|
||||
|
||||
if transformer.patched_linear:
|
||||
for name, param in transformer.named_parameters():
|
||||
if "loras" in name or "controlnet" in name:
|
||||
continue
|
||||
module = transformer
|
||||
subnames = name.split('.')
|
||||
for subname in subnames[:-1]:
|
||||
module = getattr(module, subname)
|
||||
attr_name = subnames[-1]
|
||||
if param.data.is_floating_point():
|
||||
meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
|
||||
setattr(module, attr_name, meta_param)
|
||||
elif isinstance(param.data, GGUFParameter):
|
||||
quant_type = getattr(param, 'quant_type', None)
|
||||
setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type))
|
||||
else:
|
||||
pass
|
||||
if remove_lora:
|
||||
remove_lora_from_module(transformer)
|
||||
else:
|
||||
transformer.to(offload_device)
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.kv_cache = None
|
||||
if transformer.audio_model is not None and hasattr(block, 'audio_block'):
|
||||
block.audio_block = None
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
def init_blockswap(transformer, block_swap_args, model):
|
||||
if not transformer.patched_linear:
|
||||
if block_swap_args is not None:
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name or "control_adapter" in name or "face" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
else:
|
||||
transformer.to(device)
|
||||
|
||||
def check_device_same(first_device, second_device):
|
||||
if first_device.type != second_device.type:
|
||||
return False
|
||||
@@ -108,13 +207,11 @@ def check_diffusers_version():
|
||||
except importlib.metadata.PackageNotFoundError:
|
||||
raise AssertionError("diffusers is not installed.")
|
||||
|
||||
def print_memory(device):
|
||||
memory = torch.cuda.memory_allocated(device) / 1024**3
|
||||
def print_memory(device, process="Sampling"):
|
||||
max_memory = torch.cuda.max_memory_allocated(device) / 1024**3
|
||||
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
|
||||
log.info(f"Allocated memory: {memory=:.3f} GB")
|
||||
log.info(f"Max allocated memory: {max_memory=:.3f} GB")
|
||||
log.info(f"Max reserved memory: {max_reserved=:.3f} GB")
|
||||
log.info(f"[{process}] Max allocated memory: {max_memory=:.3f} GB")
|
||||
log.info(f"[{process}] Max reserved memory: {max_reserved=:.3f} GB")
|
||||
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
|
||||
#log.info(f"Memory Summary:\n{memory_summary}")
|
||||
|
||||
@@ -125,6 +222,18 @@ def get_module_memory_mb(module):
|
||||
memory += param.nelement() * param.element_size()
|
||||
return memory / (1024 * 1024) # Convert to MB
|
||||
|
||||
def get_module_memory_mb_per_device(module):
|
||||
memory_per_device = {}
|
||||
memory = 0
|
||||
for param in module.parameters():
|
||||
if param.data is not None:
|
||||
device = str(param.device)
|
||||
memory += param.nelement() * param.element_size()
|
||||
memory_per_device[device] = memory_per_device.get(device, 0) + memory
|
||||
|
||||
memory_per_device = {dev: mem / (1024 * 1024) for dev, mem in memory_per_device.items()}
|
||||
return memory_per_device
|
||||
|
||||
def get_tensor_memory(tensor):
|
||||
memory_bytes = tensor.element_size() * tensor.nelement()
|
||||
return f"{memory_bytes / (1024 * 1024):.2f} MB"
|
||||
@@ -140,7 +249,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back
|
||||
self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update)
|
||||
|
||||
if device_to is not None:
|
||||
temp_weight = cast_to_device(weight, device_to, torch.float32, copy=True)
|
||||
temp_weight = mm.cast_to_device(weight, device_to, torch.float32, copy=True)
|
||||
else:
|
||||
temp_weight = weight.to(torch.float32, copy=True)
|
||||
if convert_func is not None:
|
||||
@@ -584,9 +693,9 @@ def check_duplicate_nodes():
|
||||
"""Check ComfyUI custom_nodes directory for duplicate installations"""
|
||||
custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0])
|
||||
current_path = Path(__file__).parent
|
||||
|
||||
|
||||
wanvideo_dirs = []
|
||||
|
||||
|
||||
# Check all directories in custom_nodes
|
||||
for path in custom_nodes_dir.iterdir():
|
||||
if (path.is_dir() and
|
||||
@@ -594,7 +703,7 @@ def check_duplicate_nodes():
|
||||
'wanvideo' in path.name.lower() and
|
||||
'wrapper' in path.name.lower()):
|
||||
wanvideo_dirs.append(str(path))
|
||||
|
||||
|
||||
return wanvideo_dirs
|
||||
|
||||
#https://github.com/temporalscorerescaling/TSR/
|
||||
|
||||
+71
-195
@@ -1,25 +1,21 @@
|
||||
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
|
||||
import torch
|
||||
from ...utils import log
|
||||
|
||||
# Flash Attention imports
|
||||
try:
|
||||
import flash_attn_interface
|
||||
FLASH_ATTN_3_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_3_AVAILABLE = False
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
|
||||
def attention_func_error(*args, **kwargs):
|
||||
raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
|
||||
|
||||
from .attention_flash import flash_attention
|
||||
|
||||
try:
|
||||
import flash_attn
|
||||
FLASH_ATTN_2_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_2_AVAILABLE = False
|
||||
|
||||
# Sage Attention imports
|
||||
# using custom ops to avoid graph breaks with torch.compile
|
||||
try:
|
||||
from sageattention import sageattn
|
||||
@torch.compiler.disable()
|
||||
def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"):
|
||||
|
||||
@torch.library.custom_op("wanvideo::sageattn", mutates_args=())
|
||||
def sageattn_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, tensor_layout: str = "HND"
|
||||
) -> torch.Tensor:
|
||||
if not (q.dtype == k.dtype == v.dtype):
|
||||
return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
|
||||
elif q.dtype == torch.float32:
|
||||
@@ -27,6 +23,13 @@ try:
|
||||
else:
|
||||
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
|
||||
|
||||
@sageattn_func.register_fake
|
||||
def _(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False, tensor_layout="HND"):
|
||||
# Return tensor with same shape as q
|
||||
return q.clone()
|
||||
|
||||
sageattn_func = torch.ops.wanvideo.sageattn
|
||||
|
||||
def sageattn_func_compiled(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"):
|
||||
if not (q.dtype == k.dtype == v.dtype):
|
||||
return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
|
||||
@@ -40,20 +43,14 @@ except Exception as e:
|
||||
log.warning("sageattention package is not installed, sageattention will not be available")
|
||||
elif isinstance(e, ImportError) and "DLL" in str(e):
|
||||
log.warning("sageattention DLL loading error, sageattention will not be available")
|
||||
sageattn_func = None
|
||||
sageattn_func = attention_func_error
|
||||
|
||||
try:
|
||||
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
|
||||
except:
|
||||
try:
|
||||
from sageattn import sageattn_blackwell
|
||||
except:
|
||||
SAGE3_AVAILABLE = False
|
||||
|
||||
try:
|
||||
from sageattention import sageattn_varlen
|
||||
@torch.compiler.disable()
|
||||
def sageattn_varlen_func(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0, is_causal=False):
|
||||
from typing import List
|
||||
|
||||
@torch.library.custom_op("wanvideo::sageattn_varlen", mutates_args=())
|
||||
def sageattn_varlen_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, q_lens: List[int], k_lens: List[int], max_seqlen_q: int, max_seqlen_k: int, dropout_p: float = 0.0, is_causal: bool = False) -> torch.Tensor:
|
||||
cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32)
|
||||
cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32)
|
||||
if not (q.dtype == k.dtype == v.dtype):
|
||||
@@ -62,180 +59,59 @@ try:
|
||||
return sageattn_varlen(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32)
|
||||
else:
|
||||
return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal)
|
||||
except:
|
||||
sageattn_varlen_func = None
|
||||
|
||||
__all__ = [
|
||||
'flash_attention',
|
||||
'attention',
|
||||
]
|
||||
@sageattn_varlen_func.register_fake
|
||||
def _(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0.0, is_causal=False):
|
||||
# Return tensor with same shape as q
|
||||
return q.clone()
|
||||
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
|
||||
except:
|
||||
sageattn_varlen_func = attention_func_error
|
||||
|
||||
# sage3
|
||||
try:
|
||||
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
|
||||
except:
|
||||
try:
|
||||
from sageattn import sageattn_blackwell
|
||||
except:
|
||||
sageattn_blackwell = attention_func_error
|
||||
|
||||
try:
|
||||
from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico
|
||||
@torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=())
|
||||
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9, frame_tokens: int = 1536
|
||||
) -> torch.Tensor:
|
||||
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor, frame_tokens=frame_tokens)
|
||||
|
||||
@sageattn_func_ultravico.register_fake
|
||||
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
|
||||
return torch.empty_like(qkv[0]).contiguous()
|
||||
sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico
|
||||
except:
|
||||
sageattn_func_ultravico = attention_func_error
|
||||
|
||||
|
||||
def flash_attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
q_lens=None,
|
||||
k_lens=None,
|
||||
dropout_p=0.,
|
||||
softmax_scale=None,
|
||||
q_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
dtype=torch.bfloat16,
|
||||
version=None,
|
||||
):
|
||||
"""
|
||||
q: [B, Lq, Nq, C1].
|
||||
k: [B, Lk, Nk, C1].
|
||||
v: [B, Lk, Nk, C2]. Nq must be divisible by Nk.
|
||||
q_lens: [B].
|
||||
k_lens: [B].
|
||||
dropout_p: float. Dropout probability.
|
||||
softmax_scale: float. The scaling of QK^T before applying softmax.
|
||||
causal: bool. Whether to apply causal attention mask.
|
||||
window_size: (left right). If not (-1, -1), apply sliding window local attention.
|
||||
deterministic: bool. If True, slightly slower and uses more memory.
|
||||
dtype: torch.dtype. Apply when dtype of q/k/v is not float16/bfloat16.
|
||||
"""
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
#assert dtype in half_dtypes
|
||||
#assert q.device.type == 'cuda' and q.size(-1) <= 256
|
||||
|
||||
# params
|
||||
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
# preprocess query
|
||||
if q_lens is None:
|
||||
q = half(q.flatten(0, 1))
|
||||
q_lens = torch.tensor(
|
||||
[lq] * b, dtype=torch.int32).to(
|
||||
device=q.device, non_blocking=True)
|
||||
else:
|
||||
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
|
||||
|
||||
# preprocess key, value
|
||||
if k_lens is None:
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
k_lens = torch.tensor(
|
||||
[lk] * b, dtype=torch.int32).to(
|
||||
device=k.device, non_blocking=True)
|
||||
else:
|
||||
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
|
||||
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
|
||||
|
||||
q = q.to(v.dtype)
|
||||
k = k.to(v.dtype)
|
||||
|
||||
if q_scale is not None:
|
||||
q = q * q_scale
|
||||
|
||||
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
|
||||
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
|
||||
|
||||
# apply attention
|
||||
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
|
||||
# Note: dropout_p, window_size are not supported in FA3 now.
|
||||
x = flash_attn_interface.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
seqused_q=None,
|
||||
seqused_k=None,
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic)[0].unflatten(0, (b, lq))
|
||||
else:
|
||||
assert FLASH_ATTN_2_AVAILABLE
|
||||
x = flash_attn.flash_attn_varlen_func(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
max_seqlen_q=lq,
|
||||
max_seqlen_k=lk,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic).unflatten(0, (b, lq))
|
||||
|
||||
# output
|
||||
return x.type(out_dtype)
|
||||
|
||||
|
||||
def attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
q_lens=None,
|
||||
k_lens=None,
|
||||
max_seqlen_q=None,
|
||||
max_seqlen_k=None,
|
||||
dropout_p=0.,
|
||||
softmax_scale=None,
|
||||
q_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
dtype=torch.bfloat16,
|
||||
attention_mode='sdpa',
|
||||
attn_mask=None,
|
||||
):
|
||||
def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.,
|
||||
softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16,
|
||||
attention_mode='sdpa', attn_mask=None, transformer_options={}, frame_tokens=1536, heads=128):
|
||||
if "flash" in attention_mode:
|
||||
if attention_mode == 'flash_attn_2':
|
||||
fa_version = 2
|
||||
elif attention_mode == 'flash_attn_3':
|
||||
fa_version = 3
|
||||
return flash_attention(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
q_lens=q_lens,
|
||||
k_lens=k_lens,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
q_scale=q_scale,
|
||||
causal=causal,
|
||||
window_size=window_size,
|
||||
deterministic=deterministic,
|
||||
dtype=dtype,
|
||||
version=fa_version,
|
||||
return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale,
|
||||
q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3,
|
||||
)
|
||||
elif attention_mode == 'sdpa':
|
||||
elif attention_mode == 'sageattn_3':
|
||||
return sageattn_blackwell(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), per_block_mean=False).transpose(1,2).contiguous()
|
||||
elif attention_mode == 'sageattn_varlen':
|
||||
return sageattn_varlen_func(q,k,v, q_lens=q_lens, k_lens=k_lens, max_seqlen_k=max_seqlen_k, max_seqlen_q=max_seqlen_q)
|
||||
elif attention_mode == 'sageattn_compiled': # for sage versions that allow torch.compile, may be redundant now as other sageattn ops are wrapper in custom ops
|
||||
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
|
||||
elif attention_mode == 'sageattn':
|
||||
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
|
||||
elif attention_mode == 'sageattn_ultravico':
|
||||
return sageattn_func_ultravico([q, k, v], multi_factor=transformer_options.get("ultravico_alpha", 0.9), frame_tokens=frame_tokens).contiguous()
|
||||
elif attention_mode == 'comfy':
|
||||
return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True)
|
||||
else: # sdpa
|
||||
if not (q.dtype == k.dtype == v.dtype):
|
||||
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype), attn_mask=attn_mask).transpose(1, 2).contiguous()
|
||||
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), attn_mask=attn_mask).transpose(1, 2).contiguous()
|
||||
elif attention_mode == 'sageattn_3':
|
||||
return sageattn_blackwell(
|
||||
q.transpose(1,2),
|
||||
k.transpose(1,2),
|
||||
v.transpose(1,2),
|
||||
per_block_mean=False #seems necessary for reasonable VRAM usage, not sure of other implications
|
||||
).transpose(1,2).contiguous()
|
||||
elif attention_mode == 'sageattn_varlen':
|
||||
return sageattn_varlen_func(
|
||||
q,k,v,
|
||||
q_lens=q_lens,
|
||||
k_lens=k_lens,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
max_seqlen_q=max_seqlen_q
|
||||
)
|
||||
elif attention_mode == 'sageattn_compiled':
|
||||
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
|
||||
else:
|
||||
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
import torch
|
||||
from ...utils import log
|
||||
|
||||
def attention_func_error(*args, **kwargs):
|
||||
raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.")
|
||||
|
||||
try:
|
||||
import flash_attn_interface
|
||||
FLASH_ATTN_3_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_3_AVAILABLE = False
|
||||
|
||||
try:
|
||||
import flash_attn
|
||||
FLASH_ATTN_2_AVAILABLE = True
|
||||
except Exception as e:
|
||||
FLASH_ATTN_2_AVAILABLE = False
|
||||
|
||||
if not FLASH_ATTN_2_AVAILABLE and not FLASH_ATTN_3_AVAILABLE:
|
||||
flash_attention = attention_func_error
|
||||
else:
|
||||
def flash_attention(q, k, v, q_lens=None, k_lens=None, dropout_p=0., softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16, version=None):
|
||||
half_dtypes = (torch.float16, torch.bfloat16)
|
||||
|
||||
# params
|
||||
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in half_dtypes else x.to(dtype)
|
||||
|
||||
# preprocess query
|
||||
if q_lens is None:
|
||||
q = half(q.flatten(0, 1))
|
||||
q_lens = torch.tensor(
|
||||
[lq] * b, dtype=torch.int32).to(
|
||||
device=q.device, non_blocking=True)
|
||||
else:
|
||||
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens)]))
|
||||
|
||||
# preprocess key, value
|
||||
if k_lens is None:
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
k_lens = torch.tensor(
|
||||
[lk] * b, dtype=torch.int32).to(
|
||||
device=k.device, non_blocking=True)
|
||||
else:
|
||||
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens)]))
|
||||
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens)]))
|
||||
|
||||
q = q.to(v.dtype)
|
||||
k = k.to(v.dtype)
|
||||
|
||||
if q_scale is not None:
|
||||
q = q * q_scale
|
||||
|
||||
if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE:
|
||||
log.warning('Flash attention 3 is not available, use flash attention 2 instead.')
|
||||
|
||||
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
|
||||
# Note: dropout_p, window_size are not supported in FA3 now.
|
||||
x = flash_attn_interface.flash_attn_varlen_func(q=q, k=k, v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
seqused_q=None, seqused_k=None, max_seqlen_q=lq, max_seqlen_k=lk,
|
||||
softmax_scale=softmax_scale, causal=causal,
|
||||
deterministic=deterministic).unflatten(0, (b, lq))
|
||||
else:
|
||||
assert FLASH_ATTN_2_AVAILABLE
|
||||
x = flash_attn.flash_attn_varlen_func(q=q, k=k, v=v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
|
||||
0, dtype=torch.int32).to(q.device, non_blocking=True),
|
||||
max_seqlen_q=lq, max_seqlen_k=lk, dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale, causal=causal, window_size=window_size,
|
||||
deterministic=deterministic).unflatten(0, (b, lq))
|
||||
return x.type(out_dtype)
|
||||
+824
-589
File diff suppressed because it is too large
Load Diff
@@ -1,12 +1,15 @@
|
||||
import torch
|
||||
from .fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas, retrieve_timesteps)
|
||||
import numpy as np
|
||||
from .fm_solvers import (FlowDPMSolverMultistepScheduler)
|
||||
from .fm_solvers_unipc import FlowUniPCMultistepScheduler
|
||||
from .basic_flowmatch import FlowMatchScheduler
|
||||
from .flowmatch_pusa import FlowMatchSchedulerPusa
|
||||
from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep
|
||||
from .ersde_scheduler import ERSDEScheduler
|
||||
from .scheduling_flow_match_lcm import FlowMatchLCMScheduler
|
||||
from .fm_sa_ode import FlowMatchSAODEStableScheduler
|
||||
from .fm_rcm import rCMFlowMatchScheduler
|
||||
from .vitb_unipc import ViBTScheduler
|
||||
from ...utils import log
|
||||
|
||||
try:
|
||||
@@ -24,25 +27,35 @@ scheduler_list = [
|
||||
"deis",
|
||||
"lcm", "lcm/beta",
|
||||
"res_multistep",
|
||||
"er_sde",
|
||||
"flowmatch_causvid",
|
||||
"flowmatch_distill",
|
||||
"flowmatch_pusa",
|
||||
"multitalk",
|
||||
"sa_ode_stable",
|
||||
"rcm"
|
||||
"rcm",
|
||||
"vibt_unipc",
|
||||
]
|
||||
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False, **kwargs):
|
||||
def _apply_custom_sigmas(sample_scheduler, sigmas, device):
|
||||
sample_scheduler.sigmas = sigmas.to(device)
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, denoise_strength=1.0, sigmas=None, log_timesteps=False, enhance_hf=False, **kwargs):
|
||||
timesteps = None
|
||||
if 'unipc' in scheduler:
|
||||
if sigmas is not None:
|
||||
steps = len(sigmas) - 1
|
||||
if scheduler == 'vibt_unipc':
|
||||
sample_scheduler = ViBTScheduler()
|
||||
sample_scheduler.set_parameters(shift=shift)
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
elif 'unipc' in scheduler:
|
||||
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
|
||||
else:
|
||||
sample_scheduler.sigmas = sigmas.to(device)
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']:
|
||||
if 'longcat' in scheduler:
|
||||
num_distill_sample_steps = 50
|
||||
@@ -57,10 +70,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas)
|
||||
else:
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
|
||||
if flowedit_args: #seems to work better
|
||||
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
else:
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif 'dpm' in scheduler:
|
||||
if 'sde' in scheduler:
|
||||
algorithm_type = "sde-dpmsolver++"
|
||||
@@ -70,16 +83,20 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler))
|
||||
else:
|
||||
sample_scheduler.sigmas = sigmas.to(device)
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif scheduler == 'deis':
|
||||
sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift)
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
sample_scheduler.sigmas[-1] = 1e-6
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
sample_scheduler.sigmas[-1] = 1e-6
|
||||
else:
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif 'lcm' in scheduler:
|
||||
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
else:
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif 'flowmatch_causvid' in scheduler:
|
||||
if sigmas is not None:
|
||||
raise NotImplementedError("This scheduler does not support custom sigmas")
|
||||
@@ -112,21 +129,50 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
|
||||
elif 'flowmatch_pusa' in scheduler:
|
||||
sample_scheduler = FlowMatchSchedulerPusa(shift=shift, sigma_min=0.0, extra_one_step=True)
|
||||
sample_scheduler.set_timesteps(steps+1, denoising_strength=denoise_strength, shift=shift,
|
||||
sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps+1, denoising_strength=denoise_strength, shift=shift)
|
||||
else:
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif scheduler == 'res_multistep':
|
||||
sample_scheduler = FlowMatchSchedulerResMultistep(shift=shift)
|
||||
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
|
||||
else:
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif scheduler == 'er_sde':
|
||||
sample_scheduler = ERSDEScheduler(shift=shift)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength)
|
||||
else:
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif "sa_ode_stable" in scheduler:
|
||||
sample_scheduler = FlowMatchSAODEStableScheduler(shift=shift, **kwargs)
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
else:
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
elif 'rcm' in scheduler:
|
||||
sample_scheduler = rCMFlowMatchScheduler()
|
||||
sample_scheduler.set_timesteps(steps, sigma_max=120)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, sigma_max=120)
|
||||
else:
|
||||
_apply_custom_sigmas(sample_scheduler, sigmas, device)
|
||||
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
|
||||
if enhance_hf:
|
||||
num_tail_uniform_steps = max(3, min(15, int(len(timesteps) * 0.2))) # Use 20% of steps for uniform tail (minimum 3, maximum 15)
|
||||
tail_uniform_start = float(timesteps.max()) * 0.5 # Split at 50% of the timestep range
|
||||
tail_uniform_end = 0
|
||||
|
||||
timesteps_uniform_tail = list(np.linspace(tail_uniform_start, tail_uniform_end, num_tail_uniform_steps, dtype=np.float32, endpoint=(tail_uniform_end != 0)))
|
||||
timesteps_uniform_tail = [torch.tensor(t, device=device).unsqueeze(0) for t in timesteps_uniform_tail]
|
||||
filtered_timesteps = [timestep.unsqueeze(0).to(device) for timestep in timesteps if timestep > tail_uniform_start]
|
||||
timesteps = torch.cat(filtered_timesteps + timesteps_uniform_tail)
|
||||
sample_scheduler.timesteps = timesteps
|
||||
sample_scheduler.sigmas = torch.cat([timesteps / 1000, torch.zeros(1, device=timesteps.device)])
|
||||
|
||||
steps = len(timesteps)
|
||||
if (isinstance(start_step, int) and end_step != -1 and start_step >= end_step) or (not isinstance(start_step, int) and start_step != -1 and end_step >= start_step):
|
||||
raise ValueError("start_step must be less than end_step")
|
||||
@@ -136,7 +182,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
end_idx = len(timesteps) - 1
|
||||
|
||||
if log_timesteps:
|
||||
log.info(f"------- Scheduler info -------")
|
||||
log.info("------- Scheduler info -------")
|
||||
log.info(f"Total timesteps: {timesteps}")
|
||||
|
||||
if isinstance(start_step, float):
|
||||
@@ -156,6 +202,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
end_idx = end_step - 1
|
||||
|
||||
# Slice timesteps and sigmas once, based on indices
|
||||
all_timesteps = timesteps
|
||||
timesteps = timesteps[start_idx:end_idx+1]
|
||||
sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone()
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
|
||||
@@ -163,9 +210,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
if log_timesteps:
|
||||
log.info(f"Using timesteps: {timesteps}")
|
||||
log.info(f"Using sigmas: {sample_scheduler.sigmas}")
|
||||
log.info(f"------------------------------")
|
||||
log.info("------------------------------")
|
||||
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
setattr(sample_scheduler, 'all_timesteps', all_timesteps)
|
||||
|
||||
return sample_scheduler, timesteps, start_idx, end_idx
|
||||
return sample_scheduler, timesteps, start_idx, end_idx
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
import torch
|
||||
|
||||
class ERSDEScheduler():
|
||||
"""Extended Reverse-Time SDE solver (VP ER-SDE-Solver-3).
|
||||
|
||||
Based on: arXiv: https://arxiv.org/abs/2309.06169
|
||||
Code reference: https://github.com/QinpengCui/ER-SDE-Solver/blob/main/er_sde_solver.py
|
||||
"""
|
||||
|
||||
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0,
|
||||
sigma_max=1.0, sigma_min=0.003 / 1.002, max_stage=3, s_noise=1.0,
|
||||
num_integration_points=200):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.shift = shift
|
||||
self.sigma_max = sigma_max
|
||||
self.sigma_min = sigma_min
|
||||
self.max_stage = max_stage
|
||||
self.s_noise = s_noise
|
||||
self.num_integration_points = num_integration_points
|
||||
self.set_timesteps(num_inference_steps)
|
||||
self.old_denoised = None
|
||||
self.old_denoised_d = None
|
||||
self.step_index = 0
|
||||
|
||||
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, sigmas=None):
|
||||
"""Generate the full sigma schedule (from max to min)."""
|
||||
full_sigmas = torch.linspace(self.sigma_max, self.sigma_min, self.num_train_timesteps)
|
||||
ss = len(full_sigmas) / num_inference_steps
|
||||
if sigmas is None:
|
||||
sigmas = []
|
||||
for x in range(num_inference_steps):
|
||||
idx = int(round(x * ss))
|
||||
sigmas.append(float(full_sigmas[idx]))
|
||||
sigmas.append(0.0)
|
||||
self.sigmas = torch.FloatTensor(sigmas)
|
||||
self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
|
||||
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
|
||||
self.step_index = 0
|
||||
self.old_denoised = None
|
||||
self.old_denoised_d = None
|
||||
|
||||
def default_er_sde_noise_scaler(self, x):
|
||||
return x * ((x ** 0.3).exp() + 10.0)
|
||||
|
||||
def step(self, model_output, timestep, sample, generator):
|
||||
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
|
||||
self.sigmas = self.sigmas.to(model_output.device)
|
||||
self.timesteps = self.timesteps.to(model_output.device)
|
||||
|
||||
if timestep.ndim == 0:
|
||||
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
|
||||
else:
|
||||
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
|
||||
noise_scaler = self.default_er_sde_noise_scaler
|
||||
|
||||
# Get current and next sigma
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
if (timestep_id + 1 >= len(self.sigmas)).any():
|
||||
sigma_next = torch.zeros_like(sigma)
|
||||
else:
|
||||
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
|
||||
er_lambda_s = sigma
|
||||
er_lambda_t = sigma_next
|
||||
|
||||
# Calculate alpha values
|
||||
alpha_s = sigma / (er_lambda_s + 1e-10)
|
||||
alpha_t = sigma_next / (er_lambda_t + 1e-10)
|
||||
r_alpha = alpha_t / (alpha_s + 1e-10)
|
||||
|
||||
# Denoised prediction (x_0 estimate)
|
||||
denoised = sample - sigma * model_output
|
||||
|
||||
# Determine which stage to use
|
||||
stage_used = min(self.max_stage, self.step_index + 1)
|
||||
|
||||
if sigma_next == 0 or (sigma_next == 0.0).all():
|
||||
# Final step - return denoised
|
||||
x = denoised
|
||||
else:
|
||||
r = noise_scaler(er_lambda_t) / (noise_scaler(er_lambda_s) + 1e-10)
|
||||
|
||||
# Stage 1: Euler step
|
||||
x = r_alpha * r * sample + alpha_t * (1 - r) * denoised
|
||||
|
||||
if stage_used >= 2 and self.old_denoised is not None:
|
||||
dt = er_lambda_t - er_lambda_s
|
||||
lambda_step_size = -dt / self.num_integration_points
|
||||
|
||||
# Create integration points
|
||||
point_indice = torch.arange(0, self.num_integration_points,
|
||||
dtype=torch.float32, device=sample.device)
|
||||
lambda_pos = er_lambda_t + point_indice * lambda_step_size
|
||||
scaled_pos = noise_scaler(lambda_pos)
|
||||
|
||||
# Stage 2: Second-order correction
|
||||
s = torch.sum(1 / (scaled_pos + 1e-10)) * lambda_step_size
|
||||
|
||||
# Get previous sigma for derivative calculation
|
||||
if timestep_id > 0:
|
||||
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1)
|
||||
er_lambda_prev = sigma_prev
|
||||
else:
|
||||
er_lambda_prev = er_lambda_s
|
||||
|
||||
denoised_d = (denoised - self.old_denoised) / ((er_lambda_s - er_lambda_prev) + 1e-10)
|
||||
x = x + alpha_t * (dt + s * noise_scaler(er_lambda_t)) * denoised_d
|
||||
|
||||
if stage_used >= 3 and self.old_denoised_d is not None:
|
||||
# Stage 3: Third-order correction
|
||||
s_u = torch.sum((lambda_pos - er_lambda_s) / (scaled_pos + 1e-10)) * lambda_step_size
|
||||
|
||||
# Get sigma from two steps ago
|
||||
if timestep_id > 1:
|
||||
sigma_prev_prev = self.sigmas[timestep_id - 2].reshape(-1, 1, 1, 1)
|
||||
er_lambda_prev_prev = sigma_prev_prev
|
||||
else:
|
||||
er_lambda_prev_prev = er_lambda_prev
|
||||
|
||||
denoised_u = (denoised_d - self.old_denoised_d) / (((er_lambda_s - er_lambda_prev_prev) / 2) + 1e-10)
|
||||
x = x + alpha_t * ((dt ** 2) / 2 + s_u * noise_scaler(er_lambda_t)) * denoised_u
|
||||
|
||||
self.old_denoised_d = denoised_d
|
||||
|
||||
# Add stochastic noise
|
||||
if self.s_noise > 0:
|
||||
noise_term = (er_lambda_t ** 2 - er_lambda_s ** 2 * r ** 2).sqrt()
|
||||
noise_term = torch.nan_to_num(noise_term, nan=0.0)
|
||||
noise = torch.randn(*x.shape, dtype=torch.float32, device=torch.device("cpu"), generator=generator).to(x)
|
||||
x = x + alpha_t * noise * self.s_noise * noise_term
|
||||
|
||||
# Store current denoised for next iteration
|
||||
self.old_denoised = denoised
|
||||
self.step_index += 1
|
||||
|
||||
return x
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
|
||||
sample = (1 - sigma) * original_samples + sigma * noise
|
||||
return sample.type_as(noise)
|
||||
@@ -35,9 +35,7 @@ class FlowMatchSchedulerResMultistep():
|
||||
self.sigmas = torch.FloatTensor(sigmas)
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
#print(f"Timesteps: {self.timesteps}, Sigmas: {self.sigmas}")
|
||||
|
||||
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
|
||||
|
||||
def step(self, model_output, timestep, sample):
|
||||
if timestep.ndim == 2:
|
||||
@@ -48,14 +46,14 @@ class FlowMatchSchedulerResMultistep():
|
||||
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
|
||||
else:
|
||||
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
|
||||
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma
|
||||
if (timestep_id + 1 >= len(self.sigmas)).any():
|
||||
sigma_next = torch.tensor(0)
|
||||
else:
|
||||
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
|
||||
|
||||
x0_pred = (sample - sigma * model_output)
|
||||
|
||||
if sigma_next == 0 or self.prev_model_output is None:
|
||||
@@ -73,7 +71,7 @@ class FlowMatchSchedulerResMultistep():
|
||||
self.old_sigma_next = sigma_next
|
||||
self.prev_model_output = x0_pred
|
||||
return x
|
||||
|
||||
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
from diffusers.schedulers import UniPCMultistepScheduler
|
||||
import torch
|
||||
|
||||
|
||||
class ViBTScheduler(UniPCMultistepScheduler):
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**{**kwargs, "use_flow_sigmas": True})
|
||||
self.set_parameters()
|
||||
|
||||
def set_parameters(self, noise_scale=1.0, shift=5.0, seed=None):
|
||||
self.noise_scale = noise_scale
|
||||
self.config.flow_shift = shift
|
||||
|
||||
def step(self, model_output, timestep, sample, generator, **kwargs):
|
||||
delta_t = (
|
||||
max(self.timesteps[self.timesteps < timestep]) - timestep
|
||||
if any(self.timesteps < timestep)
|
||||
else -timestep - 1
|
||||
) / 1000
|
||||
|
||||
current_t = (timestep + 1) / 1000.0
|
||||
eta = (-delta_t * (current_t + delta_t) / current_t) ** 0.5
|
||||
|
||||
noise = torch.randn(
|
||||
sample.shape,
|
||||
generator=generator,
|
||||
device=torch.device("cpu"),
|
||||
dtype=sample.dtype,
|
||||
).to(sample.device)
|
||||
latents = sample + delta_t * model_output + eta * self.noise_scale * noise
|
||||
|
||||
return (latents,)
|
||||
|
||||
@classmethod
|
||||
def from_scheduler(
|
||||
cls, scheduler: UniPCMultistepScheduler, noise_scale=1.0, shift_gamma=5.0
|
||||
):
|
||||
obj = cls.__new__(cls)
|
||||
obj.__dict__ = scheduler.__dict__.copy()
|
||||
obj.set_parameters(noise_scale, shift_gamma)
|
||||
return obj
|
||||
+145
-45
@@ -5,6 +5,9 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from tqdm import tqdm
|
||||
from comfy.utils import ProgressBar
|
||||
from ..utils import print_memory, log
|
||||
from comfy import model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
import comfy.ops
|
||||
ops = comfy.ops.disable_weight_init
|
||||
|
||||
@@ -254,10 +257,11 @@ class Resample38(Resample):
|
||||
|
||||
class ResidualBlock(nn.Module):
|
||||
|
||||
def __init__(self, in_dim, out_dim, dropout=0.0):
|
||||
def __init__(self, in_dim, out_dim, dropout=0.0, cpu_cache=False):
|
||||
super().__init__()
|
||||
self.in_dim = in_dim
|
||||
self.out_dim = out_dim
|
||||
self.cpu_cache = cpu_cache
|
||||
|
||||
# layers
|
||||
self.residual = nn.Sequential(
|
||||
@@ -269,6 +273,12 @@ class ResidualBlock(nn.Module):
|
||||
if in_dim != out_dim else nn.Identity()
|
||||
|
||||
def forward(self, x, feat_cache=None, feat_idx=[0]):
|
||||
if self.cpu_cache:
|
||||
return self._forward_cpu_cache(x, feat_cache, feat_idx)
|
||||
else:
|
||||
return self._forward(x, feat_cache, feat_idx)
|
||||
|
||||
def _forward(self, x, feat_cache=None, feat_idx=[0]):
|
||||
h = self.shortcut(x)
|
||||
for layer in self.residual:
|
||||
if check_is_instance(layer, CausalConv3d) and feat_cache is not None:
|
||||
@@ -288,6 +298,26 @@ class ResidualBlock(nn.Module):
|
||||
x = layer(x)
|
||||
return x + h
|
||||
|
||||
def _forward_cpu_cache(self, x, feat_cache=None, feat_idx=[0]):
|
||||
h = self.shortcut(x)
|
||||
for layer in self.residual:
|
||||
if check_is_instance(layer, CausalConv3d) and feat_cache is not None:
|
||||
idx = feat_idx[0]
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
|
||||
cached_frame = feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device)
|
||||
cache_x = torch.cat([cached_frame, cache_x], dim=2)
|
||||
|
||||
prev_cache = feat_cache[idx].to(x.device) if feat_cache[idx] is not None else None
|
||||
|
||||
x = layer(x, prev_cache)
|
||||
|
||||
feat_cache[idx] = cache_x.to("cpu", non_blocking=True)
|
||||
feat_idx[0] += 1
|
||||
else:
|
||||
x = layer(x)
|
||||
return x + h
|
||||
|
||||
|
||||
class AttentionBlock(nn.Module):
|
||||
"""
|
||||
@@ -507,7 +537,8 @@ class Encoder3d(nn.Module):
|
||||
attn_scales=[],
|
||||
temperal_downsample=[True, True, False],
|
||||
dropout=0.0,
|
||||
pruning_rate=0.0):
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -515,6 +546,7 @@ class Encoder3d(nn.Module):
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_downsample = temperal_downsample
|
||||
self.cpu_cache = cpu_cache
|
||||
|
||||
# dimensions
|
||||
dims = [dim * u for u in [1] + dim_mult]
|
||||
@@ -529,7 +561,7 @@ class Encoder3d(nn.Module):
|
||||
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
|
||||
# residual (+attention) blocks
|
||||
for _ in range(num_res_blocks):
|
||||
downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
|
||||
downsamples.append(ResidualBlock(in_dim, out_dim, dropout, cpu_cache=cpu_cache))
|
||||
if scale in attn_scales:
|
||||
downsamples.append(AttentionBlock(out_dim))
|
||||
in_dim = out_dim
|
||||
@@ -543,9 +575,9 @@ class Encoder3d(nn.Module):
|
||||
self.downsamples = nn.Sequential(*downsamples)
|
||||
|
||||
# middle blocks
|
||||
self.middle = nn.Sequential(ResidualBlock(out_dim, out_dim, dropout),
|
||||
self.middle = nn.Sequential(ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
|
||||
AttentionBlock(out_dim),
|
||||
ResidualBlock(out_dim, out_dim, dropout))
|
||||
ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache))
|
||||
|
||||
# output blocks
|
||||
self.head = nn.Sequential(RMS_norm(out_dim, images=False), nn.SiLU(),
|
||||
@@ -612,7 +644,8 @@ class Encoder3d_38(nn.Module):
|
||||
attn_scales=[],
|
||||
temperal_downsample=[False, True, True],
|
||||
dropout=0.0,
|
||||
pruning_rate=0.0):
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -620,6 +653,7 @@ class Encoder3d_38(nn.Module):
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_downsample = temperal_downsample
|
||||
self.cpu_cache = cpu_cache
|
||||
|
||||
# dimensions
|
||||
dims = [dim * u for u in [1] + dim_mult]
|
||||
@@ -650,9 +684,9 @@ class Encoder3d_38(nn.Module):
|
||||
|
||||
# middle blocks
|
||||
self.middle = nn.Sequential(
|
||||
ResidualBlock(out_dim, out_dim, dropout),
|
||||
ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
|
||||
AttentionBlock(out_dim),
|
||||
ResidualBlock(out_dim, out_dim, dropout),
|
||||
ResidualBlock(out_dim, out_dim, dropout, cpu_cache=cpu_cache),
|
||||
)
|
||||
|
||||
# # output blocks
|
||||
@@ -730,7 +764,8 @@ class Decoder3d(nn.Module):
|
||||
attn_scales=[],
|
||||
temperal_upsample=[False, True, True],
|
||||
dropout=0.0,
|
||||
pruning_rate=0.0):
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -738,6 +773,7 @@ class Decoder3d(nn.Module):
|
||||
self.num_res_blocks = num_res_blocks
|
||||
self.attn_scales = attn_scales
|
||||
self.temperal_upsample = temperal_upsample
|
||||
self.cpu_cache = cpu_cache
|
||||
|
||||
# dimensions
|
||||
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
|
||||
@@ -748,9 +784,9 @@ class Decoder3d(nn.Module):
|
||||
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
|
||||
|
||||
# middle blocks
|
||||
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout),
|
||||
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache),
|
||||
AttentionBlock(dims[0]),
|
||||
ResidualBlock(dims[0], dims[0], dropout))
|
||||
ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache))
|
||||
|
||||
# upsample blocks
|
||||
upsamples = []
|
||||
@@ -759,7 +795,7 @@ class Decoder3d(nn.Module):
|
||||
if i == 1 or i == 2 or i == 3:
|
||||
in_dim = in_dim // 2
|
||||
for _ in range(num_res_blocks + 1):
|
||||
upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
|
||||
upsamples.append(ResidualBlock(in_dim, out_dim, dropout, cpu_cache=cpu_cache))
|
||||
if scale in attn_scales:
|
||||
upsamples.append(AttentionBlock(out_dim))
|
||||
in_dim = out_dim
|
||||
@@ -838,7 +874,8 @@ class Decoder3d_38(nn.Module):
|
||||
attn_scales=[],
|
||||
temperal_upsample=[False, True, True],
|
||||
dropout=0.0,
|
||||
pruning_rate=0.0):
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -855,9 +892,9 @@ class Decoder3d_38(nn.Module):
|
||||
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
|
||||
|
||||
# middle blocks
|
||||
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout),
|
||||
self.middle = nn.Sequential(ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache),
|
||||
AttentionBlock(dims[0]),
|
||||
ResidualBlock(dims[0], dims[0], dropout))
|
||||
ResidualBlock(dims[0], dims[0], dropout, cpu_cache=cpu_cache))
|
||||
|
||||
# upsample blocks
|
||||
upsamples = []
|
||||
@@ -951,7 +988,9 @@ class VideoVAE_(nn.Module):
|
||||
dropout=0.0,
|
||||
mean=None,
|
||||
inv_std=None,
|
||||
pruning_rate=0.0):
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False,
|
||||
verbose=False):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -962,14 +1001,15 @@ class VideoVAE_(nn.Module):
|
||||
self.temperal_upsample = temperal_downsample[::-1]
|
||||
self.mean = mean
|
||||
self.inv_std = inv_std
|
||||
self.verbose = verbose
|
||||
|
||||
# modules
|
||||
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
|
||||
attn_scales, self.temperal_downsample, dropout, pruning_rate)
|
||||
attn_scales, self.temperal_downsample, dropout, pruning_rate, cpu_cache=cpu_cache)
|
||||
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
|
||||
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
|
||||
self.decoder = Decoder3d(dim, z_dim, dim_mult, num_res_blocks,
|
||||
attn_scales, self.temperal_upsample, dropout, pruning_rate)
|
||||
attn_scales, self.temperal_upsample, dropout, pruning_rate, cpu_cache=cpu_cache)
|
||||
|
||||
def forward(self, x):
|
||||
mu, log_var = self.encode(x)
|
||||
@@ -1016,10 +1056,15 @@ class VideoVAE_(nn.Module):
|
||||
def encode(self, x, pbar=True, sample=False):
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
input_shape = x.shape
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
|
||||
for i in range(iter_):
|
||||
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
|
||||
self._enc_conv_idx = [0]
|
||||
if i == 0:
|
||||
out = self.encoder(x[:, :, :1, :, :],
|
||||
@@ -1042,15 +1087,22 @@ class VideoVAE_(nn.Module):
|
||||
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
|
||||
eps = torch.randn_like(std)
|
||||
return mu + std * eps
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE encode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return mu
|
||||
|
||||
|
||||
#modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames
|
||||
def decode_2(self, z):
|
||||
# z: [b,c,t,h,w]
|
||||
|
||||
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
|
||||
|
||||
iter_ = z.shape[2]
|
||||
z_head=z[:,:,:-1,:,:]
|
||||
z_tail=z[:,:,-1,:,:].unsqueeze(2)
|
||||
@@ -1065,12 +1117,12 @@ class VideoVAE_(nn.Module):
|
||||
out_ = self.decoder(x[:, :, -1, :, :].unsqueeze(2),
|
||||
feat_cache=None,
|
||||
feat_idx=self._conv_idx)
|
||||
out = torch.cat([out, out_], 2) # may add tensor offload
|
||||
out = torch.cat([out, out_], 2)
|
||||
else:
|
||||
out_ = self.decoder(x[:, :, i:i + 1, :, :],
|
||||
feat_cache=self._feat_map,
|
||||
feat_idx=self._conv_idx)
|
||||
out = torch.cat([out, out_], 2) # may add tensor offload
|
||||
out = torch.cat([out, out_], 2)
|
||||
self.clear_cache()
|
||||
return out
|
||||
|
||||
@@ -1079,11 +1131,16 @@ class VideoVAE_(nn.Module):
|
||||
def decode(self, z, pbar=True):
|
||||
# z: [b,c,t,h,w]
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
input_shape = z.shape
|
||||
iter_ = z.shape[2]
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
x = self.conv2(z)
|
||||
for i in range(iter_):
|
||||
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
|
||||
self._conv_idx = [0]
|
||||
if i == 0:
|
||||
out = self.decoder(x[:, :, i:i + 1, :, :],
|
||||
@@ -1093,13 +1150,20 @@ class VideoVAE_(nn.Module):
|
||||
out_ = self.decoder(x[:, :, i:i + 1, :, :],
|
||||
feat_cache=self._feat_map,
|
||||
feat_idx=self._conv_idx)
|
||||
out = torch.cat([out, out_], 2) # may add tensor offload
|
||||
out = torch.cat([out, out_], 2)
|
||||
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
if pbar:
|
||||
pbar.update_absolute(0)
|
||||
self.clear_cache()
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return out
|
||||
|
||||
def reparameterize(self, mu, log_var):
|
||||
@@ -1126,11 +1190,12 @@ class VideoVAE_(nn.Module):
|
||||
|
||||
class WanVideoVAE(nn.Module):
|
||||
|
||||
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0):
|
||||
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False, verbose=False):
|
||||
super().__init__()
|
||||
|
||||
self.dtype = dtype
|
||||
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
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
|
||||
@@ -1144,7 +1209,7 @@ class WanVideoVAE(nn.Module):
|
||||
self.z_dim = z_dim
|
||||
|
||||
# init model
|
||||
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate).eval().requires_grad_(False)
|
||||
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache, verbose=self.verbose).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 8
|
||||
|
||||
|
||||
@@ -1170,7 +1235,7 @@ class WanVideoVAE(nn.Module):
|
||||
return mask
|
||||
|
||||
|
||||
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True):
|
||||
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, end_=False, pbar=True):
|
||||
_, _, T, H, W = hidden_states.shape
|
||||
size_h, size_w = tile_size
|
||||
stride_h, stride_w = tile_stride
|
||||
@@ -1187,14 +1252,20 @@ class WanVideoVAE(nn.Module):
|
||||
data_device = "cpu"
|
||||
computation_device = device
|
||||
|
||||
out_T = T * 4 - 3
|
||||
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
weight, values = None, None
|
||||
if pbar:
|
||||
pbar = ProgressBar(len(tasks))
|
||||
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
|
||||
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
|
||||
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
|
||||
if end_:
|
||||
hidden_states_batch = self.model.decode_2(hidden_states_batch).to(data_device)
|
||||
else:
|
||||
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
|
||||
|
||||
if weight is None:
|
||||
weight = torch.zeros((1, 1, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
if values is None:
|
||||
values = torch.zeros((1, 3, hidden_states_batch.shape[2], H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
|
||||
mask = self.build_mask(
|
||||
hidden_states_batch,
|
||||
@@ -1227,7 +1298,7 @@ class WanVideoVAE(nn.Module):
|
||||
|
||||
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False, pbar=True):
|
||||
_, _, T, H, W = video.shape
|
||||
|
||||
|
||||
if tile_size is None and tile_stride is None:
|
||||
size_h, size_w = H //2, W // 2
|
||||
stride_h, stride_w = size_h // 2, size_w // 2
|
||||
@@ -1338,7 +1409,7 @@ class WanVideoVAE(nn.Module):
|
||||
for hidden_state in hidden_states:
|
||||
hidden_state = hidden_state.unsqueeze(0)
|
||||
if tiled:
|
||||
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar)
|
||||
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, end_=end_, pbar=pbar)
|
||||
else:
|
||||
if end_:
|
||||
video = self.double_decode(hidden_state, device)
|
||||
@@ -1363,7 +1434,9 @@ class VideoVAE38_(VideoVAE_):
|
||||
dtype=torch.bfloat16,
|
||||
mean=None,
|
||||
inv_std=None,
|
||||
pruning_rate=0.0):
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False,
|
||||
verbose=False):
|
||||
super(VideoVAE_, self).__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -1375,24 +1448,30 @@ class VideoVAE38_(VideoVAE_):
|
||||
self.dtype = dtype
|
||||
self.mean = mean
|
||||
self.inv_std = inv_std
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
|
||||
# modules
|
||||
self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks,
|
||||
attn_scales, self.temperal_downsample, dropout, pruning_rate)
|
||||
attn_scales, self.temperal_downsample, dropout, pruning_rate, cpu_cache=cpu_cache)
|
||||
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
|
||||
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
|
||||
self.decoder = Decoder3d_38(dec_dim, z_dim, dim_mult, num_res_blocks,
|
||||
attn_scales, self.temperal_upsample, dropout, pruning_rate)
|
||||
|
||||
attn_scales, self.temperal_upsample, dropout, pruning_rate, cpu_cache=cpu_cache)
|
||||
|
||||
def encode(self, x, pbar=True, sample=False):
|
||||
input_shape = x.shape
|
||||
self.clear_cache()
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
x = patchify(x, patch_size=2)
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
for i in range(iter_):
|
||||
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
|
||||
self._enc_conv_idx = [0]
|
||||
if i == 0:
|
||||
out = self.encoder(x[:, :, :1, :, :],
|
||||
@@ -1408,18 +1487,30 @@ class VideoVAE38_(VideoVAE_):
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
self.clear_cache()
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return mu
|
||||
|
||||
|
||||
def decode(self, z, pbar=True):
|
||||
self.clear_cache()
|
||||
input_shape = z.shape
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
|
||||
|
||||
iter_ = z.shape[2]
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
x = self.conv2(z)
|
||||
for i in range(iter_):
|
||||
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
|
||||
self._conv_idx = [0]
|
||||
if i == 0:
|
||||
out = self.decoder(x[:, :, i:i + 1, :, :],
|
||||
@@ -1435,12 +1526,19 @@ class VideoVAE38_(VideoVAE_):
|
||||
pbar.update(1)
|
||||
out = unpatchify(out, patch_size=2)
|
||||
self.clear_cache()
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return out
|
||||
|
||||
|
||||
class WanVideoVAE38(WanVideoVAE):
|
||||
|
||||
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0):
|
||||
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False, verbose=False):
|
||||
super(WanVideoVAE, self).__init__()
|
||||
|
||||
mean = [
|
||||
@@ -1463,7 +1561,9 @@ class WanVideoVAE38(WanVideoVAE):
|
||||
self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1)
|
||||
self.dtype = dtype
|
||||
self.z_dim = z_dim
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
|
||||
# init model
|
||||
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 16
|
||||
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache, verbose=verbose).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 16
|
||||
|
||||
Reference in New Issue
Block a user