36 Commits
Author SHA1 Message Date
kijai 088128b224 Don't count .disabled as duplicate 2026-05-24 16:07:12 +03:00
Jukka Seppänen 126819826a Merge pull request #1948 from haosenwang1018/fix/bare-excepts
fix: replace 47 bare excepts with except Exception
2026-05-24 16:02:36 +03:00
kijai 8f1804bf72 Fix multitalk_audio_stride init 2026-05-24 15:59:40 +03:00
kijai 5437b016e3 Initial LongCatAvatar 1.5 support 2026-05-23 23:13:28 +03:00
kijai d18cdb1859 Fix offload on interrupt 2026-05-05 14:00:56 +03:00
haosenwang1018 0d78230336 fix: replace 47 bare excepts with except Exception
Bare except catches KeyboardInterrupt and SystemExit, masking real errors.
2026-02-25 06:04:52 +00:00
Jukka Seppänen df8f3e49da Delete .github/FUNDING.yml 2026-02-22 15:18:05 +02:00
Jukka Seppänen 86ad93d616 Merge pull request #1840 from little6neko/main
feat(wanvideo): add manual start reference support for WanAnimate loop
2026-02-16 15:44:15 +02:00
Jukka Seppänen 309491b269 Merge pull request #1908 from jjdejong/patch-1
Implement conditional offloading for diffusion model
2026-02-16 15:35:45 +02:00
kijai 06122f1e9d Rope offset should be disabled by default 2026-02-16 15:32:54 +02:00
kijai 5d36631795 I2V cross_attn fixes 2026-02-16 15:31:20 +02:00
kijai 3d7b49e2df Move string_to_seed import 2026-02-02 17:02:05 +02:00
kijai e091c4a774 version 1.4.7 2026-02-02 01:43:38 +02:00
kijai 60a387579e Update wanvideo_2_1_14B_I2V_SkyReelsV3_TalkingAvatar_example_01.json 2026-02-01 15:16:33 +02:00
kijai 9ae1a4f2c9 Update multitalk_loop.py 2026-02-01 15:15:37 +02:00
kijai f55b7b3d89 Update wanvideo_2_1_14B_I2V_SkyReelsV3_TalkingAvatar_example_01.json 2026-02-01 15:15:33 +02:00
kijai e4e7f413f7 Make compatible with latest ComfyUI version 2026-02-01 15:15:27 +02:00
kijai e21fe20a4d Support SkyReels TalkingAvatar (A2V) 2026-01-31 23:24:29 +02:00
kijai 2c5a04cc63 Init image_cond_mask 2026-01-30 18:59:41 +02:00
kijai d00abe52d7 version 1.4.6 2026-01-28 02:09:50 +02:00
kijai 2952f5d9dd Show T5 loading progress bar 2026-01-26 15:53:48 +02:00
kijai 339e0fec81 Make NAG inplace application optional 2026-01-23 18:35:30 +02:00
kijai 2c2a6e1889 RoPE frequency offset option for storymem
I'm not 100% sure on this, I initially tested this when I noticed the original code doesn't, but it's described in the paper... now I see the original code has added it too so it seems to be the intended way to use it.
2026-01-23 14:31:04 +02:00
kijai 8640bfad52 VRAM optimizations
Minor for almost everything, major for multitalk when using masks
2026-01-22 19:46:52 +02:00
kijai 3b4a711a40 Reduce NAG memory usage 2026-01-22 17:43:34 +02:00
Jean J. de Jong b99eac73da Implement conditional offloading for diffusion model
Add condition to skip offloading if load_device is main_device.
2026-01-21 17:33:58 +01:00
kijai 9f9a8e71c2 Revert " Fix: Fix torch.arange bounds error in context window processing"
This reverts commit 707bbcd72b.
2026-01-17 20:13:43 +02:00
小六妞儿 0fbcbed06a Merge branch 'kijai:main' into main 2026-01-16 15:52:00 +08:00
Jukka Seppänen f2e2e2550a Merge pull request #1877 from YFJack/patch-1
Fix: Fix torch.arange bounds error in context window processing
2026-01-15 17:06:31 +02:00
小六妞儿 6f9832ed47 Merge branch 'kijai:main' into main 2026-01-15 14:05:52 +08:00
Yifu Wang 707bbcd72b Fix: Fix torch.arange bounds error in context window processing
Fixed "upper bound and lower bound inconsistent with step sign" error
       when using WanVideoContextOptions with pose latents. Changed end index
       from c[-1] to c[-1] + 1 to properly include the last frame in the range.
2026-01-09 11:36:53 +08:00
Jukka Seppänen 855d103ee6 Merge pull request #1868 from vantagewithai/main
longcat avatar GGUF support.
2026-01-08 13:36:43 +02:00
kijai 64191921d4 Squashed commit of the following:
commit fdb23dec7d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jan 5 22:11:04 2026 +0200

    Update model.py

commit 07d7d8ca8e
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jan 5 22:10:02 2026 +0200

    remove prints

commit 01869d4bf5
Merge: 55c6720 bf1d77f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Jan 5 18:47:48 2026 +0200

    Merge branch 'main' into longvie2

commit 55c672028b
Merge: b551ec9 be41f67
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 29 15:39:43 2025 +0200

    Merge branch 'main' into longvie2

commit b551ec9e31
Merge: 9f019d7 19bcee6
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 29 15:03:53 2025 +0200

    Merge branch 'main' into longvie2

commit 9f019d7dfb
Merge: fc5322f c5d3fb4
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 23:40:25 2025 +0200

    Merge branch 'main' into longvie2

commit fc5322fae4
Merge: 222fc70 e75f814
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 22:04:15 2025 +0200

    Merge branch 'main' into longvie2

commit 222fc70eb7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 17:18:55 2025 +0200

    Update nodes.py

commit 8509236da1
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 23 14:20:18 2025 +0200

    init
