Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
088128b224 | ||
|
|
126819826a | ||
|
|
8f1804bf72 | ||
|
|
5437b016e3 | ||
|
|
d18cdb1859 | ||
|
|
0d78230336 | ||
|
|
df8f3e49da | ||
|
|
86ad93d616 | ||
|
|
309491b269 | ||
|
|
06122f1e9d | ||
|
|
5d36631795 | ||
|
|
3d7b49e2df | ||
|
|
e091c4a774 | ||
|
|
60a387579e | ||
|
|
9ae1a4f2c9 | ||
|
|
f55b7b3d89 | ||
|
|
e4e7f413f7 | ||
|
|
e21fe20a4d | ||
|
|
2c5a04cc63 | ||
|
|
d00abe52d7 | ||
|
|
2952f5d9dd | ||
|
|
339e0fec81 | ||
|
|
2c2a6e1889 | ||
|
|
8640bfad52 | ||
|
|
3b4a711a40 | ||
|
|
b99eac73da | ||
|
|
9f9a8e71c2 | ||
|
|
0fbcbed06a | ||
|
|
f2e2e2550a | ||
|
|
6f9832ed47 | ||
|
|
707bbcd72b | ||
|
|
855d103ee6 | ||
|
|
64191921d4 | ||
|
|
576f065073 | ||
|
|
58c1bcb7ce | ||
|
|
64cbd28e00 |
@@ -1 +0,0 @@
|
|||||||
github: [kijai]
|
|
||||||
+190
-2
@@ -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)",
|
||||||
}
|
}
|
||||||
@@ -0,0 +1,199 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
from ..wanvideo.modules.attention import attention
|
||||||
|
|
||||||
|
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
|
||||||
|
return (x * (1 + scale) + shift)
|
||||||
|
|
||||||
|
|
||||||
|
def sinusoidal_embedding_1d(dim, position):
|
||||||
|
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
|
||||||
|
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
|
||||||
|
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||||
|
return x.to(position.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0):
|
||||||
|
# 3d rope precompute
|
||||||
|
f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta)
|
||||||
|
h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
|
||||||
|
w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
|
||||||
|
return f_freqs_cis, h_freqs_cis, w_freqs_cis
|
||||||
|
|
||||||
|
|
||||||
|
def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0):
|
||||||
|
# 1d rope precompute
|
||||||
|
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
|
||||||
|
[: (dim // 2)].double() / dim))
|
||||||
|
freqs = torch.outer(torch.arange(end, device=freqs.device), freqs)
|
||||||
|
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
|
||||||
|
return freqs_cis
|
||||||
|
|
||||||
|
|
||||||
|
def rope_apply(x, freqs, num_heads):
|
||||||
|
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
|
||||||
|
x_out = torch.view_as_complex(x.to(torch.float64).reshape(
|
||||||
|
x.shape[0], x.shape[1], x.shape[2], -1, 2))
|
||||||
|
x_out = torch.view_as_real(x_out * freqs).flatten(2)
|
||||||
|
return x_out.to(x.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
class RMSNorm(nn.Module):
|
||||||
|
def __init__(self, dim, eps=1e-5):
|
||||||
|
super().__init__()
|
||||||
|
self.eps = eps
|
||||||
|
self.weight = nn.Parameter(torch.ones(dim))
|
||||||
|
|
||||||
|
def norm(self, x):
|
||||||
|
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
dtype = x.dtype
|
||||||
|
return self.norm(x.float()).to(dtype) * self.weight
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionModule(nn.Module):
|
||||||
|
def __init__(self, num_heads, head_dim):
|
||||||
|
super().__init__()
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = head_dim
|
||||||
|
|
||||||
|
def forward(self, q, k, v):
|
||||||
|
b, n, d = q.size(0), self.num_heads, self.head_dim
|
||||||
|
x = attention(
|
||||||
|
q.view(b, -1, n, d),
|
||||||
|
k.view(b, -1, n, d),
|
||||||
|
v.view(b, -1, n, d)
|
||||||
|
)
|
||||||
|
return x.flatten(2)
|
||||||
|
|
||||||
|
|
||||||
|
class SelfAttention(nn.Module):
|
||||||
|
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = dim // num_heads
|
||||||
|
|
||||||
|
self.q = nn.Linear(dim, dim)
|
||||||
|
self.k = nn.Linear(dim, dim)
|
||||||
|
self.v = nn.Linear(dim, dim)
|
||||||
|
self.o = nn.Linear(dim, dim)
|
||||||
|
self.norm_q = RMSNorm(dim, eps=eps)
|
||||||
|
self.norm_k = RMSNorm(dim, eps=eps)
|
||||||
|
|
||||||
|
self.attn = AttentionModule(self.num_heads, self.head_dim)
|
||||||
|
|
||||||
|
def forward(self, x, freqs):
|
||||||
|
q = self.norm_q(self.q(x))
|
||||||
|
k = self.norm_k(self.k(x))
|
||||||
|
v = self.v(x)
|
||||||
|
q = rope_apply(q, freqs, self.num_heads)
|
||||||
|
k = rope_apply(k, freqs, self.num_heads)
|
||||||
|
x = self.attn(q, k, v)
|
||||||
|
return self.o(x)
|
||||||
|
|
||||||
|
|
||||||
|
class CrossAttention(nn.Module):
|
||||||
|
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = dim // num_heads
|
||||||
|
|
||||||
|
self.q = nn.Linear(dim, dim)
|
||||||
|
self.k = nn.Linear(dim, dim)
|
||||||
|
self.v = nn.Linear(dim, dim)
|
||||||
|
self.o = nn.Linear(dim, dim)
|
||||||
|
self.norm_q = RMSNorm(dim, eps=eps)
|
||||||
|
self.norm_k = RMSNorm(dim, eps=eps)
|
||||||
|
|
||||||
|
|
||||||
|
self.k_img = nn.Linear(dim, dim)
|
||||||
|
self.v_img = nn.Linear(dim, dim)
|
||||||
|
self.norm_k_img = RMSNorm(dim, eps=eps)
|
||||||
|
|
||||||
|
self.attn = AttentionModule(self.num_heads, self.head_dim)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None):
|
||||||
|
ctx = y
|
||||||
|
q = self.norm_q(self.q(x))
|
||||||
|
k = self.norm_k(self.k(ctx))
|
||||||
|
v = self.v(ctx)
|
||||||
|
x = self.attn(q, k, v)
|
||||||
|
if clip_fea is not None:
|
||||||
|
k_img = self.norm_k_img(self.k_img(clip_fea))
|
||||||
|
v_img = self.v_img(clip_fea)
|
||||||
|
y = self.attn(q, k_img, v_img)
|
||||||
|
x = x + y
|
||||||
|
return self.o(x)
|
||||||
|
|
||||||
|
|
||||||
|
class GateModule(nn.Module):
|
||||||
|
def __init__(self,):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
def forward(self, x, gate, residual):
|
||||||
|
return x + gate * residual
|
||||||
|
|
||||||
|
class DiTBlock(nn.Module):
|
||||||
|
def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.ffn_dim = ffn_dim
|
||||||
|
|
||||||
|
self.self_attn = SelfAttention(dim, num_heads, eps)
|
||||||
|
self.cross_attn = CrossAttention(dim, num_heads, eps)
|
||||||
|
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
||||||
|
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
||||||
|
self.norm3 = nn.LayerNorm(dim, eps=eps)
|
||||||
|
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(
|
||||||
|
approximate='tanh'), nn.Linear(ffn_dim, dim))
|
||||||
|
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||||
|
self.gate = GateModule()
|
||||||
|
|
||||||
|
def forward(self, x, context, t_mod, freqs, clip_fea=None):
|
||||||
|
has_seq = len(t_mod.shape) == 4
|
||||||
|
chunk_dim = 2 if has_seq else 1
|
||||||
|
# msa: multi-head self-attention mlp: multi-layer perceptron
|
||||||
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||||
|
self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)
|
||||||
|
if has_seq:
|
||||||
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||||
|
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
|
||||||
|
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
|
||||||
|
)
|
||||||
|
input_x = modulate(self.norm1(x), shift_msa, scale_msa)
|
||||||
|
x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
|
||||||
|
x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea)
|
||||||
|
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
|
||||||
|
x = self.gate(x, gate_mlp, self.ffn(input_x))
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class WanModelDualControl(torch.nn.Module):
|
||||||
|
def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12):
|
||||||
|
super().__init__()
|
||||||
|
self.control_layers = control_layers
|
||||||
|
self.control_blocks_dense = nn.ModuleList([
|
||||||
|
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
|
||||||
|
for _ in range(self.control_layers)
|
||||||
|
])
|
||||||
|
|
||||||
|
self.control_blocks_sparse = nn.ModuleList([
|
||||||
|
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
|
||||||
|
for _ in range(self.control_layers)
|
||||||
|
])
|
||||||
|
|
||||||
|
self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2)
|
||||||
|
self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2)
|
||||||
|
|
||||||
|
self.control_text_linear = torch.nn.Linear(dim, dim//2)
|
||||||
|
self.control_t_mod = torch.nn.Linear(dim, dim//2)
|
||||||
|
|
||||||
|
self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)])
|
||||||
|
head_dim = dim // num_heads
|
||||||
|
self.freqs = precompute_freqs_cis_3d(head_dim)
|
||||||
@@ -0,0 +1,88 @@
|
|||||||
|
import torch
|
||||||
|
from ..utils import log
|
||||||
|
import comfy.model_management as mm
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
|
||||||
|
class WanVideoAddDualControlEmbeds:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||||
|
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||||
|
"first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}),
|
||||||
|
"sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}),
|
||||||
|
"prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||||
|
RETURN_NAMES = ("image_embeds",)
|
||||||
|
FUNCTION = "add"
|
||||||
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
|
def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None):
|
||||||
|
updated = dict(embeds)
|
||||||
|
updated.setdefault("dual_control", {})
|
||||||
|
|
||||||
|
if dense is None and sparse is None:
|
||||||
|
raise ValueError("At least one of dense or sparse inputs must be provided.")
|
||||||
|
|
||||||
|
num_frames = dense.shape[0] if dense is not None else sparse.shape[0]
|
||||||
|
height = dense.shape[1] if dense is not None else sparse.shape[1]
|
||||||
|
width = dense.shape[2] if dense is not None else sparse.shape[2]
|
||||||
|
msk = torch.ones(1, num_frames, height//8, width//8, device=device)
|
||||||
|
msk[:, 1:] = 0
|
||||||
|
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
|
||||||
|
msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
|
||||||
|
msk = msk.transpose(1, 2)
|
||||||
|
|
||||||
|
dense_input_latent = sparse_input_latent = None
|
||||||
|
|
||||||
|
vae.to(device)
|
||||||
|
if dense is not None:
|
||||||
|
dense_images = 1 - dense[..., :3] # Invert colors for depth to match the usual range in comfy
|
||||||
|
dense_images = dense_images.permute(3, 0, 1, 2) * 2 - 1
|
||||||
|
dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False)
|
||||||
|
dense_first = (dense_images[:, :1]).to(device, vae.dtype)
|
||||||
|
vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
|
||||||
|
dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False)
|
||||||
|
dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1)
|
||||||
|
dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1)
|
||||||
|
if sparse is not None:
|
||||||
|
sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1
|
||||||
|
sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False)
|
||||||
|
sparse_first = (sparse_images[:, :1]).to(device, vae.dtype)
|
||||||
|
vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
|
||||||
|
sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False)
|
||||||
|
sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1)
|
||||||
|
sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1)
|
||||||
|
|
||||||
|
if prev_images is not None:
|
||||||
|
prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1
|
||||||
|
prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False)
|
||||||
|
updated["dual_control"]["prev_latent"] = prev_video_latent[0]
|
||||||
|
|
||||||
|
vae.to(offload_device)
|
||||||
|
updated["dual_control"]["dense_input_latent"] = dense_input_latent
|
||||||
|
updated["dual_control"]["sparse_input_latent"] = sparse_input_latent
|
||||||
|
updated["dual_control"]["strength"] = strength
|
||||||
|
updated["dual_control"]["start_percent"] = start_percent
|
||||||
|
updated["dual_control"]["end_percent"] = end_percent
|
||||||
|
updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level
|
||||||
|
return (updated,)
|
||||||
|
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds,
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds",
|
||||||
|
}
|
||||||
+1
-1
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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",
|
||||||
}
|
}
|
||||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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,
|
||||||
|
|||||||
@@ -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.")
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user