2026-01-05 22:11:20 +02:00
Vantage with AI 576f065073 Update tensor name condition for multitalk audio proj, to support longcat video avatar GGUF 2026-01-04 21:13:54 +05:30
小六妞儿 58c1bcb7ce Merge branch 'kijai:main' into main 2025-12-28 12:59:47 +08:00
小六妞儿 64cbd28e00 feat(wanvideo): add manual start reference support for WanAnimate loop 2025-12-27 16:04:00 +08:00
25 changed files with 4858 additions and 269 deletions
-1
View File
@@ -1 +0,0 @@
github: [kijai]
+190 -2
View File
@@ -1,4 +1,5 @@
import torch import torch
import torch.nn.functional as F
from ..utils import log from ..utils import log
import comfy.model_management as mm import comfy.model_management as mm
from comfy_api.latest import io from comfy_api.latest import io
@@ -24,6 +25,8 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
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.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("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"), io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"),
io.Custom("IMAGE").Input("prev_images", optional=True, tooltip="LongCat-Avatar-1.5: decoded frames from the previous segment. When provided together with `vae`, the trailing `overlap` frames are re-encoded through the VAE and used as the overlap conditioning (matches v1.5's use_vcond=False behavior). Leave disconnected for v1.0."),
io.Custom("WANVAE").Input("vae", optional=True, tooltip="LongCat-Avatar-1.5: VAE used to re-encode `prev_images` for the overlap region. Only used when `prev_images` is also provided."),
], ],
outputs=[ outputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"), io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
@@ -32,7 +35,7 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
) )
@classmethod @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: 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, prev_images=None, vae=None) -> io.NodeOutput:
new_audio_embed = audio_embeds.copy() new_audio_embed = audio_embeds.copy()
@@ -55,6 +58,19 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
prev_samples = prev_latents["samples"].clone() prev_samples = prev_latents["samples"].clone()
if overlap != 0: if overlap != 0:
latent_overlap = (overlap - 1) // 4 + 1 latent_overlap = (overlap - 1) // 4 + 1
if prev_images is not None and vae is not None:
# LongCat-Avatar-1.5 path: re-encodes instead of just slicing
img = prev_images[-overlap:]
if img.shape[-1] == 4:
img = img[..., :3]
img = img.to(vae.dtype).to(device) * 2.0 - 1.0
img = img.permute(3, 0, 1, 2).unsqueeze(0).contiguous() # [T, H, W, C] -> [B, C, T, H, W]
vae.to(device)
prev_samples = vae.encode(img, device=device).to(prev_samples)
vae.to(offload_device)
mm.soft_empty_cache()
log.info(f"Re-encoded {overlap} overlap frames -> latent shape {tuple(prev_samples.shape)}")
else:
prev_samples = prev_samples[:, :, -latent_overlap:] prev_samples = prev_samples[:, :, -latent_overlap:]
ref_sample = None ref_sample = None
@@ -65,7 +81,7 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
new_latent_frames = (num_frames - 1) // 4 + 1 new_latent_frames = (num_frames - 1) // 4 + 1
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1]) target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
audio_stride = 2 audio_stride = new_audio_embed.get("audio_stride", 2)
indices = torch.arange(2 * 2 + 1) - 2 indices = torch.arange(2 * 2 + 1) - 2
if frames_processed == 0: if frames_processed == 0:
@@ -112,9 +128,181 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
return io.NodeOutput(embeds, samples_slice) return io.NodeOutput(embeds, samples_slice)
class LongCatAvatarWhisperEmbeds:
"""Audio embeds for LongCat-Video-Avatar-1.5 (Whisper-large-v3).
Produces a MULTITALK_EMBEDS dict whose audio_features are shaped [T, 5, 1280]
(5 grouped Whisper layers, 1280-d hidden state), matching the audio stream
the v1.5 AudioProjModel expects. audio_stride is set to 1 to signal v1.5
timing to the consumer nodes (vs. 2 for the v1.0 wav2vec2 path).
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"whisper_model": ("WHISPERMODEL",),
"audio_1": ("AUDIO",),
"normalize_loudness": ("BOOLEAN", {"default": True, "tooltip": "Normalize audio loudness to -23 LUFS before encoding (matches the v1.5 reference pipeline)"}),
"num_frames": ("INT", {"default": 93, "min": 1, "max": 10000, "step": 1, "tooltip": "Total frame count to generate; bounds how much audio is consumed"}),
"fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 60.0, "step": 0.1, "tooltip": "Target video fps. LongCat-Video-Avatar-1.5 is trained at 25 fps."}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done"}),
"multi_audio_type": (["para", "add"], {"default": "para", "tooltip": "'para' overlays speakers in parallel (equal length); 'add' concatenates speakers sequentially with silence padding"}),
},
"optional": {
"audio_2": ("AUDIO",),
"audio_3": ("AUDIO",),
"audio_4": ("AUDIO",),
"ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space, one per speaker"}),
},
}
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", "INT",)
RETURN_NAMES = ("multitalk_embeds", "audio", "num_frames",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, whisper_model, audio_1, normalize_loudness, num_frames, fps,
audio_scale, audio_cfg_scale, multi_audio_type,
audio_2=None, audio_3=None, audio_4=None, ref_target_masks=None):
import torchaudio
import numpy as np
from ..multitalk.nodes import loudness_norm
model = whisper_model["model"]
feature_extractor = whisper_model["feature_extractor"]
dtype = whisper_model["dtype"]
sr = 16000
MEL_CHUNK = 750 * 640 # 480000 samples = 30s at 16kHz; matches Whisper's chunk_length
ENC_CHUNK = 3000 # encoder window in mel frames
ENC_FPS = 50 # whisper encoder output frames per second
def linear_interp(features, output_len):
features = features.transpose(1, 2) # [B, D, T]
out = F.interpolate(features, size=output_len, align_corners=True, mode='linear')
return out.transpose(1, 2)
audio_inputs = [a for a in [audio_1, audio_2, audio_3, audio_4] if a is not None]
audio_features_list = []
seq_lengths = []
audio_outputs = []
end_time = num_frames / float(fps)
end_sample = int(end_time * sr)
for audio in audio_inputs:
audio_input = audio["waveform"]
sample_rate = audio["sample_rate"]
if sample_rate != sr:
audio_input = torchaudio.functional.resample(audio_input, sample_rate, sr)
audio_input = audio_input[0][0]
audio_segment = audio_input[:end_sample].cpu().numpy().astype(np.float32)
if normalize_loudness:
audio_segment = loudness_norm(audio_segment, sr=sr)
audio_duration = len(audio_segment) / sr
video_length = int(audio_duration * fps)
if video_length < 1:
continue
mel_chunks = []
for i in range(0, len(audio_segment), MEL_CHUNK):
mel = feature_extractor(audio_segment[i:i + MEL_CHUNK], sampling_rate=sr,
return_tensors="pt").input_features
mel_chunks.append(mel)
mel_features = torch.cat(mel_chunks, dim=-1).to(device=device, dtype=dtype)
model.to(device)
enc_chunks = []
with torch.no_grad():
for i in range(0, mel_features.shape[-1], ENC_CHUNK):
chunk = mel_features[:, :, i:i + ENC_CHUNK]
chunk_hs = model.encoder(chunk, output_hidden_states=True).hidden_states
enc_chunks.append(torch.stack(chunk_hs, dim=2)) # [1, T_enc, n_layers+1, D]
model.to(offload_device)
audio_prompts = torch.cat(enc_chunks, dim=1)
audio_prompts = audio_prompts[:, :video_length * 2]
feat0 = linear_interp(audio_prompts[:, :, 0:8].mean(dim=2), video_length)
feat1 = linear_interp(audio_prompts[:, :, 8:16].mean(dim=2), video_length)
feat2 = linear_interp(audio_prompts[:, :, 16:24].mean(dim=2), video_length)
feat3 = linear_interp(audio_prompts[:, :, 24:32].mean(dim=2), video_length)
feat4 = linear_interp(audio_prompts[:, :, 32], video_length)
audio_emb = torch.stack([feat0, feat1, feat2, feat3, feat4], dim=2)[0] # [T, 5, 1280]
audio_features_list.append(audio_emb.cpu().detach())
seq_lengths.append(audio_emb.shape[0])
waveform_tensor = torch.from_numpy(audio_segment).float().unsqueeze(0).unsqueeze(0)
audio_outputs.append({"waveform": waveform_tensor, "sample_rate": sr})
if len(audio_features_list) == 0:
raise RuntimeError("No valid Whisper audio embeddings extracted, please check inputs")
if len(audio_features_list) > 1:
if multi_audio_type == "para":
max_len = max(seq_lengths)
padded = []
for emb in audio_features_list:
if emb.shape[0] < max_len:
pad = torch.zeros(max_len - emb.shape[0], *emb.shape[1:], dtype=emb.dtype)
emb = torch.cat([emb, pad], dim=0)
padded.append(emb)
audio_features_list = padded
else: # "add"
total_len = sum(seq_lengths)
full_list = []
offset = 0
for emb, length in zip(audio_features_list, seq_lengths):
full = torch.zeros(total_len, *emb.shape[1:], dtype=emb.dtype)
full[offset:offset + length] = emb
full_list.append(full)
offset += length
audio_features_list = full_list
multitalk_embeds = {
"audio_features": audio_features_list,
"audio_scale": audio_scale,
"audio_cfg_scale": audio_cfg_scale,
"ref_target_masks": ref_target_masks,
"audio_stride": 1,
"audio_encoder_type": "whisper",
}
if len(audio_outputs) == 1:
out_audio = audio_outputs[0]
elif multi_audio_type == "para":
max_len = max(a["waveform"].shape[-1] for a in audio_outputs)
mixed = torch.zeros(1, 1, max_len, dtype=audio_outputs[0]["waveform"].dtype)
for a in audio_outputs:
w = a["waveform"]
if w.shape[-1] < max_len:
w = F.pad(w, (0, max_len - w.shape[-1]))
mixed += w
out_audio = {"waveform": mixed, "sample_rate": sr}
else:
total_len = sum(a["waveform"].shape[-1] for a in audio_outputs)
mixed = torch.zeros(1, 1, total_len, dtype=audio_outputs[0]["waveform"].dtype)
offset = 0
for a in audio_outputs:
w = a["waveform"]
mixed[:, :, offset:offset + w.shape[-1]] += w
offset += w.shape[-1]
out_audio = {"waveform": mixed, "sample_rate": sr}
return (multitalk_embeds, out_audio, num_frames)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds, "WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
"LongCatAvatarWhisperEmbeds": LongCatAvatarWhisperEmbeds,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds", "WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
"LongCatAvatarWhisperEmbeds": "LongCat Avatar Whisper Embeds (v1.5)",
} }
+199
View File
@@ -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)
+88
View File
@@ -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",
}
+1 -1
View File
@@ -31,7 +31,7 @@ def check_jit_script_function():
f" Qualified name: {qualname}\n" f" Qualified name: {qualname}\n"
f" Defined in: {code_file}:{code_line}\n" f" Defined in: {code_file}:{code_line}\n"
f"This may cause issues with the NLF model.") f"This may cause issues with the NLF model.")
except: except Exception:
log.warning("--------------------------------") log.warning("--------------------------------")
log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, " 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.") f"this has been modified by another custom node. This may cause issues with the NLF model.")
+2 -1
View File
@@ -6,7 +6,7 @@ try:
for dir_path in duplicate_dirs: for dir_path in duplicate_dirs:
warning_msg += f" - {color_text(dir_path, 'yellow')}\n" warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red")) log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
except: except Exception:
pass pass
from .utils import log from .utils import log
@@ -49,6 +49,7 @@ OPTIONAL_MODULES = [
(".WanMove.nodes", "WanMove"), (".WanMove.nodes", "WanMove"),
(".SCAIL.nodes", "SCAIL"), (".SCAIL.nodes", "SCAIL"),
(".LongCat.nodes", "LongCat"), (".LongCat.nodes", "LongCat"),
(".LongVie2.nodes", "LongVie2"),
] ]
def register_nodes(module_path: str, name: str, optional: bool) -> None: def register_nodes(module_path: str, name: str, optional: bool) -> None:
+1 -1
View File
@@ -56,7 +56,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
module_prefix = module_prefix.replace("_orig_mod.", "") module_prefix = module_prefix.replace("_orig_mod.", "")
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert) _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert)
if isinstance(module, nn.Linear) and "loras" not in module_prefix and name not in modules_to_not_convert: if isinstance(module, nn.Linear) and "loras" not in module_prefix and "dual_controller" not in module_prefix and name not in modules_to_not_convert:
weight_key = module_prefix + "weight" weight_key = module_prefix + "weight"
if weight_key not in state_dict: if weight_key not in state_dict:
continue continue
File diff suppressed because one or more lines are too long
+1 -1
View File
@@ -167,7 +167,7 @@ class FantasyTalkingWav2VecEmbeds:
try: try:
audio_segment = audio_input[start_sample:end_sample] audio_segment = audio_input[start_sample:end_sample]
except: except Exception:
audio_segment = audio_input audio_segment = audio_input
print("audio_segment.shape", audio_segment.shape) print("audio_segment.shape", audio_segment.shape)
+1 -1
View File
@@ -85,7 +85,7 @@ def get_previewer(device, latent_format):
taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device) taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device)
previewer = TAESDPreviewerImpl(taesd) previewer = TAESDPreviewerImpl(taesd)
previewer = WrappedPreviewer(previewer, rate=16) previewer = WrappedPreviewer(previewer, rate=16)
except: except Exception:
log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors") log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors")
log.info("Using Latent2RGB previewer instead.") log.info("Using Latent2RGB previewer instead.")
method = LatentPreviewMethod.Latent2RGB method = LatentPreviewMethod.Latent2RGB
+40 -35
View File
@@ -42,41 +42,43 @@ def rotate_half(x):
x = torch.stack((-x2, x1), dim=-1) x = torch.stack((-x2, x1), dim=-1)
return rearrange(x, "... d r -> ... (d r)") return rearrange(x, "... d r -> ... (d r)")
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', attn_bias=None): def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, split_num=4):
ref_k = ref_k.to(visual_q.dtype).to(visual_q.device)
scale = 1.0 / visual_q.shape[-1] ** 0.5 scale = 1.0 / visual_q.shape[-1] ** 0.5
visual_q = visual_q * scale visual_q = visual_q.transpose(1, 2) * scale
visual_q = visual_q.transpose(1, 2)
ref_k = ref_k.transpose(1, 2)
attn = visual_q @ ref_k.transpose(-2, -1)
if attn_bias is not None:
attn = attn + attn_bias
x_ref_attn_map_source = attn.softmax(-1) # B, H, x_seqlens, ref_seqlens
B, H, x_seqlens, K = visual_q.shape
x_ref_attn_maps = [] x_ref_attn_maps = []
ref_target_masks = ref_target_masks.to(visual_q.dtype)
x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype)
for class_idx, ref_target_mask in enumerate(ref_target_masks): for class_idx, ref_target_mask in enumerate(ref_target_masks):
ref_target_mask = ref_target_mask[None, None, None, ...] ref_target_mask = ref_target_mask.view(1, 1, 1, -1)
x_ref_attnmap = x_ref_attn_map_source * ref_target_mask
x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens
x_ref_attnmap = x_ref_attnmap.permute(0, 2, 1) # B, x_seqlens, H
if mode == 'mean': x_ref_attnmap = torch.zeros(B, H, x_seqlens, device=visual_q.device, dtype=visual_q.dtype)
x_ref_attnmap = x_ref_attnmap.mean(-1) # B, x_seqlens chunk_size = min(max(x_seqlens // split_num, 1), x_seqlens)
elif mode == 'max':
x_ref_attnmap = x_ref_attnmap.max(-1) # B, x_seqlens
for i in range(0, x_seqlens, chunk_size):
end_i = min(i + chunk_size, x_seqlens)
attn_chunk = visual_q[:, :, i:end_i] @ ref_k.permute(0, 2, 3, 1) # B, H, chunk, ref_seqlens
# Apply softmax
attn_max = attn_chunk.max(dim=-1, keepdim=True).values
attn_chunk = (attn_chunk - attn_max).exp()
attn_sum = attn_chunk.sum(dim=-1, keepdim=True)
attn_chunk = attn_chunk / (attn_sum + 1e-8)
# Apply mask and sum
masked_attn = attn_chunk * ref_target_mask
x_ref_attnmap[:, :, i:end_i] = masked_attn.sum(-1) / (ref_target_mask.sum() + 1e-8)
del attn_chunk, masked_attn
# Average across heads
x_ref_attnmap = x_ref_attnmap.mean(dim=1) # B, x_seqlens
x_ref_attn_maps.append(x_ref_attnmap) x_ref_attn_maps.append(x_ref_attnmap)
del attn, x_ref_attn_map_source del visual_q, ref_k
return torch.concat(x_ref_attn_maps, dim=0) return torch.cat(x_ref_attn_maps, dim=0)
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2): def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2):
"""Args: """Args:
@@ -129,27 +131,30 @@ class RotaryPositionalEmbedding1D(nn.Module):
query with the same shape as input. query with the same shape as input.
""" """
freqs_cis = self.precompute_freqs_cis_1d(pos_indices) freqs_cis = self.precompute_freqs_cis_1d(pos_indices)
in_dtype = x.dtype
x_ = x.float() x = x.float()
freqs_cis = freqs_cis.float().to(x.device) freqs_cis = freqs_cis.float().to(x.device)
cos, sin = freqs_cis.cos(), freqs_cis.sin() cos = rearrange(freqs_cis.cos(), 'n d -> 1 1 n d')
cos, sin = rearrange(cos, 'n d -> 1 1 n d'), rearrange(sin, 'n d -> 1 1 n d') sin = rearrange(freqs_cis.sin(), 'n d -> 1 1 n d')
x_ = (x_ * cos) + (rotate_half(x_) * sin)
return x_.type_as(x) # In-place rotation to save memory
x_rotated = rotate_half(x)
x.mul_(cos).add_(x_rotated * sin)
return x.to(in_dtype)
class AudioProjModel(nn.Module): class AudioProjModel(nn.Module):
def __init__( def __init__(
self, self,
seq_len=5, seq_len=5,
seq_len_vf=12, seq_len_vf=8,
blocks=12, blocks=12,
channels=768, channels=768,
intermediate_dim=512, intermediate_dim=512,
output_dim=768, output_dim=768,
context_tokens=32, context_tokens=32,
norm_output_audio=False, norm_output_audio=True,
): ):
super().__init__() super().__init__()
@@ -273,9 +278,9 @@ class SingleStreamMultiAttention(SingleStreamAttention):
def __init__( def __init__(
self, self,
dim: int, dim: int,
encoder_hidden_states_dim: int,
num_heads: int, num_heads: int,
qkv_bias: bool, qkv_bias: bool = True,
encoder_hidden_states_dim: int = 768,
class_range: int = 24, class_range: int = 24,
class_interval: int = 4, class_interval: int = 4,
attention_mode: str = 'sdpa', attention_mode: str = 'sdpa',
+108 -22
View File
@@ -6,7 +6,7 @@ import numpy as np
from ..latent_preview import prepare_callback from ..latent_preview import prepare_callback
from ..wanvideo.schedulers import get_scheduler from ..wanvideo.schedulers import get_scheduler
from .multitalk import timestep_transform, add_noise from .multitalk import timestep_transform, add_noise
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap, match_and_blend_colors
from comfy.utils import load_torch_file from comfy.utils import load_torch_file
from ..nodes_model_loading import load_weights from ..nodes_model_loading import load_weights
from ..HuMo.nodes import get_audio_emb_window from ..HuMo.nodes import get_audio_emb_window
@@ -48,7 +48,13 @@ def multitalk_loop(self, **kwargs):
mode = image_embeds.get("multitalk_mode", "multitalk") mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto": if mode == "auto":
mode = transformer.multitalk_model_type.lower() mode = transformer.multitalk_model_type.lower()
elif mode == "skyreelsv3":
num_pseudo_frames = 5
pseudo_frames = reference_keyframes = None
keyframe_index = 0
reference_video = image_embeds.get("reference_video", None)
log.info(f"Multitalk mode: {mode}") log.info(f"Multitalk mode: {mode}")
drop_frames = image_embeds.get("drop_frames", 0)
cond_frame = None cond_frame = None
offload = image_embeds.get("force_offload", False) offload = image_embeds.get("force_offload", False)
offloaded = False offloaded = False
@@ -62,7 +68,9 @@ def multitalk_loop(self, **kwargs):
motion_frame = image_embeds.get("motion_frame", 25) motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None) target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None) target_h = image_embeds.get("target_h", None)
original_images = cond_image = image_embeds.get("multitalk_start_image", None) original_images = image_embeds.get("multitalk_start_image", None)
cond_image = original_images.clone() if original_images is not None else None
original_color_reference = cond_image.clone() if cond_image is not None else None
if original_images is None: if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device) original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
@@ -94,7 +102,6 @@ def multitalk_loop(self, **kwargs):
audio_embedding = multitalk_audio_embeds audio_embedding = multitalk_audio_embeds
human_num = len(audio_embedding) human_num = len(audio_embedding)
audio_embs = None audio_embs = None
cond_frame = None
uni3c_data = None uni3c_data = None
if uni3c_embeds is not None: if uni3c_embeds is not None:
@@ -106,13 +113,60 @@ def multitalk_loop(self, **kwargs):
try: try:
silence_path = os.path.join(script_directory, "encoded_silence.safetensors") silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype) encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except: except Exception:
log.warning("No encoded silence file found, padding with end of audio embedding instead.") log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0]) total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame) + 1 estimated_iterations = total_frames // (frame_num - motion_frame - drop_frames) + 1
callback = prepare_callback(patcher, estimated_iterations) callback = prepare_callback(patcher, estimated_iterations)
# If reference_video is provided, extract keyframes from it
if mode == "skyreelsv3" and reference_video is not None:
ref_video_length = reference_video.shape[1] # (C, T, H, W)
if colormatch == "reinhard_torch":
reference_video = match_and_blend_colors(reference_video, original_color_reference, 1.0)
if ref_video_length >= total_frames:
# Reference is long enough - extract keyframes at the expected positions
segment_interval = frame_num - motion_frame - drop_frames
generate_idx = []
current_idx = frame_num - 1
while current_idx < total_frames:
generate_idx.append(min(current_idx, ref_video_length - 1))
current_idx += segment_interval
else:
# Calculate target indices then map to reference video
audio_length = total_frames
generate_idx_target = [0]
segment_interval = frame_num - motion_frame - drop_frames
current_idx = frame_num - 1
while current_idx < audio_length - 1:
generate_idx_target.append(current_idx)
current_idx += segment_interval
if generate_idx_target[-1] != audio_length - 1:
generate_idx_target.append(audio_length - 1)
# Map target indices to reference video
generate_idx_target = np.array(generate_idx_target, dtype=np.int16)
original_max = generate_idx_target[-1]
original_min = generate_idx_target[0]
if original_max > original_min:
generate_idx_float = (generate_idx_target.astype(np.float64) - original_min) * (ref_video_length - 1) / (original_max - original_min)
generate_idx = np.clip(np.round(generate_idx_float), 0, ref_video_length - 1).astype(np.int32).tolist()
else:
generate_idx = [0]
generate_idx = generate_idx[1:]
log.info(f"Reference video ({ref_video_length} frames) mapped to target ({total_frames} frames). Keyframe indices: {generate_idx}")
# Extract keyframes from reference video
# reference_video shape: (C, T, H, W) from nodes.py processing
# Select keyframes and add batch dimension: (C, num_keyframes, H, W) -> (1, C, num_keyframes, H, W)
selected_keyframes = reference_video[:, generate_idx] # (C, num_keyframes, H, W)
reference_keyframes = selected_keyframes.unsqueeze(0).cpu() # (1, C, num_keyframes, H, W)
log.info(f"Extracted {len(generate_idx)} keyframes from provided reference video at indices {generate_idx}, shape: {reference_keyframes.shape}")
log.info(f"Reference video total frames: {reference_video.shape[1]}, will generate {total_frames} total frames with {estimated_iterations} windows")
if frame_num >= total_frames: if frame_num >= total_frames:
arrive_last_frame = True arrive_last_frame = True
estimated_iterations = 1 estimated_iterations = 1
@@ -122,6 +176,14 @@ def multitalk_loop(self, **kwargs):
while True: # start video generation iteratively while True: # start video generation iteratively
self.cache_state = [None, None] self.cache_state = [None, None]
if mode == "skyreelsv3" and reference_keyframes is not None:
clamped_index = min(keyframe_index, reference_keyframes.shape[2] - 1) # Clamp keyframe_index to reuse last keyframe if we run out
pseudo_frames = reference_keyframes[:, :, clamped_index:clamped_index+1].repeat(1, 1, num_pseudo_frames, 1, 1) # Use one keyframe and repeat it 5 times
log.info(f"Window {iteration_count}: using keyframe {clamped_index}/{reference_keyframes.shape[2]-1} for pseudo frames.")
keyframe_index += 1
else:
pseudo_frames = None
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4) cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk": if mode == "infinitetalk":
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
@@ -133,15 +195,13 @@ def multitalk_loop(self, **kwargs):
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1) 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_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
audio_embs.append(audio_emb) audio_embs.append(audio_emb)
audio_embs = torch.concat(audio_embs, dim=0).to(dtype) audio_embs = torch.cat(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) 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] lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
latent_frame_num = (frame_num - 1) // 4 + 1 latent_frame_num = (frame_num - 1) // 4 + 1
noise = torch.randn( noise = torch.randn(16, latent_frame_num, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
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 # Calculate the correct latent slice based on current iteration
if is_first_clip: if is_first_clip:
@@ -198,22 +258,39 @@ def multitalk_loop(self, **kwargs):
if cond_image is not None or cond_frame is not None: 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_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
cond_frame_num = cond_.shape[2] cond_frame_num = cond_.shape[2]
# Prepare pseudo frames if enabled and available from reference_video
if mode == "skyreelsv3" and pseudo_frames is not None:
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num-num_pseudo_frames, target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames, pseudo_frames.to(device, vae.dtype)], dim=2)
else:
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype) 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) padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames], dim=2)
# encode # encode
vae.to(device) vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0] y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
if mode == "multitalk": if mode == "infinitetalk":
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
else:
cond_ = cond_image if is_first_clip else cond_frame 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] latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
else:
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
vae.to(offload_device) vae.to(offload_device)
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1 #motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
if mode == "skyreelsv3" and pseudo_frames is not None:
# create mask in pixel space, then transform
msk_pixel = torch.ones(1, frame_num, lat_h, lat_w, device=device)
msk_pixel[:, cur_motion_frames_num : -num_pseudo_frames] = 0
msk_pixel = torch.cat([
torch.repeat_interleave(msk_pixel[:, 0:1], repeats=4, dim=1),
msk_pixel[:, 1:],
], dim=1)
msk_pixel = msk_pixel.view(1, msk_pixel.shape[1] // 4, 4, lat_h, lat_w)
msk = msk_pixel.transpose(1, 2).squeeze(0).to(dtype) # 4 T H W
else:
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype) msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
msk[:, :1] = 1 msk[:, :1] = 1
y = torch.cat([msk, y]) # 4+C T H W y = torch.cat([msk, y]) # 4+C T H W
@@ -258,11 +335,12 @@ def multitalk_loop(self, **kwargs):
latent = noise latent = noise
# injecting motion frames # injecting motion frames
if not is_first_clip and mode == "multitalk": if not is_first_clip and mode != "infinitetalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device) 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() 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]) add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
latent[:, :add_latent.shape[1]] = add_latent latent[:, :add_latent.shape[1]] = add_latent
del motion_add_noise, add_latent
if offloaded: if offloaded:
# Load weights # Load weights
@@ -370,12 +448,13 @@ def multitalk_loop(self, **kwargs):
latent = image_latent * mask + latent * (1-mask) latent = image_latent * mask + latent * (1-mask)
# injecting motion frames # injecting motion frames
if not is_first_clip and mode == "multitalk": if not is_first_clip and mode != "infinitetalk":
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device) 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() 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]) add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
latent[:, :add_latent.shape[1]] = add_latent latent[:, :add_latent.shape[1]] = add_latent
else: del motion_add_noise, add_latent
elif mode == "infinitetalk":
if humo_image_cond is None or not is_first_clip: if humo_image_cond is None or not is_first_clip:
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
@@ -385,24 +464,31 @@ def multitalk_loop(self, **kwargs):
offloaded = True offloaded = True
if humo_image_cond is not None and humo_reference_count > 0: if humo_image_cond is not None and humo_reference_count > 0:
latent = latent[:,:-humo_reference_count] latent = latent[:,:-humo_reference_count]
vae.to(device) vae.to(device)
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
vae.to(offload_device) vae.to(offload_device)
sampling_pbar.close() sampling_pbar.close()
# crop drop_frames from end if enabled
if mode == "skyreelsv3" and drop_frames > 0 and not arrive_last_frame:
videos = videos[:, :-drop_frames]
# optional color correction (less relevant for InfiniteTalk) # optional color correction (less relevant for InfiniteTalk)
if colormatch != "disabled": if colormatch != "disabled":
if colormatch == "reinhard_torch":
videos = match_and_blend_colors(videos, original_color_reference, 1.0)
else:
videos = videos.permute(1, 2, 3, 0).float().numpy() videos = videos.permute(1, 2, 3, 0).float().numpy()
from color_matcher import ColorMatcher from color_matcher import ColorMatcher
cm = ColorMatcher() cm = ColorMatcher()
cm_result_list = [] cm_result_list = []
for img in videos: for img in videos:
if mode == "multitalk": if mode == "infinitetalk":
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 = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
else:
cm_result = cm.transfer(src=img, ref=original_images[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)) 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) videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
@@ -441,7 +527,7 @@ def multitalk_loop(self, **kwargs):
# Repeat audio emb # Repeat audio emb
if multitalk_embeds is not None: if multitalk_embeds is not None:
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count) audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count - drop_frames)
audio_end_idx = audio_start_idx + clip_length audio_end_idx = audio_start_idx + clip_length
if audio_end_idx >= len(audio_embedding[0]): if audio_end_idx >= len(audio_embedding[0]):
arrive_last_frame = True arrive_last_frame = True
@@ -478,6 +564,6 @@ def multitalk_loop(self, **kwargs):
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path}, return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
+89 -1
View File
@@ -128,7 +128,7 @@ class MultiTalkModelLoader:
def loudness_norm(audio_array, sr=16000, lufs=-23): def loudness_norm(audio_array, sr=16000, lufs=-23):
try: try:
import pyloudnorm import pyloudnorm
except: except Exception:
raise ImportError("pyloudnorm package is not installed") raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr) meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array) loudness = meter.integrated_loudness(audio_array)
@@ -462,12 +462,99 @@ class WanVideoImageToVideoMultiTalk:
return (image_embeds, output_path) return (image_embeds, output_path)
class WanVideoImageToVideoSkyreelsv3_audio:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"vae": ("WANVAE",),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the generation"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the generation"}),
"frame_window_size": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "The number of frames to process at once, should be a value the model is generally good at."}),
"motion_frame": ("INT", {"default": 5, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation. Basically the overlap length."}),
"drop_frames": ("INT", {"default": 12, "min": 0, "max": 10000, "step": 1, "tooltip": "Additional frames to drop when advancing the audio window. Higher values = less overlap = faster generation but potentially less smooth transitions."}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
"force_offload": ("BOOLEAN", {"default": False, "tooltip": "Whether to force offload the model within the loop for VAE operations, enable if you encounter memory issues."}),
"colormatch": (
[
'disabled',
'reinhard_torch',
'mkl',
'hm',
'reinhard',
'mvgd',
'hm-mvgd-hm',
'hm-mkl-hm',
], {
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
},),
},
"optional": {
"start_image": ("IMAGE", {"tooltip": "Images to encode"}),
"reference_video": ("IMAGE", {"tooltip": "Optional: Pre-generated reference video to use for keyframes instead of extracting from first generation. Should be color-matched to source image."}),
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
"output_path": ("STRING", {"default": "", "tooltip": "If set, will save each window's resulting frames to this folder, also DISABLES returning the final video tensor to save memory"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "STRING",)
RETURN_NAMES = ("image_embeds", "output_path")
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Enables Multi/InfiniteTalk long video generation sampling method, the video is created in windows with overlapping frames. Not compatible or necessary to be used with context windows and many other features besides Multi/InfiniteTalk."
def process(self, vae, width, height, frame_window_size, motion_frame, drop_frames, force_offload, colormatch, start_image=None,
tiled_vae=False, clip_embeds=None, mode="multitalk", output_path="", reference_video=None):
H, W = height, width
num_frames = ((frame_window_size - 1) // 4) * 4 + 1
# Resize and rearrange the input image dimensions
if start_image is not None:
resized_start_image = common_upscale(start_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
resized_start_image = resized_start_image * 2 - 1
resized_start_image = resized_start_image.unsqueeze(0)
target_shape = (16, (num_frames - 1) // 4 + 1, height // 8, width // 8)
if output_path:
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
output_path = os.path.join(output_path, f"{timestamp}_{mode}_output")
os.makedirs(output_path, exist_ok=True)
processed_reference_video = None
if reference_video is not None:
processed_reference_video = common_upscale(reference_video.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
processed_reference_video = processed_reference_video * 2 - 1
image_embeds = {
"multitalk_sampling": True,
"multitalk_start_image": resized_start_image if start_image is not None else None,
"frame_window_size": num_frames,
"motion_frame": motion_frame,
"drop_frames": drop_frames,
"use_pseudo_frames": True,
"reference_video": processed_reference_video,
"target_h": H,
"target_w": W,
"tiled_vae": tiled_vae,
"force_offload": force_offload,
"vae": vae,
"target_shape": target_shape,
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
"colormatch": colormatch,
"multitalk_mode": "skyreelsv3",
"output_path": output_path
}
return (image_embeds, output_path)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"MultiTalkModelLoader": MultiTalkModelLoader, "MultiTalkModelLoader": MultiTalkModelLoader,
"MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds, "MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds,
"WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk, "WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk,
"Wav2VecModelLoader": Wav2VecModelLoader, "Wav2VecModelLoader": Wav2VecModelLoader,
"MultiTalkSilentEmbeds": MultiTalkSilentEmbeds, "MultiTalkSilentEmbeds": MultiTalkSilentEmbeds,
"WanVideoImageToVideoSkyreelsv3_audio": WanVideoImageToVideoSkyreelsv3_audio,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
@@ -476,4 +563,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk", "WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk",
"Wav2VecModelLoader": "Wav2vec2 Model Loader", "Wav2VecModelLoader": "Wav2vec2 Model Loader",
"MultiTalkSilentEmbeds": "MultiTalk Silent Embeds", "MultiTalkSilentEmbeds": "MultiTalk Silent Embeds",
"WanVideoImageToVideoSkyreelsv3_audio": "WanVideo Long SkyReelsV3 A2V",
} }
+32 -8
View File
@@ -2,6 +2,7 @@ import os, gc, math
import torch import torch
import torch.nn.functional as F import torch.nn.functional as F
import hashlib import hashlib
from tqdm import tqdm
from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device) from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, set_module_tensor_to_device)
from .taehv import TAEHV from .taehv import TAEHV
@@ -326,7 +327,7 @@ class WanVideoTextEncode:
try: try:
log.info(f"Moving video model to {offload_device}") log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device) model_to_offload.model.to(offload_device)
except: except Exception:
pass pass
encoder = t5["model"] encoder = t5["model"]
@@ -366,10 +367,18 @@ class WanVideoTextEncode:
cast_dtype = encoder.dtype cast_dtype = encoder.dtype
params_to_keep = {'norm', 'pos_embedding', 'token_embedding'} params_to_keep = {'norm', 'pos_embedding', 'token_embedding'}
for name, param in encoder.model.named_parameters(): if hasattr(encoder, 'state_dict'):
model_state_dict = encoder.state_dict
else:
model_state_dict = encoder.model.state_dict()
params_list = list(encoder.model.named_parameters())
pbar = tqdm(params_list, desc="Loading T5 parameters", leave=True)
for name, param in pbar:
dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype
value = encoder.state_dict[name] if hasattr(encoder, 'state_dict') else encoder.model.state_dict()[name] value = model_state_dict[name]
set_module_tensor_to_device(encoder.model, name, device=device_to, dtype=dtype_to_use, value=value) set_module_tensor_to_device(encoder.model, name, device=device_to, dtype=dtype_to_use, value=value)
del model_state_dict
if hasattr(encoder, 'state_dict'): if hasattr(encoder, 'state_dict'):
del encoder.state_dict del encoder.state_dict
mm.soft_empty_cache() mm.soft_empty_cache()
@@ -493,7 +502,7 @@ class WanVideoTextEncodeSingle:
log.info(f"Moving video model to {offload_device}") log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device) model_to_offload.model.to(offload_device)
mm.soft_empty_cache() mm.soft_empty_cache()
except: except Exception:
pass pass
encoder = t5["model"] encoder = t5["model"]
@@ -550,6 +559,9 @@ class WanVideoApplyNAG:
"nag_tau": ("FLOAT", {"default": 2.5, "min": 0.0, "max": 10.0, "step": 0.1}), "nag_tau": ("FLOAT", {"default": 2.5, "min": 0.0, "max": 10.0, "step": 0.1}),
"nag_alpha": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}), "nag_alpha": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}),
}, },
"optional": {
"inplace": ("BOOLEAN", {"default": True, "tooltip": "If true, modifies tensors in place to save memory. Leads to different numerical results which may change the output slightly."}),
}
} }
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", ) RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
@@ -558,7 +570,7 @@ class WanVideoApplyNAG:
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Adds NAG prompt embeds to original prompt embeds: 'https://github.com/ChenDarYen/Normalized-Attention-Guidance'" DESCRIPTION = "Adds NAG prompt embeds to original prompt embeds: 'https://github.com/ChenDarYen/Normalized-Attention-Guidance'"
def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha): def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha, inplace=True):
prompt_embeds_dict_copy = original_text_embeds.copy() prompt_embeds_dict_copy = original_text_embeds.copy()
prompt_embeds_dict_copy.update({ prompt_embeds_dict_copy.update({
"nag_prompt_embeds": nag_text_embeds["prompt_embeds"], "nag_prompt_embeds": nag_text_embeds["prompt_embeds"],
@@ -566,6 +578,7 @@ class WanVideoApplyNAG:
"nag_scale": nag_scale, "nag_scale": nag_scale,
"nag_tau": nag_tau, "nag_tau": nag_tau,
"nag_alpha": nag_alpha, "nag_alpha": nag_alpha,
"inplace": inplace,
} }
}) })
return (prompt_embeds_dict_copy,) return (prompt_embeds_dict_copy,)
@@ -895,6 +908,8 @@ class WanVideoAddStoryMemLatents:
"vae": ("WANVAE",), "vae": ("WANVAE",),
"embeds": ("WANVIDIMAGE_EMBEDS",), "embeds": ("WANVIDIMAGE_EMBEDS",),
"memory_images": ("IMAGE",), "memory_images": ("IMAGE",),
"rope_negative_offset": ("BOOLEAN", {"default": False, "tooltip": "Use positive RoPE frequency offset for the memory latents"}),
"rope_negative_offset_frames": ("INT", {"default": 5, "min": 0, "max": 100, "step": 1, "tooltip": "RoPE frequency offset for the memory latents"}),
} }
} }
@@ -903,10 +918,11 @@ class WanVideoAddStoryMemLatents:
FUNCTION = "add" FUNCTION = "add"
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def add(self, vae, embeds, memory_images): def add(self, vae, embeds, memory_images, rope_negative_offset, rope_negative_offset_frames):
updated = dict(embeds) updated = dict(embeds)
story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images) 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] updated["story_mem_latents"] = story_mem_latents["samples"].squeeze(2).permute(1, 0, 2, 3) # [C, T, H, W]
updated["rope_negative_offset_frames"] = rope_negative_offset_frames if rope_negative_offset else 0
return (updated,) return (updated,)
@@ -1191,6 +1207,7 @@ class WanVideoAnimateEmbeds:
"face_images": ("IMAGE", {"tooltip": "end frame"}), "face_images": ("IMAGE", {"tooltip": "end frame"}),
"bg_images": ("IMAGE", {"tooltip": "background images"}), "bg_images": ("IMAGE", {"tooltip": "background images"}),
"mask": ("MASK", {"tooltip": "mask"}), "mask": ("MASK", {"tooltip": "mask"}),
"start_ref_image": ("IMAGE", {"tooltip": "start ref image"}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
} }
} }
@@ -1201,7 +1218,7 @@ class WanVideoAnimateEmbeds:
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength, def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength,
ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None): ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None, start_ref_image=None):
W = (width // 16) * 16 W = (width // 16) * 16
H = (height // 16) * 16 H = (height // 16) * 16
@@ -1212,7 +1229,7 @@ class WanVideoAnimateEmbeds:
num_refs = ref_images.shape[0] if ref_images is not None else 0 num_refs = ref_images.shape[0] if ref_images is not None else 0
num_frames = ((num_frames - 1) // 4) * 4 + 1 num_frames = ((num_frames - 1) // 4) * 4 + 1
looping = num_frames > frame_window_size looping = num_frames > frame_window_size or start_ref_image is not None
if num_frames < frame_window_size: if num_frames < frame_window_size:
frame_window_size = num_frames frame_window_size = num_frames
@@ -1310,6 +1327,12 @@ class WanVideoAnimateEmbeds:
resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0) resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0)
resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype) resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype)
if start_ref_image is not None:
if start_ref_image.shape[1] != H or start_ref_image.shape[2] != W:
resized_start_ref_image = common_upscale(start_ref_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
else:
resized_start_ref_image = start_ref_image.permute(3, 0, 1, 2) # C, T, H, W
resized_start_ref_image = resized_start_ref_image[:3] * 2 - 1
seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1]) seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1])
@@ -1329,6 +1352,7 @@ class WanVideoAnimateEmbeds:
"is_masked": mask is not None, "is_masked": mask is not None,
"ref_latent": ref_latent, "ref_latent": ref_latent,
"ref_image": resized_ref_images if ref_images is not None else None, "ref_image": resized_ref_images if ref_images is not None else None,
"start_ref_image": resized_start_ref_image if start_ref_image is not None else None,
"face_pixels": resized_face_images if face_images is not None else None, "face_pixels": resized_face_images if face_images is not None else None,
"num_frames": num_frames, "num_frames": num_frames,
"target_shape": target_shape, "target_shape": target_shape,
+100 -74
View File
@@ -23,7 +23,7 @@ from comfy.sd import load_lora_for_models
try: try:
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
from gguf import GGMLQuantizationType from gguf import GGMLQuantizationType
except: except Exception:
pass pass
script_directory = os.path.dirname(os.path.abspath(__file__)) script_directory = os.path.dirname(os.path.abspath(__file__))
@@ -33,7 +33,7 @@ offload_device = mm.unet_offload_device()
try: try:
from server import PromptServer from server import PromptServer
except: except Exception:
PromptServer = None PromptServer = None
attention_modes = ["sdpa", "flash_attn_2", "flash_attn_3", "sageattn", "sageattn_3", "radial_sage_attention", "sageattn_compiled", attention_modes = ["sdpa", "flash_attn_2", "flash_attn_3", "sageattn", "sageattn_3", "radial_sage_attention", "sageattn_compiled",
@@ -52,17 +52,6 @@ def update_folder_names_and_paths(key, targets=[]):
log.warning(f"Unknown file list already present on key {key}: {base}") log.warning(f"Unknown file list already present on key {key}: {base}")
update_folder_names_and_paths("unet_gguf", ["diffusion_models", "unet"]) update_folder_names_and_paths("unet_gguf", ["diffusion_models", "unet"])
class WanVideoModel(comfy.model_base.BaseModel):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.pipeline = {}
def __getitem__(self, k):
return self.pipeline[k]
def __setitem__(self, k, v):
self.pipeline[k] = v
try: try:
from comfy.latent_formats import Wan21, Wan22 from comfy.latent_formats import Wan21, Wan22
latent_format = Wan21 latent_format = Wan21
@@ -71,16 +60,27 @@ except: #for backwards compatibility
from comfy.latent_formats import HunyuanVideo from comfy.latent_formats import HunyuanVideo
latent_format = HunyuanVideo latent_format = HunyuanVideo
class WanVideoModel(torch.nn.Module):
def __init__(self, model_config, transformer, device=None):
super().__init__()
self.latent_format = model_config.latent_format
self.model_config = model_config
self.device = device
self.current_patcher = None
self.diffusion_model = transformer
self.pipeline = {}
def __getitem__(self, k):
return self.pipeline[k]
def __setitem__(self, k, v):
self.pipeline[k] = v
class WanVideoModelConfig: class WanVideoModelConfig:
def __init__(self, dtype, latent_format=latent_format): def __init__(self, latent_format=latent_format):
self.unet_config = {} self.unet_config = {}
self.unet_extra_config = {} self.unet_extra_config = {}
self.latent_format = latent_format self.latent_format = latent_format
#self.latent_format.latent_channels = 16
self.manual_cast_dtype = dtype
self.sampling_settings = {"multiplier": 1.0}
self.memory_usage_factor = 2.0
self.unet_config["disable_unet_model_creation"] = True
def filter_state_dict_by_blocks(state_dict, blocks_mapping, layer_filter=[]): def filter_state_dict_by_blocks(state_dict, blocks_mapping, layer_filter=[]):
filtered_dict = {} filtered_dict = {}
@@ -414,7 +414,7 @@ class WanVideoLoraSelect:
try: try:
lora_path = folder_paths.get_full_path_or_raise("loras", lora) lora_path = folder_paths.get_full_path_or_raise("loras", lora)
except: except Exception:
lora_path = lora lora_path = lora
# Load metadata from the safetensors file # Load metadata from the safetensors file
@@ -810,7 +810,6 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"} "adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"}
param_count = sum(1 for _ in transformer.named_parameters()) param_count = sum(1 for _ in transformer.named_parameters())
pbar = ProgressBar(param_count) pbar = ProgressBar(param_count)
cnt = 0
block_idx = vace_block_idx = None block_idx = vace_block_idx = None
if gguf: if gguf:
@@ -830,7 +829,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
all_tensors.extend(r.tensors) all_tensors.extend(r.tensors)
for tensor in all_tensors: for tensor in all_tensors:
name = rename_fuser_block(tensor.name) name = rename_fuser_block(tensor.name)
if "glob" not in name and "audio_proj" in name: if "glob" not in name and "multitalk_audio_proj" not in name and "audio_proj" in name:
name = name.replace("audio_proj", "multitalk_audio_proj") name = name.replace("audio_proj", "multitalk_audio_proj")
load_device = device load_device = device
if "vace_blocks." in name: if "vace_blocks." in name:
@@ -920,9 +919,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
load_device = offload_device load_device = offload_device
# Set tensor to device # Set tensor to device
set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=value) set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=value)
cnt += 1 pbar.update(1)
if cnt % 100 == 0:
pbar.update(100)
#[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()] #[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()]
memory_on_device = get_module_memory_mb_per_device(transformer) memory_on_device = get_module_memory_mb_per_device(transformer)
@@ -931,6 +928,8 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
for dev, mem_mb in memory_on_device.items(): for dev, mem_mb in memory_on_device.items():
log.info(f"Device: {dev:8s} | Memory: {mem_mb:,.2f} MB") log.info(f"Device: {dev:8s} | Memory: {mem_mb:,.2f} MB")
if hasattr(pbar, "_last_sent_value"):
pbar._last_sent_value = -1
pbar.update_absolute(0) pbar.update_absolute(0)
def patch_control_lora(transformer, device): def patch_control_lora(transformer, device):
@@ -1152,7 +1151,7 @@ class WanVideoModelLoader:
try: try:
if hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation"): if hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation"):
torch.backends.cuda.matmul.allow_fp16_accumulation = False torch.backends.cuda.matmul.allow_fp16_accumulation = False
except: except Exception:
pass pass
@@ -1512,7 +1511,57 @@ class WanVideoModelLoader:
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False) block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False) block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
if multitalk_model is not None: # LongCat Avatar
proj1_key = "multitalk_audio_proj.proj1.weight" if "multitalk_audio_proj.proj1.weight" in sd \
else "multitalk_audio_proj.proj1.weight_int8" if "multitalk_audio_proj.proj1.weight_int8" in sd \
else None
if proj1_key is not None and ("blocks.0.audio_cross_attn.q_norm.weight" in sd or "blocks.0.audio_cross_attn.q_norm.weight_int8" in sd):
log.info("MultiTalk/InfiniteTalk model detected, patching model...")
from .multitalk.multitalk import AudioProjModel
from .wanvideo.modules.model import WanLayerNorm
from .LongCat.layers import SingleStreamAttention
# Detect LongCat-Avatar audio encoder variant from proj1 input dim:
# v1.0 (wav2vec2): seq_len * blocks * channels = 5 * 12 * 768 = 46080
# v1.5 (whisper): seq_len * blocks * channels = 5 * 5 * 1280 = 32000
proj1_in = sd[proj1_key].shape[1]
if proj1_in == 32000:
audio_proj_blocks, audio_proj_channels = 5, 1280
log.info("LongCat-Avatar-1.5 (Whisper) audio proj detected")
else:
audio_proj_blocks, audio_proj_channels = 12, 768
for block in transformer.blocks:
with init_empty_weights():
if "blocks.0.audio_modulation.1.weight" in sd:
block.audio_modulation = nn.Sequential(nn.SiLU(), nn.Linear(512, 3 * dim, bias=True))
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamAttention(
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
qkv_bias=True,
qk_norm=True,
class_range=24,
class_interval=4,
attention_mode=attention_mode,
)
multitalk_proj_model = AudioProjModel(blocks=audio_proj_blocks, channels=audio_proj_channels)
transformer.multitalk_audio_proj = multitalk_proj_model
# SkyreelsV3
elif "blocks.1.audio_cross_attn.kv_linear.weight" in sd and "audio_proj.proj1.weight" in sd:
sd = {k.replace("audio_proj", "multitalk_audio_proj"): v for k, v in sd.items()}
# init audio module
from .multitalk.multitalk import SingleStreamMultiAttention, AudioProjModel
from .wanvideo.modules.model import WanLayerNorm
for block in transformer.blocks:
with init_empty_weights():
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamMultiAttention(dim=dim, num_heads=num_heads, attention_mode=attention_mode)
transformer.multitalk_audio_proj = AudioProjModel()
elif multitalk_model is not None:
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk") multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
log.info(f"{multitalk_model_type} detected, patching model...") log.info(f"{multitalk_model_type} detected, patching model...")
@@ -1529,15 +1578,7 @@ class WanVideoModelLoader:
for block in transformer.blocks: for block in transformer.blocks:
with init_empty_weights(): with init_empty_weights():
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamMultiAttention( block.audio_cross_attn = SingleStreamMultiAttention(dim=dim, num_heads=num_heads, attention_mode=attention_mode)
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
qkv_bias=True,
class_range=24,
class_interval=4,
attention_mode=attention_mode,
)
transformer.multitalk_audio_proj = multitalk_model["proj_model"] transformer.multitalk_audio_proj = multitalk_model["proj_model"]
transformer.multitalk_model_type = multitalk_model_type transformer.multitalk_model_type = multitalk_model_type
@@ -1555,39 +1596,6 @@ class WanVideoModelLoader:
sd.update(extra_sd) sd.update(extra_sd)
del extra_sd del extra_sd
elif "multitalk_audio_proj.proj1.weight" in sd:
log.info("MultiTalk/InfiniteTalk model detected, patching model...")
from .multitalk.multitalk import AudioProjModel
from .wanvideo.modules.model import WanLayerNorm
from .LongCat.layers import SingleStreamAttention
audio_window = 5
vae_scale = 4
for block in transformer.blocks:
with init_empty_weights():
if "blocks.0.audio_modulation.1.weight" in sd:
block.audio_modulation = nn.Sequential(nn.SiLU(), nn.Linear(512, 3 * dim, bias=True))
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
block.audio_cross_attn = SingleStreamAttention(
dim=dim,
encoder_hidden_states_dim=768,
num_heads=num_heads,
qkv_bias=True,
qk_norm=True,
class_range=24,
class_interval=4,
attention_mode=attention_mode,
)
multitalk_proj_model = AudioProjModel(
seq_len=audio_window,
seq_len_vf=audio_window+vae_scale-1,
intermediate_dim=512,
output_dim=768,
context_tokens=32,
norm_output_audio=True,
)
transformer.multitalk_audio_proj = multitalk_proj_model
sd = {k.replace(".weight_scale", ".scale_weight"): v for k, v in sd.items()} sd = {k.replace(".weight_scale", ".scale_weight"): v for k, v in sd.items()}
@@ -1613,11 +1621,7 @@ class WanVideoModelLoader:
transformer.text_projection = nn.Sequential(nn.Linear(sd["text_projection.0.weight"].shape[1], text_dim), nn.GELU(approximate='tanh'), nn.Linear(text_dim, text_dim)) transformer.text_projection = nn.Sequential(nn.Linear(sd["text_projection.0.weight"].shape[1], text_dim), nn.GELU(approximate='tanh'), nn.Linear(text_dim, text_dim))
latent_format=Wan22 if dim == 3072 else Wan21 latent_format=Wan22 if dim == 3072 else Wan21
comfy_model = WanVideoModel( comfy_model = WanVideoModel(WanVideoModelConfig(latent_format=latent_format), device=device, transformer=transformer)
WanVideoModelConfig(base_dtype, latent_format=latent_format),
model_type=comfy.model_base.ModelType.FLOW,
device=device,
)
# SteadyDancer # SteadyDancer
if "condition_embedding_align.cross_attn.in_proj_bias" in sd: if "condition_embedding_align.cross_attn.in_proj_bias" in sd:
@@ -1681,6 +1685,24 @@ class WanVideoModelLoader:
block.ref_attn_v_img = nn.Linear(in_features, out_features) block.ref_attn_v_img = nn.Linear(in_features, out_features)
block.ref_attn_norm_k_img = WanRMSNorm(out_features, eps=1e-6) block.ref_attn_norm_k_img = WanRMSNorm(out_features, eps=1e-6)
if "blocks.0.control_blocks_dense.cross_attn.k.weight" in sd:
log.info("LongVie2 model detected, patching model...")
from .LongVie2.modules import WanModelDualControl
control_layers = 12
with init_empty_weights():
dual_controller = WanModelDualControl(dim=5120, ffn_dim=13824, eps=1e-06, num_heads=40, control_layers=control_layers)
for b in range(control_layers):
transformer.blocks[b].control_blocks_dense = dual_controller.control_blocks_dense[b]
transformer.blocks[b].control_blocks_sparse = dual_controller.control_blocks_sparse[b]
transformer.blocks[b].control_combine_linears = dual_controller.control_combine_linears[b]
transformer.dual_controller = nn.Module()
transformer.dual_controller.control_initial_combine_linear_dense = dual_controller.control_initial_combine_linear_dense
transformer.dual_controller.control_initial_combine_linear_sparse = dual_controller.control_initial_combine_linear_sparse
transformer.dual_controller.control_t_mod = dual_controller.control_t_mod
transformer.dual_controller.control_text_linear = dual_controller.control_text_linear
transformer.dual_controller_freqs = dual_controller.freqs
comfy_model.diffusion_model = transformer comfy_model.diffusion_model = transformer
comfy_model.load_device = transformer_load_device comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
@@ -1785,10 +1807,14 @@ class WanVideoModelLoader:
) )
if merge_loras and lora is not None: if merge_loras and lora is not None:
# Skip offloading if load_device is main_device (for unified memory systems like AMD Strix Halo)
if load_device != "main_device":
log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}") log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}")
patcher.model.diffusion_model.to(offload_device) patcher.model.diffusion_model.to(offload_device)
gc.collect() gc.collect()
mm.soft_empty_cache() mm.soft_empty_cache()
else:
log.info(f"Skipping offload (load_device=main_device, keeping model on {patcher.model.diffusion_model.device})")
patcher.model["base_dtype"] = base_dtype patcher.model["base_dtype"] = base_dtype
patcher.model["weight_dtype"] = weight_dtype patcher.model["weight_dtype"] = weight_dtype
+54 -15
View File
@@ -185,6 +185,7 @@ class WanVideoSampler:
is_pusa = "pusa" in sample_scheduler.__class__.__name__.lower() is_pusa = "pusa" in sample_scheduler.__class__.__name__.lower()
if scheduler != "multitalk":
scheduler_step_args = {"generator": seed_g} scheduler_step_args = {"generator": seed_g}
step_sig = inspect.signature(sample_scheduler.step) step_sig = inspect.signature(sample_scheduler.step)
for arg in list(scheduler_step_args.keys()): for arg in list(scheduler_step_args.keys()):
@@ -225,6 +226,7 @@ class WanVideoSampler:
#I2V #I2V
story_mem_latents = image_embeds.get("story_mem_latents", None) story_mem_latents = image_embeds.get("story_mem_latents", None)
image_cond = image_embeds.get("image_embeds", None) image_cond = image_embeds.get("image_embeds", None)
image_cond_mask = None
if image_cond is not None: if image_cond is not None:
if transformer.in_dim == 16: if transformer.in_dim == 16:
raise ValueError("T2V (text to video) model detected, encoded images only work with I2V (Image to video) models") raise ValueError("T2V (text to video) model detected, encoded images only work with I2V (Image to video) models")
@@ -562,6 +564,7 @@ class WanVideoSampler:
# MultiTalk # MultiTalk
multitalk_audio_embeds = audio_emb_slice = audio_features_in = None multitalk_audio_embeds = audio_emb_slice = audio_features_in = None
multitalk_audio_stride = None
multitalk_embeds = image_embeds.get("multitalk_embeds", multitalk_embeds) multitalk_embeds = image_embeds.get("multitalk_embeds", multitalk_embeds)
if multitalk_embeds is not None: if multitalk_embeds is not None:
@@ -582,6 +585,7 @@ class WanVideoSampler:
audio_scale = multitalk_embeds.get("audio_scale", 1.0) audio_scale = multitalk_embeds.get("audio_scale", 1.0)
audio_cfg_scale = multitalk_embeds.get("audio_cfg_scale", 1.0) audio_cfg_scale = multitalk_embeds.get("audio_cfg_scale", 1.0)
ref_target_masks = multitalk_embeds.get("ref_target_masks", None) ref_target_masks = multitalk_embeds.get("ref_target_masks", None)
multitalk_audio_stride = multitalk_embeds.get("audio_stride", None)
if not isinstance(audio_cfg_scale, list): if not isinstance(audio_cfg_scale, list):
audio_cfg_scale = [audio_cfg_scale] * (steps + 1) audio_cfg_scale = [audio_cfg_scale] * (steps + 1)
@@ -815,6 +819,10 @@ class WanVideoSampler:
latent_video_length += insert_len latent_video_length += insert_len
longcat_num_cond_latents = len(clean_latent_indices) longcat_num_cond_latents = len(clean_latent_indices)
log.info(f"LongCat num_cond_latents: {longcat_num_cond_latents} num_ref_latents: {longcat_num_ref_latents}") log.info(f"LongCat num_cond_latents: {longcat_num_cond_latents} num_ref_latents: {longcat_num_ref_latents}")
# v1.5 (Whisper) embeds set audio_stride=1; v1.0 (wav2vec2) uses 2 for LongCat
if multitalk_audio_stride is not None:
audio_stride = multitalk_audio_stride
else:
audio_stride = 2 if transformer.is_longcat else 1 audio_stride = 2 if transformer.is_longcat else 1
#controlnet #controlnet
@@ -1148,6 +1156,18 @@ class WanVideoSampler:
if context_options is None: if context_options is None:
image_cond = replace_feature(image_cond.unsqueeze(0).clone(), track_pos.unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0] image_cond = replace_feature(image_cond.unsqueeze(0).clone(), track_pos.unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
# LongVie2 dual control
dual_control_embeds = image_embeds.get("dual_control", None)
if dual_control_embeds is not None and context_options is None:
dual_control_input = dict_to_device(dual_control_embeds.copy(), device, dtype) if dual_control_embeds is not None else None
prev_latents = dual_control_input.get("prev_latent", None)
if prev_latents is not None:
_sigma = dual_control_embeds.get("first_frame_noise_level", 0.925926)
log.info(f"Using dual control previous latents with first frame noise level: {_sigma}")
latent[:, :1] = (1 - _sigma) * prev_latents[:, -1:].to(latent) + _sigma * noise[:, :1]
prev_ones = torch.ones(20, *prev_latents.shape[1:], device=device, dtype=dtype)
dual_control_input["prev_latent"] = torch.cat([prev_ones, prev_latents]).unsqueeze(0)
#region model pred #region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
@@ -1400,6 +1420,19 @@ class WanVideoSampler:
if wanmove_embeds is not None and context_window is not None: if wanmove_embeds is not None and context_window is not None:
image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0] image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
dual_control_in = None
if dual_control_embeds is not None:
if context_window is not None:
dual_control_in = dual_control_embeds.copy()
dense_input_latent = dual_control_embeds.get("dense_input_latent", None)
if dense_input_latent is not None:
dual_control_in["dense_input_latent"] = dual_control_embeds["dense_input_latent"][:, :, context_window]
sparse_input_latent = dual_control_embeds.get("sparse_input_latent", None)
if sparse_input_latent is not None:
dual_control_in["sparse_input_latent"] = dual_control_embeds["sparse_input_latent"][:, :, context_window]
else:
dual_control_in = dual_control_input
base_params = { base_params = {
'x': [z], # latent 'x': [z], # latent
'y': [image_cond_input] if image_cond_input is not None else None, # image cond 'y': [image_cond_input] if image_cond_input is not None else None, # image cond
@@ -1462,7 +1495,10 @@ class WanVideoSampler:
"one_to_all_input": one_to_all_data, # One-to-All input "one_to_all_input": one_to_all_data, # One-to-All input
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0, "one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
"scail_input": scail_data_in, # SCAIL input "scail_input": scail_data_in, # SCAIL input
"transformer_options": transformer_options "dual_control_input": dual_control_in, # LongVie2 dual control input
"transformer_options": transformer_options,
"rope_negative_offset": image_embeds.get("rope_negative_offset_frames", 0), # StoryMem rope negative offset
"num_memory_frames": story_mem_latents.shape[1] if story_mem_latents is not None else 0, # StoryMem memory frames
} }
batch_size = 1 batch_size = 1
@@ -1700,7 +1736,7 @@ class WanVideoSampler:
gc.collect() gc.collect()
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
# Main sampling loop with FreeInit iterations # Main sampling loop with FreeInit iterations
@@ -2152,7 +2188,7 @@ class WanVideoSampler:
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return {"video": gen_video_samples}, return {"video": gen_video_samples},
# region wananimate loop # region wananimate loop
@@ -2176,7 +2212,11 @@ class WanVideoSampler:
bg_images = image_embeds.get("bg_images", None) bg_images = image_embeds.get("bg_images", None)
pose_images = image_embeds.get("pose_images", None) pose_images = image_embeds.get("pose_images", None)
current_ref_images = face_images = face_images_in = None current_ref_images = image_embeds.get("start_ref_image", None)
if current_ref_images is not None:
log.info(
"WanAnimate: Detected manual start reference image, enabling continuous generation across windows.")
face_images = face_images_in = None
if wananim_face_pixels is not None: if wananim_face_pixels is not None:
face_images = tensor_pingpong_pad(wananim_face_pixels, target_len) face_images = tensor_pingpong_pad(wananim_face_pixels, target_len)
@@ -2215,6 +2255,9 @@ class WanVideoSampler:
mm.soft_empty_cache() mm.soft_empty_cache()
if current_ref_images is not None:
mask_reft_len = refert_num
else:
mask_reft_len = 0 if start == 0 else refert_num mask_reft_len = 0 if start == 0 else refert_num
self.cache_state = [None, None] self.cache_state = [None, None]
@@ -2394,7 +2437,7 @@ class WanVideoSampler:
videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
del latent del latent
if start != 0: if start != 0 or current_ref_images is not None:
videos = videos[:, refert_num:] videos = videos[:, refert_num:]
sampling_pbar.close() sampling_pbar.close()
@@ -2446,7 +2489,7 @@ class WanVideoSampler:
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path}, return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
@@ -2556,10 +2599,10 @@ class WanVideoSampler:
except Exception as e: except Exception as e:
log.error(f"Error during sampling: {e}") log.error(f"Error during sampling: {e}")
if force_offload: raise
if not model["auto_cpu_offload"]: finally:
if force_offload and not model["auto_cpu_offload"]:
offload_transformer(transformer) offload_transformer(transformer)
raise e
if phantom_latents is not None: if phantom_latents is not None:
latent = latent[:,:-phantom_latents.shape[1]] latent = latent[:,:-phantom_latents.shape[1]]
@@ -2583,14 +2626,10 @@ class WanVideoSampler:
"magcache_state": transformer.magcache_state, "magcache_state": transformer.magcache_state,
} }
if force_offload:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return ({ return ({
"samples": latent.unsqueeze(0).cpu(), "samples": latent.unsqueeze(0).cpu(),
@@ -2734,7 +2773,7 @@ class WanVideoScheduler:
import io import io
import base64 import base64
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
except: except Exception:
PromptServer = None PromptServer = None
if unique_id and PromptServer is not None: if unique_id and PromptServer is not None:
try: try:
+4 -4
View File
@@ -9,7 +9,7 @@ from einops import rearrange
try: try:
from server import PromptServer from server import PromptServer
except: except Exception:
PromptServer = None PromptServer = None
VAE_STRIDE = (4, 8, 8) VAE_STRIDE = (4, 8, 8)
@@ -256,7 +256,7 @@ class CreateCFGScheduleFloatList:
f"{cfg_list}", f"{cfg_list}",
unique_id unique_id
) )
except: except Exception:
pass pass
return (cfg_list,) return (cfg_list,)
@@ -319,7 +319,7 @@ class CreateScheduleFloatList:
f"{cfg_list}", f"{cfg_list}",
unique_id unique_id
) )
except: except Exception:
pass pass
return (cfg_list,) return (cfg_list,)
@@ -454,7 +454,7 @@ class NormalizeAudioLoudness:
def loudness_norm(self, audio_array, sr=16000, lufs=-23): def loudness_norm(self, audio_array, sr=16000, lufs=-23):
try: try:
import pyloudnorm import pyloudnorm
except: except Exception:
raise ImportError("pyloudnorm package is not installed") raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr) meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array) loudness = meter.integrated_loudness(audio_array)
+1 -1
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "ComfyUI-WanVideoWrapper" name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo" description = "ComfyUI wrapper nodes for WanVideo"
version = "1.4.5" version = "1.4.7"
license = {file = "LICENSE"} license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"] dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.17.0", "ftfy", "gguf >= 0.17.1", "pyloudnorm"]
+2 -2
View File
@@ -548,7 +548,7 @@ class WanVideoDiffusionForcingSampler:
gc.collect() gc.collect()
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
#region main loop start #region main loop start
@@ -615,7 +615,7 @@ class WanVideoDiffusionForcingSampler:
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return ({ return ({
+3 -3
View File
@@ -200,7 +200,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
if ref_image is not None: if ref_image is not None:
try: try:
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold) pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
except: except Exception:
raise ValueError("No pose detected in reference image") raise ValueError("No pose detected in reference image")
prev_pose = None prev_pose = None
for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)): for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)):
@@ -208,7 +208,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
pose = dwpose_model(img, score_threshold=score_threshold) pose = dwpose_model(img, score_threshold=score_threshold)
if handle_not_detected == "repeat": if handle_not_detected == "repeat":
prev_pose = pose prev_pose = pose
except: except Exception:
if prev_pose is not None: if prev_pose is not None:
pose = prev_pose pose = prev_pose
else: else:
@@ -675,7 +675,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size,
draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head) draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
result = torch.from_numpy(dwpose_woface) result = torch.from_numpy(dwpose_woface)
#except: #except Exception:
# result = torch.zeros((height, width, 3), dtype=torch.uint8) # result = torch.zeros((height, width, 3), dtype=torch.uint8)
dwpose_woface_list.append(result) dwpose_woface_list.append(result)
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0) dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
+65 -7
View File
@@ -7,9 +7,14 @@ from pathlib import Path
import gc import gc
import types, collections import types, collections
from comfy.utils import ProgressBar, copy_to_param, set_attr_param from comfy.utils import ProgressBar, copy_to_param, set_attr_param
from comfy.model_patcher import get_key_weight, string_to_seed from comfy.model_patcher import get_key_weight
from comfy.lora import calculate_weight from comfy.lora import calculate_weight
try:
from comfy.utils import string_to_seed
except Exception:
from comfy.model_patcher import string_to_seed
from comfy.float import stochastic_rounding from comfy.float import stochastic_rounding
from .custom_linear import remove_lora_from_module from .custom_linear import remove_lora_from_module
import folder_paths import folder_paths
@@ -22,7 +27,7 @@ offload_device = mm.unet_offload_device()
try: try:
from .gguf.gguf import GGUFParameter from .gguf.gguf import GGUFParameter
except: except Exception:
pass pass
COLOR_CODES = { COLOR_CODES = {
@@ -190,9 +195,9 @@ def set_module_tensor_to_device(module, tensor_name, device, value=None, dtype=N
device = device_quantization device = device_quantization
if is_buffer: if is_buffer:
module._buffers[tensor_name] = new_value module._buffers[tensor_name] = new_value
elif value is not None or not check_device_same(torch.device(device), module._parameters[tensor_name].device): elif value is not None or not check_device_same(device, module._parameters[tensor_name].device):
param_cls = type(module._parameters[tensor_name]) param_cls = type(module._parameters[tensor_name])
new_value = param_cls(new_value, requires_grad=False).to(device) new_value = param_cls(new_value, requires_grad=False)
module._parameters[tensor_name] = new_value module._parameters[tensor_name] = new_value
#if device != "cpu": #if device != "cpu":
@@ -304,7 +309,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
key = f"{name.replace('diffusion_model.', '')}.{param}" key = f"{name.replace('diffusion_model.', '')}.{param}"
try: try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key]) set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key])
except: except Exception:
continue continue
key = f"{name}.{param}" key = f"{name}.{param}"
if scale_weights is not None: if scale_weights is not None:
@@ -318,7 +323,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
if low_mem_load: if low_mem_load:
try: try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key]) set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key])
except: except Exception:
continue continue
m.comfy_patched_weights = True m.comfy_patched_weights = True
cnt += 1 cnt += 1
@@ -347,7 +352,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
dtype_to_use = torch.float32 dtype_to_use = torch.float32
try: try:
set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name]) set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name])
except: except Exception:
continue continue
return model return model
@@ -700,6 +705,7 @@ def check_duplicate_nodes():
for path in custom_nodes_dir.iterdir(): for path in custom_nodes_dir.iterdir():
if (path.is_dir() and if (path.is_dir() and
path != current_path and path != current_path and
not path.name.endswith('.disabled') and
'wanvideo' in path.name.lower() and 'wanvideo' in path.name.lower() and
'wrapper' in path.name.lower()): 'wrapper' in path.name.lower()):
wanvideo_dirs.append(str(path)) wanvideo_dirs.append(str(path))
@@ -718,3 +724,55 @@ def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.
if not t == 1.0: if not t == 1.0:
model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t) model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t)
return model_output return model_output
def match_and_blend_colors(
source_chunk: torch.Tensor, # (C, T, H, W), range [-1, 1]
reference_image: torch.Tensor, # (C, 1, H, W), range [-1, 1]
strength: float,
) -> torch.Tensor:
import kornia
if strength == 0.0:
return source_chunk
source_chunk = source_chunk.unsqueeze(0) # (1, C, T, H, W)
# shapes
B, C, T, H, W = source_chunk.shape
input_dtype = source_chunk.dtype
# [-1,1] -> [0,1]
src_01 = (source_chunk + 1.0) * 0.5
ref_01 = (reference_image + 1.0) * 0.5
src32 = src_01.to(torch.float32)
ref32 = ref_01.to(torch.float32)
# (B, C, T, H, W) -> (B*T, C, H, W)
src_bt = src32.permute(0, 2, 1, 3, 4).contiguous().view(B * T, C, H, W)
ref_bchw = ref32[:, :, 0, :, :].contiguous()
# RGB->Lab
src_lab = kornia.color.rgb_to_lab(src_bt) # (B*T, C, H, W)
ref_lab = kornia.color.rgb_to_lab(ref_bchw) # (B, C, H, W)
src_lab_flat = src_lab.view(B * T, C, -1) # (B*T, C, HW)
ref_lab_flat = ref_lab.view(B, C, -1) # (B, C, HW)
src_std, src_mean = torch.std_mean(src_lab_flat, dim=-1, keepdim=True, unbiased=False)
ref_std, ref_mean = torch.std_mean(ref_lab_flat, dim=-1, keepdim=True, unbiased=False)
src_std = src_std.clamp_min_(1e-6)
ref_mean_bt = ref_mean.repeat_interleave(T, dim=0) # (B*T, C, 1)
ref_std_bt = ref_std.repeat_interleave(T, dim=0) # (B*T, C, 1)
corrected_lab_flat = (src_lab_flat - src_mean) * (ref_std_bt / src_std) + ref_mean_bt
corrected_lab = corrected_lab_flat.view(B * T, C, H, W)
# Lab->RGB
corrected_rgb_01 = kornia.color.lab_to_rgb(corrected_lab) # (B*T, C, H, W)
blended_rgb_01 = (1.0 - strength) * src_bt + strength * corrected_rgb_01
# (B, C, T, H, W)
blended_rgb_01 = blended_rgb_01.view(B, T, C, H, W).permute(0, 2, 1, 3, 4).contiguous()
# [0,1] -> [-1,1]
return (blended_rgb_01 * 2.0 - 1.0)[0].to(dtype=input_dtype)
+4 -4
View File
@@ -65,16 +65,16 @@ try:
# Return tensor with same shape as q # Return tensor with same shape as q
return q.clone() return q.clone()
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
except: except Exception:
sageattn_varlen_func = attention_func_error sageattn_varlen_func = attention_func_error
# sage3 # sage3
try: try:
from sageattn3 import sageattn3_blackwell as sageattn_blackwell from sageattn3 import sageattn3_blackwell as sageattn_blackwell
except: except Exception:
try: try:
from sageattn import sageattn_blackwell from sageattn import sageattn_blackwell
except: except Exception:
sageattn_blackwell = attention_func_error sageattn_blackwell = attention_func_error
try: try:
@@ -88,7 +88,7 @@ try:
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9): def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
return torch.empty_like(qkv[0]).contiguous() return torch.empty_like(qkv[0]).contiguous()
sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico
except: except Exception:
sageattn_func_ultravico = attention_func_error sageattn_func_ultravico = attention_func_error
+133 -36
View File
@@ -10,7 +10,7 @@ from contextlib import nullcontext
try: try:
from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense, MaskMap from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense, MaskMap
except: except Exception:
pass pass
from .attention import attention from .attention import attention
@@ -558,39 +558,53 @@ class WanSelfAttention(nn.Module):
# output # output
return self.o(x.flatten(2)) return self.o(x.flatten(2))
def normalized_attention_guidance(self, b, n, d, q, context, nag_context=None, nag_params={}): def nag_attention(self, b, n, d, q, context, nag_context=None):
k_positive = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
v_positive = self.v(context).view(b, -1, n, d)
x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode, heads=self.num_heads)
del k_positive, v_positive
k_negative = self.norm_k(self.k(nag_context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
v_negative = self.v(nag_context).view(b, -1, n, d)
x_negative = attention(q, k_negative, v_negative, attention_mode=self.attention_mode, heads=self.num_heads)
del k_negative, v_negative
return x_positive.flatten(2), x_negative.flatten(2)
def normalized_attention_guidance(self, x_positive, x_negative,nag_params={}):
# NAG text attention # NAG text attention
context_positive = context
context_negative = nag_context
nag_scale = nag_params['nag_scale'] nag_scale = nag_params['nag_scale']
nag_alpha = nag_params['nag_alpha'] nag_alpha = nag_params['nag_alpha']
nag_tau = nag_params['nag_tau'] nag_tau = nag_params['nag_tau']
inplace = nag_params.get('inplace', True)
k_positive = self.norm_k(self.k(context_positive).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype) if inplace:
v_positive = self.v(context_positive).view(b, -1, n, d) nag_guidance = x_negative.mul_(nag_scale - 1).neg_().add_(x_positive, alpha=nag_scale)
k_negative = self.norm_k(self.k(context_negative).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype) else:
v_negative = self.v(context_negative).view(b, -1, n, d)
x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode, heads=self.num_heads)
x_positive = x_positive.flatten(2)
x_negative = attention(q, k_negative, v_negative, attention_mode=self.attention_mode, heads=self.num_heads)
x_negative = x_negative.flatten(2)
nag_guidance = x_positive * nag_scale - x_negative * (nag_scale - 1) nag_guidance = x_positive * nag_scale - x_negative * (nag_scale - 1)
del x_negative
norm_positive = torch.norm(x_positive, p=1, dim=-1, keepdim=True) norm_positive = torch.norm(x_positive, p=1, dim=-1, keepdim=True)
norm_guidance = torch.norm(nag_guidance, p=1, dim=-1, keepdim=True) norm_guidance = torch.norm(nag_guidance, p=1, dim=-1, keepdim=True)
scale = norm_guidance / norm_positive scale = norm_guidance / norm_positive
scale = torch.nan_to_num(scale, nan=10.0) torch.nan_to_num_(scale, nan=10.0)
mask = scale > nag_tau mask = scale > nag_tau
del scale
adjustment = (norm_positive * nag_tau) / (norm_guidance + 1e-7) adjustment = (norm_positive * nag_tau) / (norm_guidance + 1e-7)
nag_guidance = torch.where(mask, nag_guidance * adjustment, nag_guidance) del norm_positive, norm_guidance
nag_guidance.mul_(torch.where(mask, adjustment, 1.0))
del mask, adjustment del mask, adjustment
return nag_guidance * nag_alpha + x_positive * (1 - nag_alpha) if inplace:
nag_guidance.sub_(x_positive).mul_(nag_alpha).add_(x_positive)
else:
nag_guidance = nag_guidance * nag_alpha + x_positive * (1 - nag_alpha)
del x_positive
return nag_guidance
class LoRALinearLayer(nn.Module): class LoRALinearLayer(nn.Module):
def __init__( def __init__(
@@ -633,7 +647,7 @@ class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0, def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy", num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy",
inner_t=None, inner_c=None, cross_freqs=None, inner_t=None, inner_c=None, cross_freqs=None,
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, longcat_num_cond_latents=None, **kwargs): adapter_proj=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, longcat_num_cond_latents=None, **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim b, n, d = x.size(0), self.num_heads, self.head_dim
s = x.size(1) s = x.size(1)
# compute query # compute query
@@ -648,7 +662,10 @@ class WanT2VCrossAttention(WanSelfAttention):
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d) q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d)
if nag_context is not None: if nag_context is not None:
x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) x_positive, x_negative = self.nag_attention(b, n, d, q, context, nag_context)
del q
x = self.normalized_attention_guidance(x_positive, x_negative, nag_params)
del x_positive, x_negative
else: else:
if is_longcat: if is_longcat:
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype).view(b, -1, n, d)).to(x.dtype) k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype).view(b, -1, n, d)).to(x.dtype)
@@ -728,7 +745,7 @@ class WanI2VCrossAttention(WanSelfAttention):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy", audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy",
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs): adapter_proj=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
r""" r"""
Args: Args:
x(Tensor): Shape [B, L1, C] x(Tensor): Shape [B, L1, C]
@@ -740,21 +757,22 @@ class WanI2VCrossAttention(WanSelfAttention):
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype) q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype)
if nag_context is not None: if nag_context is not None:
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) x_positive, x_negative = self.nag_attention(b, n, d, q, context, nag_context)
x = self.normalized_attention_guidance(x_positive, x_negative, nag_params)
del x_positive, x_negative
else: else:
# text attention # text attention
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype) k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype)
v = self.v(context).view(b, -1, n, d) v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) x = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
del k, v
#img attention #img attention
if clip_embed is not None: if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype) k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype)
v_img = self.v_img(clip_embed).view(b, -1, n, d) v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) x.add_(attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2))
x = x_text + img_x del k_img, v_img
else:
x = x_text
# FantasyTalking audio attention # FantasyTalking audio attention
if audio_proj is not None: if audio_proj is not None:
@@ -787,7 +805,7 @@ class WanI2VCrossAttention(WanSelfAttention):
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads) adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads)
adapter_x = adapter_x.flatten(2) adapter_x = adapter_x.flatten(2)
x = x + adapter_x * ip_scale x = x + adapter_x * ip_scale
del q
return self.o(x) return self.o(x)
class WanHuMoCrossAttention(WanSelfAttention): class WanHuMoCrossAttention(WanSelfAttention):
@@ -1280,7 +1298,7 @@ class WanAttentionBlock(nn.Module):
y[:, tr_end:] * gate_msa y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype) ], dim=1).to(input_dtype)
else: else:
x = x.addcmul(y, gate_msa) x.addcmul_(y, gate_msa)
del y, gate_msa del y, gate_msa
# cross-attention & ffn function # cross-attention & ffn function
@@ -1310,11 +1328,10 @@ class WanAttentionBlock(nn.Module):
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes) x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
return x, x_ip, lynx_ref_feature, x_ovi return x, x_ip, lynx_ref_feature, x_ovi
else: else:
x = x + self.cross_attn(self.norm3(x.to(self.norm3.weight.dtype)).to(input_dtype), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale, x += self.cross_attn(self.norm3(x.to(self.norm3.weight.dtype)).to(input_dtype), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context,
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs, rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, longcat_num_cond_latents=longcat_num_cond_latents) adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, longcat_num_cond_latents=longcat_num_cond_latents).to(input_dtype)
x = x.to(input_dtype)
# MultiTalk # MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock): if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
@@ -1328,7 +1345,8 @@ class WanAttentionBlock(nn.Module):
else: else:
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding, x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num) shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
x = x.add(x_audio, alpha=audio_scale) x.add_(x_audio, alpha=audio_scale)
del x_audio
# MTV-Crafter Motion Attention # MTV-Crafter Motion Attention
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None: if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
@@ -2188,7 +2206,7 @@ class WanModel(torch.nn.Module):
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, ref_frame_shape=None, pose_frame_shape=None, def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, ref_frame_shape=None, pose_frame_shape=None,
steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None,
ref_frame_index=10, longcat_num_ref_latents=0): ref_frame_index=10, longcat_num_ref_latents=0, num_memory_frames=3, rope_negative_offset=0):
patch_size = self.patch_size patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0]) t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
@@ -2212,6 +2230,15 @@ class WanModel(torch.nn.Module):
torch.arange(0, steps_t - longcat_num_ref_latents, dtype=dtype, device=device) torch.arange(0, steps_t - longcat_num_ref_latents, dtype=dtype, device=device)
], dim=0) ], dim=0)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1) img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1)
elif num_memory_frames > 0 and rope_negative_offset > 0:
# Negative RoPE shift for memory frames
# Memory frames get negative indices: {-f_m*S, -(f_m-1)*S, ..., -S}
# Current video frames start from 0: {0, 1, ..., f-1}
memory_indices = torch.arange(-num_memory_frames * rope_negative_offset, 0, rope_negative_offset, dtype=dtype, device=device)
current_indices = torch.arange(0, steps_t - num_memory_frames, dtype=dtype, device=device)
grid_t = torch.cat([memory_indices, current_indices], dim=0)
log.info(f"{num_memory_frames} memory frames, temporal rope positions: {grid_t}")
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1)
else: else:
# Standard temporal encoding # Standard temporal encoding
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start+freq_offset + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1) img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start+freq_offset + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
@@ -2319,7 +2346,10 @@ class WanModel(torch.nn.Module):
sdancer_input=None, # SteadyDancer sdancer_input=None, # SteadyDancer
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
scail_input=None, # SCAIL pose scail_input=None, # SCAIL pose
dual_control_input=None, # LongVie2 dual controlnet
transformer_options={}, transformer_options={},
rope_negative_offset=0,
num_memory_frames=0,
): ):
r""" r"""
Forward pass through the diffusion model Forward pass through the diffusion model
@@ -2546,6 +2576,16 @@ class WanModel(torch.nn.Module):
x = [u.flatten(2).transpose(1, 2) for u in x] x = [u.flatten(2).transpose(1, 2) for u in x]
self.original_seq_len = x[0].shape[1] self.original_seq_len = x[0].shape[1]
prev_latent = None
if dual_control_input is not None:
prev_latent = dual_control_input.get("prev_latent", None)
if prev_latent is not None:
F += prev_latent.shape[2]
prev_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in prev_latent]
prev_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in prev_x]
seq_len += prev_x[0].shape[1]
x = [torch.cat([u, v], dim=1) for u, v in zip(prev_x, x)]
# SCAIL pose # SCAIL pose
if scail_input is not None: if scail_input is not None:
scail_pose_latents = scail_input.get("pose_latent", None) scail_pose_latents = scail_input.get("pose_latent", None)
@@ -2554,6 +2594,7 @@ class WanModel(torch.nn.Module):
scail_x = [u.flatten(2).transpose(1, 2) * scail_input.get("pose_strength", 1) for u in scail_x] scail_x = [u.flatten(2).transpose(1, 2) * scail_input.get("pose_strength", 1) for u in scail_x]
x = [torch.cat([u, v], dim=1) for u, v in zip(x, scail_x)] x = [torch.cat([u, v], dim=1) for u, v in zip(x, scail_x)]
seq_len += scail_x[0].shape[1] seq_len += scail_x[0].shape[1]
del scail_x
pose_frame_shape = scail_pose_latents.shape pose_frame_shape = scail_pose_latents.shape
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32) seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
@@ -2632,6 +2673,8 @@ class WanModel(torch.nn.Module):
self.rope_embedder.k, self.rope_embedder.k,
tuple(ntk_alphas), tuple(ntk_alphas),
longcat_num_ref_latents, longcat_num_ref_latents,
rope_negative_offset,
num_memory_frames,
) )
# Check cache using key comparison # Check cache using key comparison
@@ -2647,6 +2690,8 @@ class WanModel(torch.nn.Module):
ref_frame_shape=ref_frame_shape, ref_frame_shape=ref_frame_shape,
pose_frame_shape=pose_frame_shape, pose_frame_shape=pose_frame_shape,
longcat_num_ref_latents=longcat_num_ref_latents, longcat_num_ref_latents=longcat_num_ref_latents,
rope_negative_offset=rope_negative_offset,
num_memory_frames=num_memory_frames,
device=x.device, device=x.device,
dtype=x.dtype dtype=x.dtype
) )
@@ -2836,6 +2881,44 @@ class WanModel(torch.nn.Module):
chunked_self_attention = False chunked_self_attention = False
seq_chunks = 0 seq_chunks = 0
# dual control
if dual_control_input is not None and dual_control_input["start_percent"] <= current_step_percentage <= dual_control_input["end_percent"]:
dense_latent = dual_control_input["dense_input_latent"]
print("dense_latent shape:", dense_latent.shape)
sparse_latent = dual_control_input["sparse_input_latent"]
if dense_latent is None and sparse_latent is None:
raise ValueError("At least one of dense_input_latent or sparse_input_latent must be provided in dual_control_input")
if dense_latent is not None:
dense_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in dense_latent]
dense_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in dense_x]
dense = self.dual_controller.control_initial_combine_linear_dense(dense_x[0])
if sparse_latent is not None:
sparse_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in sparse_latent]
sparse_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in sparse_x]
sparse = self.dual_controller.control_initial_combine_linear_sparse(sparse_x[0])
if dense_latent is None:
dense = torch.zeros_like(sparse)
elif sparse_latent is None:
sparse = torch.zeros_like(dense)
control_context = clip_fea_control = None
if context != []:
control_context = self.dual_controller.control_text_linear(context)
if clip_embed is not None:
clip_fea_control = self.dual_controller.control_text_linear(clip_embed)
control_t_mod = self.dual_controller.control_t_mod(e0)
control_freqs = torch.cat([
self.dual_controller_freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
self.dual_controller_freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
self.dual_controller_freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
], dim=-1).reshape(f * h * w, 1, -1).to(x.device)
else:
dual_control_input = None
# MultiTalk # MultiTalk
if multitalk_audio is not None: if multitalk_audio is not None:
self.multitalk_audio_proj.to(self.main_device) self.multitalk_audio_proj.to(self.main_device)
@@ -3191,6 +3274,18 @@ class WanModel(torch.nn.Module):
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_onetoall_ref=x_onetoall_ref, onetoall_freqs=onetoall_freqs, attention_mode_override=attention_mode, **kwargs) x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_onetoall_ref=x_onetoall_ref, onetoall_freqs=onetoall_freqs, attention_mode_override=attention_mode, **kwargs)
# ---post block----# # ---post block----#
# dual controlnet
if dual_control_input is not None and (hasattr(block, "control_blocks_dense") or hasattr(block, "control_blocks_sparse")):
if dense_latent is not None and hasattr(block, "control_blocks_dense"):
dense = block.control_blocks_dense(dense, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control)
if sparse_latent is not None and hasattr(block, "control_blocks_sparse"):
sparse = block.control_blocks_sparse(sparse, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control)
if prev_latent is not None:
x[:, -self.original_seq_len:] += block.control_combine_linears(dense + sparse) * dual_control_input["strength"]
else:
x += block.control_combine_linears(dense + sparse) * dual_control_input["strength"]
if self.audio_injector is not None and s2v_audio_input is not None: if self.audio_injector is not None and s2v_audio_input is not None:
x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v
if block.has_face_fuser_block and motion_vec is not None: if block.has_face_fuser_block and motion_vec is not None:
@@ -3293,7 +3388,9 @@ class WanModel(torch.nn.Module):
# x = x[:, :self.original_seq_len] # x = x[:, :self.original_seq_len]
#grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) #grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
if prev_latent is not None:
x = x[:, -self.original_seq_len:]
else:
x = x[:, :self.original_seq_len] x = x[:, :self.original_seq_len]
x = self.head(x, e.to(x.device), temp_length=F, x = self.head(x, e.to(x.device), temp_length=F,
+3 -3
View File
@@ -4,15 +4,15 @@ import torch
try: try:
from spas_sage_attn import block_sparse_sage2_attn_cuda from spas_sage_attn import block_sparse_sage2_attn_cuda
sparse_attn_func = block_sparse_sage2_attn_cuda sparse_attn_func = block_sparse_sage2_attn_cuda
except: except Exception:
try: try:
from sparse_sageattn import sparse_sageattn from sparse_sageattn import sparse_sageattn
sparse_attn_func = sparse_sageattn sparse_attn_func = sparse_sageattn
except: except Exception:
try: try:
from .sparse_sage.core import sparse_sageattn from .sparse_sage.core import sparse_sageattn
sparse_attn_func = sparse_sageattn sparse_attn_func = sparse_sageattn
except: except Exception:
sparse_sageattn = None sparse_sageattn = None
raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.") raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.")
+8 -8
View File
@@ -1061,7 +1061,7 @@ class VideoVAE_(nn.Module):
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar): for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
@@ -1092,7 +1092,7 @@ class VideoVAE_(nn.Module):
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}") log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE encode") print_memory(device, process="WanVAE encode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return mu return mu
@@ -1137,7 +1137,7 @@ class VideoVAE_(nn.Module):
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
x = self.conv2(z) x = self.conv2(z)
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar): for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
@@ -1162,7 +1162,7 @@ class VideoVAE_(nn.Module):
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}") log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode") print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return out return out
@@ -1464,7 +1464,7 @@ class VideoVAE38_(VideoVAE_):
self.clear_cache() self.clear_cache()
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
x = patchify(x, patch_size=2) x = patchify(x, patch_size=2)
t = x.shape[2] t = x.shape[2]
@@ -1492,7 +1492,7 @@ class VideoVAE38_(VideoVAE_):
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}") log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode") print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return mu return mu
@@ -1502,7 +1502,7 @@ class VideoVAE38_(VideoVAE_):
input_shape = z.shape input_shape = z.shape
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
z = z / self.inv_std.to(z) + self.mean.to(z) z = z / self.inv_std.to(z) + self.mean.to(z)
@@ -1531,7 +1531,7 @@ class VideoVAE38_(VideoVAE_):
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}") log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode") print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return out return out