Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fd32b14fdc | ||
|
|
1776695e26 | ||
|
|
ef36204fa8 | ||
|
|
c6f32c1424 | ||
|
|
6d4a0f6e53 | ||
|
|
e7e00061e5 | ||
|
|
eb5ec262a0 | ||
|
|
7c0ba84a26 | ||
|
|
06a86923e7 | ||
|
|
dca3106f10 | ||
|
|
175418b8d2 | ||
|
|
4a6e2d3c6c |
@@ -0,0 +1 @@
|
||||
github: [kijai]
|
||||
+3890
-3709
File diff suppressed because one or more lines are too long
+5
-194
@@ -1,5 +1,4 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
from comfy_api.latest import io
|
||||
@@ -25,8 +24,6 @@ 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.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.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=[
|
||||
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
|
||||
@@ -35,21 +32,20 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
|
||||
)
|
||||
|
||||
@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, prev_images=None, vae=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) -> io.NodeOutput:
|
||||
|
||||
new_audio_embed = audio_embeds.copy()
|
||||
|
||||
audio_features = torch.stack(new_audio_embed["audio_features"])
|
||||
num_audio_features = audio_features.shape[1]
|
||||
if audio_features.shape[1] < frames_processed + num_frames:
|
||||
deficit = frames_processed + num_frames - audio_features.shape[1]
|
||||
if if_not_enough_audio == "pad_with_start":
|
||||
pad = audio_features[:, :1].repeat(1, deficit, 1, 1)
|
||||
pad = audio_features[:, :1].repeat(1, deficit, 1, 1, 1)
|
||||
audio_features = torch.cat([audio_features, pad], dim=1)
|
||||
elif if_not_enough_audio == "mirror_from_end":
|
||||
to_add = audio_features[:, -deficit:, :].flip(dims=[1])
|
||||
audio_features = torch.cat([audio_features, to_add], dim=1)
|
||||
log.warning(f"Not enough audio features, padded with strategy '{if_not_enough_audio}' from {num_audio_features} to {audio_features.shape[1]} frames")
|
||||
log.info(f"Not enough audio features, extended from {new_audio_embed['audio_features'].shape[1]} to {audio_features.shape[1]} frames.")
|
||||
|
||||
ref_target_masks = new_audio_embed.get("ref_target_masks", None)
|
||||
if ref_target_masks is not None:
|
||||
@@ -58,20 +54,7 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
|
||||
prev_samples = prev_latents["samples"].clone()
|
||||
if overlap != 0:
|
||||
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
|
||||
if ref_latent is not None:
|
||||
@@ -81,7 +64,7 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
|
||||
new_latent_frames = (num_frames - 1) // 4 + 1
|
||||
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
|
||||
|
||||
audio_stride = new_audio_embed.get("audio_stride", 2)
|
||||
audio_stride = 2
|
||||
indices = torch.arange(2 * 2 + 1) - 2
|
||||
|
||||
if frames_processed == 0:
|
||||
@@ -128,181 +111,9 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
|
||||
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 = {
|
||||
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
|
||||
"LongCatAvatarWhisperEmbeds": LongCatAvatarWhisperEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
|
||||
"LongCatAvatarWhisperEmbeds": "LongCat Avatar Whisper Embeds (v1.5)",
|
||||
}
|
||||
@@ -1,199 +0,0 @@
|
||||
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)
|
||||
@@ -1,88 +0,0 @@
|
||||
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" Defined in: {code_file}:{code_line}\n"
|
||||
f"This may cause issues with the NLF model.")
|
||||
except Exception:
|
||||
except:
|
||||
log.warning("--------------------------------")
|
||||
log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, "
|
||||
f"this has been modified by another custom node. This may cause issues with the NLF model.")
|
||||
|
||||
+5
-6
@@ -1,12 +1,12 @@
|
||||
try:
|
||||
from .utils import check_duplicate_nodes, log, color_text
|
||||
from .utils import check_duplicate_nodes, log
|
||||
duplicate_dirs = check_duplicate_nodes()
|
||||
if duplicate_dirs:
|
||||
warning_msg = f"WARNING: Found {len(duplicate_dirs)} other WanVideoWrapper directories:\n"
|
||||
for dir_path in duplicate_dirs:
|
||||
warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
|
||||
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
|
||||
except Exception:
|
||||
warning_msg += f" - {dir_path}\n"
|
||||
log.warning(warning_msg + "Please remove duplicates to avoid possible conflicts.")
|
||||
except:
|
||||
pass
|
||||
|
||||
from .utils import log
|
||||
@@ -49,7 +49,6 @@ OPTIONAL_MODULES = [
|
||||
(".WanMove.nodes", "WanMove"),
|
||||
(".SCAIL.nodes", "SCAIL"),
|
||||
(".LongCat.nodes", "LongCat"),
|
||||
(".LongVie2.nodes", "LongVie2"),
|
||||
]
|
||||
|
||||
def register_nodes(module_path: str, name: str, optional: bool) -> None:
|
||||
@@ -72,4 +71,4 @@ for module_path, name in REQUIRED_MODULES:
|
||||
for module_path, name in OPTIONAL_MODULES:
|
||||
register_nodes(module_path, name, optional=True)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
+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.", "")
|
||||
_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 "dual_controller" not in module_prefix and name not in modules_to_not_convert:
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix and name not in modules_to_not_convert:
|
||||
weight_key = module_prefix + "weight"
|
||||
if weight_key not in state_dict:
|
||||
continue
|
||||
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,780 @@
|
||||
{
|
||||
"id": "8b7a9a57-2303-4ef5-9fc2-bf41713bd1fc",
|
||||
"revision": 0,
|
||||
"last_node_id": 46,
|
||||
"last_link_id": 58,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 33,
|
||||
"type": "Note",
|
||||
"pos": [
|
||||
227.3764190673828,
|
||||
-205.28524780273438
|
||||
],
|
||||
"size": [
|
||||
351.70458984375,
|
||||
88
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"Models:\nhttps://huggingface.co/Kijai/WanVideo_comfy/tree/main"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 11,
|
||||
"type": "LoadWanVideoT5TextEncoder",
|
||||
"pos": [
|
||||
224.15325927734375,
|
||||
-34.481563568115234
|
||||
],
|
||||
"size": [
|
||||
377.1661376953125,
|
||||
130
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "wan_t5_model",
|
||||
"type": "WANTEXTENCODER",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
15
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadWanVideoT5TextEncoder"
|
||||
},
|
||||
"widgets_values": [
|
||||
"umt5-xxl-enc-bf16.safetensors",
|
||||
"bf16",
|
||||
"offload_device",
|
||||
"disabled"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 28,
|
||||
"type": "WanVideoDecode",
|
||||
"pos": [
|
||||
1692.973876953125,
|
||||
-404.8614501953125
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
174
|
||||
],
|
||||
"flags": {},
|
||||
"order": 12,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "WANVAE",
|
||||
"link": 43
|
||||
},
|
||||
{
|
||||
"name": "samples",
|
||||
"type": "LATENT",
|
||||
"link": 33
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
48
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoDecode"
|
||||
},
|
||||
"widgets_values": [
|
||||
true,
|
||||
272,
|
||||
272,
|
||||
144,
|
||||
128
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 38,
|
||||
"type": "WanVideoVAELoader",
|
||||
"pos": [
|
||||
1687.4093017578125,
|
||||
-582.2750854492188
|
||||
],
|
||||
"size": [
|
||||
416.25482177734375,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "WANVAE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
43
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoVAELoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"wanvideo\\Wan2_1_VAE_bf16.safetensors",
|
||||
"bf16"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 42,
|
||||
"type": "GetImageSizeAndCount",
|
||||
"pos": [
|
||||
1708.7301025390625,
|
||||
-140.99705505371094
|
||||
],
|
||||
"size": [
|
||||
277.20001220703125,
|
||||
86
|
||||
],
|
||||
"flags": {},
|
||||
"order": 13,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 48
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
56
|
||||
]
|
||||
},
|
||||
{
|
||||
"label": "832 width",
|
||||
"name": "width",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"label": "480 height",
|
||||
"name": "height",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"label": "257 count",
|
||||
"name": "count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "GetImageSizeAndCount"
|
||||
},
|
||||
"widgets_values": []
|
||||
},
|
||||
{
|
||||
"id": 16,
|
||||
"type": "WanVideoTextEncode",
|
||||
"pos": [
|
||||
675.8850708007812,
|
||||
-36.032100677490234
|
||||
],
|
||||
"size": [
|
||||
420.30511474609375,
|
||||
261.5306701660156
|
||||
],
|
||||
"flags": {},
|
||||
"order": 10,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "t5",
|
||||
"type": "WANTEXTENCODER",
|
||||
"link": 15
|
||||
},
|
||||
{
|
||||
"name": "model_to_offload",
|
||||
"shape": 7,
|
||||
"type": "WANVIDEOMODEL",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "text_embeds",
|
||||
"type": "WANVIDEOTEXTEMBEDS",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
30
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoTextEncode"
|
||||
},
|
||||
"widgets_values": [
|
||||
"high quality nature video featuring a red panda balancing on a bamboo stem while a bird lands on it's head, on the background there is a waterfall",
|
||||
"色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走",
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 30,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
2127.120849609375,
|
||||
-511.9014587402344
|
||||
],
|
||||
"size": [
|
||||
873.2135620117188,
|
||||
840.2385864257812
|
||||
],
|
||||
"flags": {},
|
||||
"order": 14,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 56
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 16,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "WanVideo2_1_T2V",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"trim_to_audio": false,
|
||||
"pingpong": false,
|
||||
"save_output": true,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "WanVideo2_1_T2V_00412.mp4",
|
||||
"subfolder": "",
|
||||
"type": "output",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 16,
|
||||
"workflow": "WanVideo2_1_T2V_00412.png",
|
||||
"fullpath": "N:\\AI\\ComfyUI\\output\\WanVideo2_1_T2V_00412.mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 37,
|
||||
"type": "WanVideoEmptyEmbeds",
|
||||
"pos": [
|
||||
1305.26708984375,
|
||||
-571.7843627929688
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "image_embeds",
|
||||
"type": "WANVIDIMAGE_EMBEDS",
|
||||
"links": [
|
||||
42
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoEmptyEmbeds"
|
||||
},
|
||||
"widgets_values": [
|
||||
832,
|
||||
480,
|
||||
257
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 35,
|
||||
"type": "WanVideoTorchCompileSettings",
|
||||
"pos": [
|
||||
193.47103881835938,
|
||||
-614.6900024414062
|
||||
],
|
||||
"size": [
|
||||
390.5999755859375,
|
||||
178
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "torch_compile_args",
|
||||
"type": "WANCOMPILEARGS",
|
||||
"slot_index": 0,
|
||||
"links": []
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoTorchCompileSettings"
|
||||
},
|
||||
"widgets_values": [
|
||||
"inductor",
|
||||
false,
|
||||
"default",
|
||||
false,
|
||||
64,
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 45,
|
||||
"type": "WanVideoTeaCache",
|
||||
"pos": [
|
||||
931.4036865234375,
|
||||
-792.5159912109375
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
154
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cache_args",
|
||||
"type": "CACHEARGS",
|
||||
"links": [
|
||||
58
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoTeaCache"
|
||||
},
|
||||
"widgets_values": [
|
||||
0.10000000000000002,
|
||||
1,
|
||||
-1,
|
||||
"offload_device",
|
||||
true
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 36,
|
||||
"type": "Note",
|
||||
"pos": [
|
||||
796.0189208984375,
|
||||
-521.5020751953125
|
||||
],
|
||||
"size": [
|
||||
298.2554016113281,
|
||||
108.62744140625
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"sdpa should work too, haven't tested flaash\n\nfp8_fast seems to cause huge quality degradation"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 46,
|
||||
"type": "Note",
|
||||
"pos": [
|
||||
937.9556274414062,
|
||||
-940.750244140625
|
||||
],
|
||||
"size": [
|
||||
297.4364013671875,
|
||||
88
|
||||
],
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"TeaCache with context windows is VERY experimental and lower values than normal should be used."
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
},
|
||||
{
|
||||
"id": 27,
|
||||
"type": "WanVideoSampler",
|
||||
"pos": [
|
||||
1315.2401123046875,
|
||||
-401.48028564453125
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
574.1923217773438
|
||||
],
|
||||
"flags": {},
|
||||
"order": 11,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "WANVIDEOMODEL",
|
||||
"link": 29
|
||||
},
|
||||
{
|
||||
"name": "text_embeds",
|
||||
"type": "WANVIDEOTEXTEMBEDS",
|
||||
"link": 30
|
||||
},
|
||||
{
|
||||
"name": "image_embeds",
|
||||
"type": "WANVIDIMAGE_EMBEDS",
|
||||
"link": 42
|
||||
},
|
||||
{
|
||||
"name": "samples",
|
||||
"shape": 7,
|
||||
"type": "LATENT",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "feta_args",
|
||||
"shape": 7,
|
||||
"type": "FETAARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "context_options",
|
||||
"shape": 7,
|
||||
"type": "WANVIDCONTEXT",
|
||||
"link": 57
|
||||
},
|
||||
{
|
||||
"name": "cache_args",
|
||||
"shape": 7,
|
||||
"type": "CACHEARGS",
|
||||
"link": 58
|
||||
},
|
||||
{
|
||||
"name": "flowedit_args",
|
||||
"shape": 7,
|
||||
"type": "FLOWEDITARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "slg_args",
|
||||
"shape": 7,
|
||||
"type": "SLGARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "loop_args",
|
||||
"shape": 7,
|
||||
"type": "LOOPARGS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "samples",
|
||||
"type": "LATENT",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
33
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoSampler"
|
||||
},
|
||||
"widgets_values": [
|
||||
30,
|
||||
6,
|
||||
5,
|
||||
1057359483639288,
|
||||
"fixed",
|
||||
true,
|
||||
"unipc",
|
||||
0,
|
||||
1,
|
||||
"",
|
||||
"comfy"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 43,
|
||||
"type": "WanVideoContextOptions",
|
||||
"pos": [
|
||||
1307.9542236328125,
|
||||
-855.8865356445312
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
226
|
||||
],
|
||||
"flags": {},
|
||||
"order": 8,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "WANVAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "context_options",
|
||||
"type": "WANVIDCONTEXT",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
57
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoContextOptions"
|
||||
},
|
||||
"widgets_values": [
|
||||
"uniform_standard",
|
||||
81,
|
||||
4,
|
||||
16,
|
||||
true,
|
||||
false,
|
||||
6,
|
||||
2
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 22,
|
||||
"type": "WanVideoModelLoader",
|
||||
"pos": [
|
||||
620.3950805664062,
|
||||
-357.8426818847656
|
||||
],
|
||||
"size": [
|
||||
477.4410095214844,
|
||||
226.43276977539062
|
||||
],
|
||||
"flags": {},
|
||||
"order": 9,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "compile_args",
|
||||
"shape": 7,
|
||||
"type": "WANCOMPILEARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "block_swap_args",
|
||||
"shape": 7,
|
||||
"type": "BLOCKSWAPARGS",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "lora",
|
||||
"shape": 7,
|
||||
"type": "WANVIDLORA",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vram_management_args",
|
||||
"shape": 7,
|
||||
"type": "VRAM_MANAGEMENTARGS",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "model",
|
||||
"type": "WANVIDEOMODEL",
|
||||
"slot_index": 0,
|
||||
"links": [
|
||||
29
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "WanVideoModelLoader"
|
||||
},
|
||||
"widgets_values": [
|
||||
"WanVideo\\wan2.1_t2v_1.3B_fp16.safetensors",
|
||||
"fp16",
|
||||
"disabled",
|
||||
"offload_device",
|
||||
"sdpa"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
15,
|
||||
11,
|
||||
0,
|
||||
16,
|
||||
0,
|
||||
"WANTEXTENCODER"
|
||||
],
|
||||
[
|
||||
29,
|
||||
22,
|
||||
0,
|
||||
27,
|
||||
0,
|
||||
"WANVIDEOMODEL"
|
||||
],
|
||||
[
|
||||
30,
|
||||
16,
|
||||
0,
|
||||
27,
|
||||
1,
|
||||
"WANVIDEOTEXTEMBEDS"
|
||||
],
|
||||
[
|
||||
33,
|
||||
27,
|
||||
0,
|
||||
28,
|
||||
1,
|
||||
"LATENT"
|
||||
],
|
||||
[
|
||||
42,
|
||||
37,
|
||||
0,
|
||||
27,
|
||||
2,
|
||||
"WANVIDIMAGE_EMBEDS"
|
||||
],
|
||||
[
|
||||
43,
|
||||
38,
|
||||
0,
|
||||
28,
|
||||
0,
|
||||
"VAE"
|
||||
],
|
||||
[
|
||||
48,
|
||||
28,
|
||||
0,
|
||||
42,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
56,
|
||||
42,
|
||||
0,
|
||||
30,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
57,
|
||||
43,
|
||||
0,
|
||||
27,
|
||||
5,
|
||||
"WANVIDCONTEXT"
|
||||
],
|
||||
[
|
||||
58,
|
||||
45,
|
||||
0,
|
||||
27,
|
||||
6,
|
||||
"TEACACHEARGS"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8140274938684471,
|
||||
"offset": [
|
||||
-122.25834160503663,
|
||||
993.5739491626379
|
||||
]
|
||||
},
|
||||
"node_versions": {
|
||||
"ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd",
|
||||
"ComfyUI-KJNodes": "a5bd3c86c8ed6b83c55c2d0e7a59515b15a0137f",
|
||||
"ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8"
|
||||
},
|
||||
"VHS_latentpreview": true,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -167,7 +167,7 @@ class FantasyTalkingWav2VecEmbeds:
|
||||
|
||||
try:
|
||||
audio_segment = audio_input[start_sample:end_sample]
|
||||
except Exception:
|
||||
except:
|
||||
audio_segment = audio_input
|
||||
|
||||
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)
|
||||
previewer = TAESDPreviewerImpl(taesd)
|
||||
previewer = WrappedPreviewer(previewer, rate=16)
|
||||
except Exception:
|
||||
except:
|
||||
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.")
|
||||
method = LatentPreviewMethod.Latent2RGB
|
||||
|
||||
+40
-45
@@ -42,43 +42,41 @@ def rotate_half(x):
|
||||
x = torch.stack((-x2, x1), dim=-1)
|
||||
return rearrange(x, "... d r -> ... (d r)")
|
||||
|
||||
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, split_num=4):
|
||||
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', attn_bias=None):
|
||||
|
||||
ref_k = ref_k.to(visual_q.dtype).to(visual_q.device)
|
||||
scale = 1.0 / visual_q.shape[-1] ** 0.5
|
||||
visual_q = visual_q.transpose(1, 2) * scale
|
||||
visual_q = visual_q * 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 = []
|
||||
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):
|
||||
ref_target_mask = ref_target_mask.view(1, 1, 1, -1)
|
||||
|
||||
x_ref_attnmap = torch.zeros(B, H, x_seqlens, device=visual_q.device, dtype=visual_q.dtype)
|
||||
chunk_size = min(max(x_seqlens // split_num, 1), 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
|
||||
ref_target_mask = ref_target_mask[None, None, None, ...]
|
||||
x_ref_attnmap = x_ref_attn_map_source * ref_target_mask
|
||||
x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens
|
||||
x_ref_attnmap = x_ref_attnmap.permute(0, 2, 1) # B, x_seqlens, H
|
||||
|
||||
if mode == 'mean':
|
||||
x_ref_attnmap = x_ref_attnmap.mean(-1) # B, x_seqlens
|
||||
elif mode == 'max':
|
||||
x_ref_attnmap = x_ref_attnmap.max(-1) # B, x_seqlens
|
||||
|
||||
x_ref_attn_maps.append(x_ref_attnmap)
|
||||
|
||||
del attn, x_ref_attn_map_source
|
||||
|
||||
del visual_q, ref_k
|
||||
|
||||
return torch.cat(x_ref_attn_maps, dim=0)
|
||||
return torch.concat(x_ref_attn_maps, dim=0)
|
||||
|
||||
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2):
|
||||
"""Args:
|
||||
@@ -131,30 +129,27 @@ class RotaryPositionalEmbedding1D(nn.Module):
|
||||
query with the same shape as input.
|
||||
"""
|
||||
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)
|
||||
cos = rearrange(freqs_cis.cos(), 'n d -> 1 1 n d')
|
||||
sin = rearrange(freqs_cis.sin(), 'n d -> 1 1 n d')
|
||||
cos, sin = freqs_cis.cos(), freqs_cis.sin()
|
||||
cos, sin = rearrange(cos, 'n d -> 1 1 n d'), rearrange(sin, 'n d -> 1 1 n d')
|
||||
x_ = (x_ * cos) + (rotate_half(x_) * sin)
|
||||
|
||||
# In-place rotation to save memory
|
||||
x_rotated = rotate_half(x)
|
||||
x.mul_(cos).add_(x_rotated * sin)
|
||||
|
||||
return x.to(in_dtype)
|
||||
return x_.type_as(x)
|
||||
|
||||
class AudioProjModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
seq_len=5,
|
||||
seq_len_vf=8,
|
||||
blocks=12,
|
||||
channels=768,
|
||||
seq_len_vf=12,
|
||||
blocks=12,
|
||||
channels=768,
|
||||
intermediate_dim=512,
|
||||
output_dim=768,
|
||||
context_tokens=32,
|
||||
norm_output_audio=True,
|
||||
norm_output_audio=False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -278,9 +273,9 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
encoder_hidden_states_dim: int,
|
||||
num_heads: int,
|
||||
qkv_bias: bool = True,
|
||||
encoder_hidden_states_dim: int = 768,
|
||||
qkv_bias: bool,
|
||||
class_range: int = 24,
|
||||
class_interval: int = 4,
|
||||
attention_mode: str = 'sdpa',
|
||||
|
||||
@@ -1,569 +0,0 @@
|
||||
import torch
|
||||
import os
|
||||
import gc
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from ..latent_preview import prepare_callback
|
||||
from ..wanvideo.schedulers import get_scheduler
|
||||
from .multitalk import timestep_transform, add_noise
|
||||
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap, match_and_blend_colors
|
||||
from comfy.utils import load_torch_file
|
||||
from ..nodes_model_loading import load_weights
|
||||
from ..HuMo.nodes import get_audio_emb_window
|
||||
import comfy.model_management as mm
|
||||
from tqdm import tqdm
|
||||
import copy
|
||||
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
PATCH_SIZE = (1, 2, 2)
|
||||
vae_upscale_factor = 8
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
def multitalk_loop(self, **kwargs):
|
||||
# Unpack kwargs into local variables
|
||||
(latent, total_steps, steps, start_step, end_step, shift, cfg, denoise_strength,
|
||||
sigmas, weight_dtype, transformer, patcher, block_swap_args, model, vae, dtype,
|
||||
scheduler, scheduler_step_args, text_embeds, image_embeds, multitalk_embeds,
|
||||
multitalk_audio_embeds, unianim_data, dwpose_data, unianimate_poses, uni3c_embeds,
|
||||
humo_image_cond, humo_image_cond_neg, humo_audio, humo_reference_count,
|
||||
add_noise_to_samples, audio_stride, use_tsr, tsr_k, tsr_sigma, fantasy_portrait_input,
|
||||
noise, timesteps, force_offload, add_cond, control_latents, audio_proj,
|
||||
control_camera_latents, samples, masks, seed_g, gguf_reader, predict_func
|
||||
) = (kwargs.get(k) for k in (
|
||||
'latent', 'total_steps', 'steps', 'start_step', 'end_step', 'shift', 'cfg',
|
||||
'denoise_strength', 'sigmas', 'weight_dtype', 'transformer', 'patcher',
|
||||
'block_swap_args', 'model', 'vae', 'dtype', 'scheduler', 'scheduler_step_args',
|
||||
'text_embeds', 'image_embeds', 'multitalk_embeds', 'multitalk_audio_embeds',
|
||||
'unianim_data', 'dwpose_data', 'unianimate_poses', 'uni3c_embeds',
|
||||
'humo_image_cond', 'humo_image_cond_neg', 'humo_audio', 'humo_reference_count',
|
||||
'add_noise_to_samples', 'audio_stride', 'use_tsr', 'tsr_k', 'tsr_sigma',
|
||||
'fantasy_portrait_input', 'noise', 'timesteps', 'force_offload', 'add_cond',
|
||||
'control_latents', 'audio_proj', 'control_camera_latents', 'samples', 'masks',
|
||||
'seed_g', 'gguf_reader', 'predict_with_cfg'
|
||||
))
|
||||
|
||||
mode = image_embeds.get("multitalk_mode", "multitalk")
|
||||
if mode == "auto":
|
||||
mode = transformer.multitalk_model_type.lower()
|
||||
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}")
|
||||
drop_frames = image_embeds.get("drop_frames", 0)
|
||||
cond_frame = None
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
offloaded = False
|
||||
tiled_vae = image_embeds.get("tiled_vae", False)
|
||||
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
|
||||
|
||||
clip_embeds = image_embeds.get("clip_context", None)
|
||||
if clip_embeds is not None:
|
||||
clip_embeds = clip_embeds.to(dtype)
|
||||
colormatch = image_embeds.get("colormatch", "disabled")
|
||||
motion_frame = image_embeds.get("motion_frame", 25)
|
||||
target_w = image_embeds.get("target_w", None)
|
||||
target_h = image_embeds.get("target_h", None)
|
||||
original_images = 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:
|
||||
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
|
||||
|
||||
output_path = image_embeds.get("output_path", "")
|
||||
img_counter = 0
|
||||
|
||||
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
|
||||
face_scale = 0.1
|
||||
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
|
||||
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
|
||||
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
|
||||
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
|
||||
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
|
||||
human_mask2[x_min:x_max, righty_min:righty_max] = 1
|
||||
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
|
||||
human_masks = [human_mask1, human_mask2, background_mask]
|
||||
ref_target_masks = torch.stack(human_masks, dim=0)
|
||||
multitalk_embeds['ref_target_masks'] = ref_target_masks
|
||||
|
||||
gen_video_list = []
|
||||
is_first_clip = True
|
||||
arrive_last_frame = False
|
||||
cur_motion_frames_num = 1
|
||||
audio_start_idx = iteration_count = step_iteration_count = 0
|
||||
audio_end_idx = (audio_start_idx + clip_length) * audio_stride
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
current_condframe_index = 0
|
||||
|
||||
audio_embedding = multitalk_audio_embeds
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
|
||||
uni3c_data = None
|
||||
if uni3c_embeds is not None:
|
||||
transformer.controlnet = uni3c_embeds["controlnet"]
|
||||
uni3c_data = uni3c_embeds.copy()
|
||||
|
||||
encoded_silence = None
|
||||
|
||||
try:
|
||||
silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
|
||||
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
|
||||
except Exception:
|
||||
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
|
||||
|
||||
total_frames = len(audio_embedding[0])
|
||||
estimated_iterations = total_frames // (frame_num - motion_frame - drop_frames) + 1
|
||||
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:
|
||||
arrive_last_frame = True
|
||||
estimated_iterations = 1
|
||||
|
||||
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
|
||||
while True: # start video generation iteratively
|
||||
self.cache_state = [None, None]
|
||||
|
||||
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)
|
||||
if mode == "infinitetalk":
|
||||
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
|
||||
if multitalk_embeds is not None:
|
||||
audio_embs = []
|
||||
# split audio with window size
|
||||
for human_idx in range(human_num):
|
||||
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
|
||||
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1)
|
||||
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
|
||||
audio_embs.append(audio_emb)
|
||||
audio_embs = torch.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)
|
||||
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
|
||||
latent_frame_num = (frame_num - 1) // 4 + 1
|
||||
|
||||
noise = torch.randn(16, latent_frame_num, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
|
||||
|
||||
# Calculate the correct latent slice based on current iteration
|
||||
if is_first_clip:
|
||||
latent_start_idx = 0
|
||||
latent_end_idx = noise.shape[1]
|
||||
else:
|
||||
new_frames_per_iteration = frame_num - motion_frame
|
||||
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
|
||||
latent_start_idx = iteration_count * new_latent_frames_per_iteration
|
||||
latent_end_idx = latent_start_idx + noise.shape[1]
|
||||
|
||||
if samples is not None:
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
input_samples = samples["samples"]
|
||||
if input_samples is not None:
|
||||
input_samples = input_samples.squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if latent_end_idx > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = latent_end_idx - input_samples.shape[1]
|
||||
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
input_samples = torch.cat([input_samples, last_frame], dim=1)
|
||||
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
|
||||
if noise_mask is not None:
|
||||
original_image = input_samples.to(device)
|
||||
|
||||
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
|
||||
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
# diff diff prep
|
||||
if noise_mask is not None:
|
||||
if len(noise_mask.shape) == 4:
|
||||
noise_mask = noise_mask.squeeze(1)
|
||||
if audio_end_idx > noise_mask.shape[0]:
|
||||
noise_mask = noise_mask.repeat(audio_end_idx // noise_mask.shape[0], 1, 1)
|
||||
noise_mask = noise_mask[audio_start_idx:audio_end_idx]
|
||||
noise_mask = torch.nn.functional.interpolate(
|
||||
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
).repeat(1, noise.shape[0], 1, 1, 1)
|
||||
|
||||
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
|
||||
thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device)
|
||||
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
|
||||
|
||||
# zero padding and vae encode for img cond
|
||||
if cond_image is not None or cond_frame is not None:
|
||||
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
|
||||
cond_frame_num = cond_.shape[2]
|
||||
|
||||
# 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)
|
||||
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames], dim=2)
|
||||
|
||||
# encode
|
||||
vae.to(device)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
|
||||
|
||||
if mode == "infinitetalk":
|
||||
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]
|
||||
else:
|
||||
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
#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[:, :1] = 1
|
||||
y = torch.cat([msk, y]) # 4+C T H W
|
||||
mm.soft_empty_cache()
|
||||
else:
|
||||
y = None
|
||||
latent_motion_frames = noise[:, :1]
|
||||
|
||||
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
|
||||
if humo_image_cond is not None:
|
||||
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
|
||||
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
|
||||
if y is not None:
|
||||
partial_humo_cond_input[:, :1] = y[:, :1]
|
||||
if humo_reference_count > 0:
|
||||
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
|
||||
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
|
||||
|
||||
if humo_audio is not None:
|
||||
if is_first_clip:
|
||||
audio_embs = None
|
||||
|
||||
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
|
||||
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
|
||||
partial_humo_audio[-humo_reference_count:] = 0
|
||||
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
|
||||
|
||||
if scheduler == "multitalk":
|
||||
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
|
||||
timesteps.append(0.)
|
||||
timesteps = [torch.tensor([t], device=device) for t in timesteps]
|
||||
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
|
||||
else:
|
||||
if isinstance(scheduler, dict):
|
||||
sample_scheduler = copy.deepcopy(scheduler["sample_scheduler"])
|
||||
timesteps = scheduler["timesteps"]
|
||||
else:
|
||||
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, denoise_strength, sigmas=sigmas)
|
||||
timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)]
|
||||
|
||||
# sample videos
|
||||
latent = noise
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip and mode != "infinitetalk":
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
|
||||
latent[:, :add_latent.shape[1]] = add_latent
|
||||
del motion_add_noise, add_latent
|
||||
|
||||
if offloaded:
|
||||
# Load weights
|
||||
if transformer.patched_linear and gguf_reader is None:
|
||||
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
|
||||
elif gguf_reader is not None: #handle GGUF
|
||||
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
|
||||
#blockswap init
|
||||
init_blockswap(transformer, block_swap_args, model)
|
||||
|
||||
# Use the appropriate prompt for this section
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
|
||||
positive = [text_embeds["prompt_embeds"][prompt_index]]
|
||||
log.info(f"Using prompt index: {prompt_index}")
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
# uni3c slices
|
||||
if uni3c_embeds is not None:
|
||||
vae.to(device)
|
||||
# Pad original_images if needed
|
||||
num_frames = original_images.shape[2]
|
||||
if audio_end_idx > num_frames:
|
||||
pad_len = audio_end_idx - num_frames
|
||||
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
|
||||
padded_images = torch.cat([original_images, last_frame], dim=2)
|
||||
else:
|
||||
padded_images = original_images
|
||||
render_latent = vae.encode(
|
||||
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
|
||||
device=device, tiled=tiled_vae
|
||||
).to(dtype)
|
||||
|
||||
vae.to(offload_device)
|
||||
uni3c_data['render_latent'] = render_latent
|
||||
|
||||
# unianimate slices
|
||||
partial_unianim_data = None
|
||||
if unianim_data is not None:
|
||||
partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx]
|
||||
partial_unianim_data = {
|
||||
"dwpose": partial_dwpose,
|
||||
"random_ref": unianim_data["random_ref"],
|
||||
"strength": unianimate_poses["strength"],
|
||||
"start_percent": unianimate_poses["start_percent"],
|
||||
"end_percent": unianimate_poses["end_percent"]
|
||||
}
|
||||
|
||||
# fantasy portrait slices
|
||||
partial_fantasy_portrait_input = None
|
||||
if fantasy_portrait_input is not None:
|
||||
adapter_proj = fantasy_portrait_input["adapter_proj"]
|
||||
if latent_end_idx > adapter_proj.shape[1]:
|
||||
pad_len = latent_end_idx - adapter_proj.shape[1]
|
||||
last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1)
|
||||
padded_proj = torch.cat([adapter_proj, last_frame], dim=1)
|
||||
else:
|
||||
padded_proj = adapter_proj
|
||||
partial_fantasy_portrait_input = fantasy_portrait_input.copy()
|
||||
partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx]
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
# sampling loop
|
||||
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
|
||||
for i in range(len(timesteps)-1):
|
||||
timestep = timesteps[i]
|
||||
latent_model_input = latent.to(device)
|
||||
if mode == "infinitetalk":
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
noise_pred, _, self.cache_state = predict_func(
|
||||
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, i, y, clip_embeds, control_latents, None, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input,
|
||||
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg,
|
||||
uni3c_data = uni3c_data)
|
||||
|
||||
if callback is not None:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * timestep.to(device) / 1000).detach().permute(1,0,2,3)
|
||||
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
|
||||
del callback_latent
|
||||
|
||||
sampling_pbar.update(1)
|
||||
step_iteration_count += 1
|
||||
|
||||
# update latent
|
||||
if use_tsr:
|
||||
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
|
||||
if scheduler == "multitalk":
|
||||
noise_pred = -noise_pred
|
||||
dt = (timesteps[i] - timesteps[i + 1]) / 1000
|
||||
latent = latent + noise_pred * dt[:, None, None, None]
|
||||
else:
|
||||
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
|
||||
del noise_pred, latent_model_input, timestep
|
||||
|
||||
# differential diffusion inpaint
|
||||
if masks is not None:
|
||||
if i < len(timesteps) - 1:
|
||||
image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1])
|
||||
mask = masks[i].to(latent)
|
||||
latent = image_latent * mask + latent * (1-mask)
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip and mode != "infinitetalk":
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
|
||||
latent[:, :add_latent.shape[1]] = add_latent
|
||||
del motion_add_noise, add_latent
|
||||
elif mode == "infinitetalk":
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
del noise, latent_motion_frames
|
||||
if offload:
|
||||
offload_transformer(transformer, remove_lora=False)
|
||||
offloaded = True
|
||||
if humo_image_cond is not None and humo_reference_count > 0:
|
||||
latent = latent[:,:-humo_reference_count]
|
||||
|
||||
vae.to(device)
|
||||
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
|
||||
vae.to(offload_device)
|
||||
|
||||
sampling_pbar.close()
|
||||
|
||||
# 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)
|
||||
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()
|
||||
from color_matcher import ColorMatcher
|
||||
cm = ColorMatcher()
|
||||
cm_result_list = []
|
||||
for img in videos:
|
||||
if mode == "infinitetalk":
|
||||
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))
|
||||
|
||||
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
|
||||
|
||||
# optionally save generated samples to disk
|
||||
if output_path:
|
||||
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
|
||||
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
|
||||
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
|
||||
start_idx = 0 if is_first_clip else cur_motion_frames_num
|
||||
for i in range(start_idx, video_np.shape[0]):
|
||||
im = Image.fromarray(video_np[i])
|
||||
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
|
||||
img_counter += 1
|
||||
else:
|
||||
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
|
||||
|
||||
current_condframe_index += 1
|
||||
iteration_count += 1
|
||||
|
||||
# decide whether is done
|
||||
if arrive_last_frame:
|
||||
break
|
||||
|
||||
# update next condition frames
|
||||
is_first_clip = False
|
||||
cur_motion_frames_num = motion_frame
|
||||
|
||||
cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0)
|
||||
if mode == "infinitetalk":
|
||||
cond_frame = cond_
|
||||
else:
|
||||
cond_image = cond_
|
||||
|
||||
del videos, latent
|
||||
|
||||
# Repeat audio emb
|
||||
if multitalk_embeds is not None:
|
||||
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count - drop_frames)
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
if audio_end_idx >= len(audio_embedding[0]):
|
||||
arrive_last_frame = True
|
||||
miss_lengths = []
|
||||
source_frames = []
|
||||
for human_inx in range(human_num):
|
||||
source_frame = len(audio_embedding[human_inx])
|
||||
source_frames.append(source_frame)
|
||||
if audio_end_idx >= len(audio_embedding[human_inx]):
|
||||
log.warning(f"Audio embedding for subject {human_inx} not long enough: {len(audio_embedding[human_inx])}, need {audio_end_idx}, padding...")
|
||||
miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3
|
||||
log.warning(f"Padding length: {miss_length}")
|
||||
if encoded_silence is not None:
|
||||
add_audio_emb = encoded_silence[-1*miss_length:]
|
||||
else:
|
||||
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
|
||||
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb.to(device, dtype)], dim=0)
|
||||
miss_lengths.append(miss_length)
|
||||
else:
|
||||
miss_lengths.append(0)
|
||||
if mode == "infinitetalk" and current_condframe_index >= original_images.shape[2]:
|
||||
last_frame = original_images[:, :, -1:, :, :]
|
||||
miss_length = 1
|
||||
original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2)
|
||||
|
||||
if not output_path:
|
||||
gen_video_samples = torch.cat(gen_video_list, dim=1)
|
||||
else:
|
||||
gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output
|
||||
|
||||
if force_offload:
|
||||
if not model["auto_cpu_offload"]:
|
||||
offload_transformer(transformer)
|
||||
try:
|
||||
print_memory(device)
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
pass
|
||||
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
|
||||
+2
-90
@@ -128,7 +128,7 @@ class MultiTalkModelLoader:
|
||||
def loudness_norm(audio_array, sr=16000, lufs=-23):
|
||||
try:
|
||||
import pyloudnorm
|
||||
except Exception:
|
||||
except:
|
||||
raise ImportError("pyloudnorm package is not installed")
|
||||
meter = pyloudnorm.Meter(sr)
|
||||
loudness = meter.integrated_loudness(audio_array)
|
||||
@@ -461,100 +461,13 @@ class WanVideoImageToVideoMultiTalk:
|
||||
}
|
||||
|
||||
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 = {
|
||||
"MultiTalkModelLoader": MultiTalkModelLoader,
|
||||
"MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds,
|
||||
"WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk,
|
||||
"Wav2VecModelLoader": Wav2VecModelLoader,
|
||||
"MultiTalkSilentEmbeds": MultiTalkSilentEmbeds,
|
||||
"WanVideoImageToVideoSkyreelsv3_audio": WanVideoImageToVideoSkyreelsv3_audio,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -563,5 +476,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk",
|
||||
"Wav2VecModelLoader": "Wav2vec2 Model Loader",
|
||||
"MultiTalkSilentEmbeds": "MultiTalk Silent Embeds",
|
||||
"WanVideoImageToVideoSkyreelsv3_audio": "WanVideo Long SkyReelsV3 A2V",
|
||||
}
|
||||
@@ -2,7 +2,6 @@ import os, gc, math
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
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 .taehv import TAEHV
|
||||
@@ -327,7 +326,7 @@ class WanVideoTextEncode:
|
||||
try:
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
encoder = t5["model"]
|
||||
@@ -367,18 +366,10 @@ class WanVideoTextEncode:
|
||||
cast_dtype = encoder.dtype
|
||||
|
||||
params_to_keep = {'norm', 'pos_embedding', 'token_embedding'}
|
||||
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:
|
||||
for name, param in encoder.model.named_parameters():
|
||||
dtype_to_use = dtype if any(keyword in name for keyword in params_to_keep) else cast_dtype
|
||||
value = model_state_dict[name]
|
||||
value = encoder.state_dict[name] if hasattr(encoder, 'state_dict') else encoder.model.state_dict()[name]
|
||||
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'):
|
||||
del encoder.state_dict
|
||||
mm.soft_empty_cache()
|
||||
@@ -502,7 +493,7 @@ class WanVideoTextEncodeSingle:
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
encoder = t5["model"]
|
||||
@@ -559,9 +550,6 @@ class WanVideoApplyNAG:
|
||||
"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}),
|
||||
},
|
||||
"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", )
|
||||
@@ -570,7 +558,7 @@ class WanVideoApplyNAG:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
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, inplace=True):
|
||||
def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha):
|
||||
prompt_embeds_dict_copy = original_text_embeds.copy()
|
||||
prompt_embeds_dict_copy.update({
|
||||
"nag_prompt_embeds": nag_text_embeds["prompt_embeds"],
|
||||
@@ -578,7 +566,6 @@ class WanVideoApplyNAG:
|
||||
"nag_scale": nag_scale,
|
||||
"nag_tau": nag_tau,
|
||||
"nag_alpha": nag_alpha,
|
||||
"inplace": inplace,
|
||||
}
|
||||
})
|
||||
return (prompt_embeds_dict_copy,)
|
||||
@@ -901,86 +888,6 @@ class WanVideoAddMTVMotion:
|
||||
updated["mtv_crafter_motion"] = new_entry
|
||||
return (updated,)
|
||||
|
||||
class WanVideoAddStoryMemLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"vae": ("WANVAE",),
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"memory_images": ("IMAGE",),
|
||||
"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"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, vae, embeds, memory_images, rope_negative_offset, rope_negative_offset_frames):
|
||||
updated = dict(embeds)
|
||||
story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images)
|
||||
updated["story_mem_latents"] = story_mem_latents["samples"].squeeze(2).permute(1, 0, 2, 3) # [C, T, H, W]
|
||||
updated["rope_negative_offset_frames"] = rope_negative_offset_frames if rope_negative_offset else 0
|
||||
return (updated,)
|
||||
|
||||
|
||||
class WanVideoSVIProEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"anchor_samples": ("LATENT", {"tooltip": "Initial start image encoded"}),
|
||||
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
|
||||
},
|
||||
"optional": {
|
||||
"prev_samples": ("LATENT", {"tooltip": "Last latent from previous generation"}),
|
||||
"motion_latent_count": ("INT", {"default": 1, "min": 0, "max": 100, "step": 1, "tooltip": "Number of latents used to continue"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, anchor_samples, num_frames, prev_samples=None, motion_latent_count=1):
|
||||
|
||||
anchor_latent = anchor_samples["samples"][0].clone()
|
||||
|
||||
C, T, H, W = anchor_latent.shape
|
||||
|
||||
total_latents = (num_frames - 1) // 4 + 1
|
||||
device = anchor_latent.device
|
||||
dtype = anchor_latent.dtype
|
||||
|
||||
if prev_samples is None or motion_latent_count == 0:
|
||||
padding_size = total_latents - anchor_latent.shape[1]
|
||||
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
|
||||
y = torch.concat([anchor_latent, padding], dim=1)
|
||||
else:
|
||||
prev_latent = prev_samples["samples"][0].clone()
|
||||
motion_latent = prev_latent[:, -motion_latent_count:]
|
||||
padding_size = total_latents - anchor_latent.shape[1] - motion_latent.shape[1]
|
||||
padding = torch.zeros(C, padding_size, H, W, dtype=dtype, device=device)
|
||||
y = torch.concat([anchor_latent, motion_latent, padding], dim=1)
|
||||
|
||||
msk = torch.ones(1, num_frames, H, W, device=device, dtype=dtype)
|
||||
msk[:, 1:] = 0
|
||||
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
|
||||
msk = msk.view(1, msk.shape[1] // 4, 4, H, W)
|
||||
msk = msk.transpose(1, 2)[0]
|
||||
|
||||
image_embeds = {
|
||||
"image_embeds": y,
|
||||
"num_frames": num_frames,
|
||||
"lat_h": H,
|
||||
"lat_w": W,
|
||||
"mask": msk
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
#region I2V encode
|
||||
class WanVideoImageToVideoEncode:
|
||||
@classmethod
|
||||
@@ -1207,7 +1114,6 @@ class WanVideoAnimateEmbeds:
|
||||
"face_images": ("IMAGE", {"tooltip": "end frame"}),
|
||||
"bg_images": ("IMAGE", {"tooltip": "background images"}),
|
||||
"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"}),
|
||||
}
|
||||
}
|
||||
@@ -1218,7 +1124,7 @@ class WanVideoAnimateEmbeds:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
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, start_ref_image=None):
|
||||
ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None):
|
||||
|
||||
W = (width // 16) * 16
|
||||
H = (height // 16) * 16
|
||||
@@ -1229,7 +1135,7 @@ class WanVideoAnimateEmbeds:
|
||||
num_refs = ref_images.shape[0] if ref_images is not None else 0
|
||||
num_frames = ((num_frames - 1) // 4) * 4 + 1
|
||||
|
||||
looping = num_frames > frame_window_size or start_ref_image is not None
|
||||
looping = num_frames > frame_window_size
|
||||
|
||||
if num_frames < frame_window_size:
|
||||
frame_window_size = num_frames
|
||||
@@ -1327,12 +1233,6 @@ class WanVideoAnimateEmbeds:
|
||||
resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0)
|
||||
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])
|
||||
|
||||
@@ -1352,7 +1252,6 @@ class WanVideoAnimateEmbeds:
|
||||
"is_masked": mask is not None,
|
||||
"ref_latent": ref_latent,
|
||||
"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,
|
||||
"num_frames": num_frames,
|
||||
"target_shape": target_shape,
|
||||
@@ -1927,7 +1826,33 @@ class WanVideoContextOptions:
|
||||
}
|
||||
|
||||
return (context_options,)
|
||||
|
||||
|
||||
class WanVideoFlowEdit:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"source_embeds": ("WANVIDEOTEXTEMBEDS", ),
|
||||
"skip_steps": ("INT", {"default": 4, "min": 0}),
|
||||
"drift_steps": ("INT", {"default": 0, "min": 0}),
|
||||
"drift_flow_shift": ("FLOAT", {"default": 3.0, "min": 1.0, "max": 30.0, "step": 0.01}),
|
||||
"source_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
|
||||
"drift_cfg": ("FLOAT", {"default": 6.0, "min": 0.0, "max": 30.0, "step": 0.01}),
|
||||
},
|
||||
"optional": {
|
||||
"source_image_embeds": ("WANVIDIMAGE_EMBEDS", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOWEDITARGS", )
|
||||
RETURN_NAMES = ("flowedit_args",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Flowedit options for WanVideo"
|
||||
|
||||
def process(self, **kwargs):
|
||||
return (kwargs,)
|
||||
|
||||
class WanVideoLoopArgs:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -2197,7 +2122,7 @@ class WanVideoEncodeLatentBatch:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Encodes a batch of images individually to create a latent video batch where each video is a single frame, useful for I2V init purposes, for example as multiple context window inits"
|
||||
|
||||
def encode(self, vae, images, enable_vae_tiling=False, tile_x=272, tile_y=272, tile_stride_x=144, tile_stride_y=128, latent_strength=1.0):
|
||||
def encode(self, vae, images, enable_vae_tiling, tile_x, tile_y, tile_stride_x, tile_stride_y, latent_strength=1.0):
|
||||
vae.to(device)
|
||||
|
||||
images = images.clone()
|
||||
@@ -2303,6 +2228,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoEnhanceAVideo": WanVideoEnhanceAVideo,
|
||||
"WanVideoContextOptions": WanVideoContextOptions,
|
||||
"WanVideoTextEmbedBridge": WanVideoTextEmbedBridge,
|
||||
"WanVideoFlowEdit": WanVideoFlowEdit,
|
||||
"WanVideoControlEmbeds": WanVideoControlEmbeds,
|
||||
"WanVideoSLG": WanVideoSLG,
|
||||
"WanVideoLoopArgs": WanVideoLoopArgs,
|
||||
@@ -2329,8 +2255,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"TextImageEncodeQwenVL": TextImageEncodeQwenVL,
|
||||
"WanVideoUniLumosEmbeds": WanVideoUniLumosEmbeds,
|
||||
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
|
||||
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
|
||||
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -2346,6 +2270,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video",
|
||||
"WanVideoContextOptions": "WanVideo Context Options",
|
||||
"WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge",
|
||||
"WanVideoFlowEdit": "WanVideo FlowEdit",
|
||||
"WanVideoControlEmbeds": "WanVideo Control Embeds",
|
||||
"WanVideoSLG": "WanVideo SLG",
|
||||
"WanVideoLoopArgs": "WanVideo Loop Args",
|
||||
@@ -2371,6 +2296,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddBindweaveEmbeds": "WanVideo Add Bindweave Embeds",
|
||||
"WanVideoUniLumosEmbeds": "WanVideo UniLumos Embeds",
|
||||
"WanVideoAddTTMLatents": "WanVideo Add TTMLatents",
|
||||
"WanVideoAddStoryMemLatents": "WanVideo Add StoryMem Latents",
|
||||
"WanVideoSVIProEmbeds": "WanVideo SVIPro Embeds",
|
||||
}
|
||||
|
||||
+102
-209
@@ -13,7 +13,7 @@ from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
|
||||
from .custom_linear import _replace_linear
|
||||
|
||||
from accelerate import init_empty_weights
|
||||
from .utils import set_module_tensor_to_device, get_module_memory_mb_per_device
|
||||
from .utils import set_module_tensor_to_device
|
||||
|
||||
import folder_paths
|
||||
import comfy.model_management as mm
|
||||
@@ -23,7 +23,7 @@ from comfy.sd import load_lora_for_models
|
||||
try:
|
||||
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
|
||||
from gguf import GGMLQuantizationType
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
@@ -33,12 +33,9 @@ offload_device = mm.unet_offload_device()
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
except Exception:
|
||||
except:
|
||||
PromptServer = None
|
||||
|
||||
attention_modes = ["sdpa", "flash_attn_2", "flash_attn_3", "sageattn", "sageattn_3", "radial_sage_attention", "sageattn_compiled",
|
||||
"sageattn_ultravico", "comfy"]
|
||||
|
||||
#from city96's gguf nodes
|
||||
def update_folder_names_and_paths(key, targets=[]):
|
||||
# check for existing key
|
||||
@@ -52,22 +49,9 @@ def update_folder_names_and_paths(key, targets=[]):
|
||||
log.warning(f"Unknown file list already present on key {key}: {base}")
|
||||
update_folder_names_and_paths("unet_gguf", ["diffusion_models", "unet"])
|
||||
|
||||
try:
|
||||
from comfy.latent_formats import Wan21, Wan22
|
||||
latent_format = Wan21
|
||||
except: #for backwards compatibility
|
||||
log.warning("WARNING: Wan21 latent format not found, update ComfyUI for better live video preview")
|
||||
from comfy.latent_formats import 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
|
||||
class WanVideoModel(comfy.model_base.BaseModel):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.pipeline = {}
|
||||
|
||||
def __getitem__(self, k):
|
||||
@@ -76,11 +60,24 @@ class WanVideoModel(torch.nn.Module):
|
||||
def __setitem__(self, k, v):
|
||||
self.pipeline[k] = v
|
||||
|
||||
try:
|
||||
from comfy.latent_formats import Wan21, Wan22
|
||||
latent_format = Wan21
|
||||
except: #for backwards compatibility
|
||||
log.warning("WARNING: Wan21 latent format not found, update ComfyUI for better live video preview")
|
||||
from comfy.latent_formats import HunyuanVideo
|
||||
latent_format = HunyuanVideo
|
||||
|
||||
class WanVideoModelConfig:
|
||||
def __init__(self, latent_format=latent_format):
|
||||
def __init__(self, dtype, latent_format=latent_format):
|
||||
self.unet_config = {}
|
||||
self.unet_extra_config = {}
|
||||
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=[]):
|
||||
filtered_dict = {}
|
||||
@@ -135,7 +132,6 @@ def standardize_lora_key_format(lora_sd):
|
||||
k = k.replace('vace_blocks.', 'diffusion_model.vace_blocks.')
|
||||
k = k.replace('.default.', '.')
|
||||
k = k.replace('.diff_m', '.modulation.diff')
|
||||
k = k.replace('base_model.model.', 'diffusion_model.')
|
||||
|
||||
# Fun LoRA format
|
||||
if k.startswith('lora_unet__'):
|
||||
@@ -181,7 +177,7 @@ def standardize_lora_key_format(lora_sd):
|
||||
|
||||
new_key += f".{component}"
|
||||
|
||||
# Handle weight type
|
||||
# Handle weight type - this is the critical fix
|
||||
if weight_type:
|
||||
if weight_type == 'alpha':
|
||||
new_key += '.alpha'
|
||||
@@ -212,12 +208,12 @@ def standardize_lora_key_format(lora_sd):
|
||||
new_key = new_key.replace('time_embedding', 'time.embedding')
|
||||
new_key = new_key.replace('time_projection', 'time.projection')
|
||||
|
||||
# Replace remaining underscores with dots
|
||||
# Replace remaining underscores with dots, carefully
|
||||
parts = new_key.split('.')
|
||||
final_parts = []
|
||||
for part in parts:
|
||||
if part in ['img_emb', 'self_attn', 'cross_attn']:
|
||||
final_parts.append(part)
|
||||
final_parts.append(part) # Keep these intact
|
||||
else:
|
||||
final_parts.append(part.replace('_', '.'))
|
||||
new_key = '.'.join(final_parts)
|
||||
@@ -277,20 +273,6 @@ def standardize_lora_key_format(lora_sd):
|
||||
new_sd[k] = v
|
||||
return new_sd
|
||||
|
||||
def compensate_rs_lora_format(lora_sd):
|
||||
rank = lora_sd["base_model.model.blocks.0.cross_attn.k.lora_A.weight"].shape[0]
|
||||
alpha = torch.tensor(rank * rank // rank ** 0.5)
|
||||
log.info(f"Detected rank stabilized peft lora format with rank {rank}, setting alpha to {alpha} to compensate.")
|
||||
new_sd = {}
|
||||
for k, v in lora_sd.items():
|
||||
if k.endswith(".lora_A.weight"):
|
||||
new_sd[k] = v
|
||||
new_k = k.replace(".lora_A.weight", ".alpha")
|
||||
new_sd[new_k] = alpha
|
||||
else:
|
||||
new_sd[k] = v
|
||||
return new_sd
|
||||
|
||||
class WanVideoBlockSwap:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -381,7 +363,7 @@ class WanVideoLoraSelect:
|
||||
"required": {
|
||||
"lora": (folder_paths.get_filename_list("loras"),
|
||||
{"tooltip": "LORA models are expected to be in ComfyUI/models/loras with .safetensors extension"}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": -1000.0, "max": 1000.0, "step": 0.0001, "tooltip": "LORA strength, set to 0.0 to unmerge the LORA"}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.0001, "tooltip": "LORA strength, set to 0.0 to unmerge the LORA"}),
|
||||
},
|
||||
"optional": {
|
||||
"prev_lora":("WANVIDLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}),
|
||||
@@ -414,7 +396,7 @@ class WanVideoLoraSelect:
|
||||
|
||||
try:
|
||||
lora_path = folder_paths.get_full_path_or_raise("loras", lora)
|
||||
except Exception:
|
||||
except:
|
||||
lora_path = lora
|
||||
|
||||
# Load metadata from the safetensors file
|
||||
@@ -773,8 +755,6 @@ class WanVideoSetLoRAs:
|
||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
||||
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
|
||||
raise NotImplementedError("Unianimate LoRA patching is not implemented in this node.")
|
||||
if "base_model.model.blocks.0.cross_attn.k.lora_A.weight" in lora_sd: # assume rs_lora
|
||||
lora_sd = compensate_rs_lora_format(lora_sd)
|
||||
|
||||
lora_sd = standardize_lora_key_format(lora_sd)
|
||||
if l["blocks"]:
|
||||
@@ -810,6 +790,7 @@ 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"}
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
block_idx = vace_block_idx = None
|
||||
|
||||
if gguf:
|
||||
@@ -829,7 +810,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
all_tensors.extend(r.tensors)
|
||||
for tensor in all_tensors:
|
||||
name = rename_fuser_block(tensor.name)
|
||||
if "glob" not in name and "multitalk_audio_proj" not in name and "audio_proj" in name:
|
||||
if "glob" not in name and "audio_proj" in name:
|
||||
name = name.replace("audio_proj", "multitalk_audio_proj")
|
||||
load_device = device
|
||||
if "vace_blocks." in name:
|
||||
@@ -863,7 +844,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
)
|
||||
transformer.gguf_patched = True
|
||||
else:
|
||||
log.info("Loading and assigning model weights to device...")
|
||||
log.info("Using accelerate to load and assign model weights to device...")
|
||||
named_params = transformer.named_parameters()
|
||||
|
||||
for name, param in tqdm(named_params,
|
||||
@@ -919,17 +900,12 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
load_device = offload_device
|
||||
# Set tensor to device
|
||||
set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=value)
|
||||
pbar.update(1)
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
|
||||
#[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()]
|
||||
memory_on_device = get_module_memory_mb_per_device(transformer)
|
||||
log.info("-" * 25)
|
||||
log.info("Transformer weights loaded:")
|
||||
for dev, mem_mb in memory_on_device.items():
|
||||
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)
|
||||
|
||||
def patch_control_lora(transformer, device):
|
||||
@@ -990,8 +966,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
from .unianimate.nodes import update_transformer
|
||||
log.info("Unianimate LoRA detected, patching model...")
|
||||
patcher.model.diffusion_model, unianimate_sd = update_transformer(patcher.model.diffusion_model, lora_sd)
|
||||
if "base_model.model.blocks.0.cross_attn.k.lora_A.weight" in lora_sd: # assume rs_lora
|
||||
lora_sd = compensate_rs_lora_format(lora_sd)
|
||||
|
||||
lora_sd = standardize_lora_key_format(lora_sd)
|
||||
|
||||
if l["blocks"]:
|
||||
@@ -1013,66 +988,6 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
|
||||
del lora_sd
|
||||
return patcher, control_lora, unianimate_sd
|
||||
|
||||
class WanVideoSetAttentionModeOverride:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("WANVIDEOMODEL", ),
|
||||
"attention_mode": (attention_modes, {"default": "sdpa"}),
|
||||
"start_step": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Step to start applying the attention mode override"}),
|
||||
"end_step": ("INT", {"default": 10000, "min": 1, "max": 10000, "step": 1, "tooltip": "Step to end applying the attention mode override"}),
|
||||
"verbose": ("BOOLEAN", {"default": False, "tooltip": "Print verbose info about attention mode override during generation"}),
|
||||
},
|
||||
"optional": {
|
||||
"blocks":("INT", {"forceInput": True} ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDEOMODEL",)
|
||||
RETURN_NAMES = ("model", )
|
||||
FUNCTION = "getmodelpath"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Override the attention mode for the model for specific step and/or block range"
|
||||
|
||||
def getmodelpath(self, model, attention_mode, start_step, end_step, verbose, blocks=None):
|
||||
model_clone = model.clone()
|
||||
attention_mode_override = {
|
||||
"mode": attention_mode,
|
||||
"start_step": start_step,
|
||||
"end_step": end_step,
|
||||
"verbose": verbose,
|
||||
}
|
||||
if blocks is not None:
|
||||
attention_mode_override["blocks"] = blocks
|
||||
model_clone.model_options['transformer_options']["attention_mode_override"] = attention_mode_override
|
||||
|
||||
return (model_clone,)
|
||||
|
||||
|
||||
class WanVideoUltraVicoSettings:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("WANVIDEOMODEL", ),
|
||||
"alpha": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.001, "tooltip": "Alpha value for the decay, higher values mean slower decay"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDEOMODEL",)
|
||||
RETURN_NAMES = ("model", )
|
||||
FUNCTION = "getmodelpath"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Set UltraVico parameters, attention mode still needs to be set to sageattn_ultravico, https://github.com/thu-ml/DiT-Extrapolation"
|
||||
|
||||
def getmodelpath(self, model, alpha):
|
||||
model_clone = model.clone()
|
||||
model_clone.model_options['transformer_options']["ultravico_alpha"] = alpha
|
||||
|
||||
return (model_clone,)
|
||||
|
||||
|
||||
#region Model loading
|
||||
class WanVideoModelLoader:
|
||||
@classmethod
|
||||
@@ -1087,7 +1002,17 @@ class WanVideoModelLoader:
|
||||
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
},
|
||||
"optional": {
|
||||
"attention_mode": (attention_modes, {"default": "sdpa"}),
|
||||
"attention_mode": ([
|
||||
"sdpa",
|
||||
"flash_attn_2",
|
||||
"flash_attn_3",
|
||||
"sageattn",
|
||||
"sageattn_3",
|
||||
"radial_sage_attention",
|
||||
"sageattn_compiled",
|
||||
"sageattn_ultravico",
|
||||
"comfy"
|
||||
], {"default": "sdpa"}),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
"lora": ("WANVIDLORA", {"default": None}),
|
||||
@@ -1151,7 +1076,7 @@ class WanVideoModelLoader:
|
||||
try:
|
||||
if hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation"):
|
||||
torch.backends.cuda.matmul.allow_fp16_accumulation = False
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
|
||||
@@ -1292,7 +1217,9 @@ class WanVideoModelLoader:
|
||||
lynx_ip_layers = "lite"
|
||||
|
||||
model_type = "t2v"
|
||||
if not "text_embedding.0.weight" in sd:
|
||||
if "audio_injector.injector.0.k.weight" in sd:
|
||||
model_type = "s2v"
|
||||
elif not "text_embedding.0.weight" in sd:
|
||||
model_type = "no_cross_attn" #minimaxremover
|
||||
elif "model_type.Wan2_1-FLF2V-14B-720P" in sd or "img_emb.emb_pos" in sd or "flf2v" in model.lower():
|
||||
model_type = "fl2v"
|
||||
@@ -1302,8 +1229,6 @@ class WanVideoModelLoader:
|
||||
model_type = "t2v"
|
||||
elif "control_adapter.conv.weight" in sd:
|
||||
model_type = "t2v"
|
||||
if "audio_injector.injector.0.k.weight" in sd:
|
||||
model_type = "s2v"
|
||||
|
||||
out_dim = 16
|
||||
if dim == 5120: #14B
|
||||
@@ -1511,57 +1436,7 @@ 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_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
|
||||
# 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:
|
||||
if multitalk_model is not None:
|
||||
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
|
||||
log.info(f"{multitalk_model_type} detected, patching model...")
|
||||
|
||||
@@ -1578,7 +1453,15 @@ class WanVideoModelLoader:
|
||||
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)
|
||||
block.audio_cross_attn = SingleStreamMultiAttention(
|
||||
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_model_type = multitalk_model_type
|
||||
|
||||
@@ -1596,6 +1479,39 @@ class WanVideoModelLoader:
|
||||
|
||||
sd.update(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()}
|
||||
|
||||
@@ -1621,7 +1537,11 @@ 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))
|
||||
|
||||
latent_format=Wan22 if dim == 3072 else Wan21
|
||||
comfy_model = WanVideoModel(WanVideoModelConfig(latent_format=latent_format), device=device, transformer=transformer)
|
||||
comfy_model = WanVideoModel(
|
||||
WanVideoModelConfig(base_dtype, latent_format=latent_format),
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
device=device,
|
||||
)
|
||||
|
||||
# SteadyDancer
|
||||
if "condition_embedding_align.cross_attn.in_proj_bias" in sd:
|
||||
@@ -1685,24 +1605,6 @@ class WanVideoModelLoader:
|
||||
block.ref_attn_v_img = nn.Linear(in_features, out_features)
|
||||
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.load_device = transformer_load_device
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
@@ -1807,14 +1709,10 @@ class WanVideoModelLoader:
|
||||
)
|
||||
|
||||
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}")
|
||||
patcher.model.diffusion_model.to(offload_device)
|
||||
gc.collect()
|
||||
mm.soft_empty_cache()
|
||||
else:
|
||||
log.info(f"Skipping offload (load_device=main_device, keeping model on {patcher.model.diffusion_model.device})")
|
||||
log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}")
|
||||
patcher.model.diffusion_model.to(offload_device)
|
||||
gc.collect()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
patcher.model["base_dtype"] = base_dtype
|
||||
patcher.model["weight_dtype"] = weight_dtype
|
||||
@@ -1887,7 +1785,6 @@ class WanVideoVAELoader:
|
||||
),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"use_cpu_cache": ("BOOLEAN", {"default": False, "tooltip": "Reduces VRAM usage, but slows the VAE down a lot"}),
|
||||
"verbose": ("BOOLEAN", {"default": False, "tooltip": "Enables memory usage logging when using the model"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1897,7 +1794,7 @@ class WanVideoVAELoader:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'"
|
||||
|
||||
def loadmodel(self, model_name, precision, compile_args=None, use_cpu_cache=False, verbose=False):
|
||||
def loadmodel(self, model_name, precision, compile_args=None, use_cpu_cache=False):
|
||||
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
|
||||
model_path = folder_paths.get_full_path_or_raise("vae", model_name)
|
||||
vae_sd = load_torch_file(model_path, safe_load=True)
|
||||
@@ -1914,9 +1811,9 @@ class WanVideoVAELoader:
|
||||
pruning_rate = 0.0
|
||||
|
||||
if vae_sd["model.conv2.weight"].shape[0] == 16:
|
||||
vae = WanVideoVAE(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache, verbose=verbose)
|
||||
vae = WanVideoVAE(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache)
|
||||
elif vae_sd["model.conv2.weight"].shape[0] == 48:
|
||||
vae = WanVideoVAE38(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache, verbose=verbose)
|
||||
vae = WanVideoVAE38(dtype=dtype, pruning_rate=pruning_rate, cpu_cache=use_cpu_cache)
|
||||
|
||||
vae.load_state_dict(vae_sd)
|
||||
del vae_sd
|
||||
@@ -2128,8 +2025,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoTorchCompileSettings": WanVideoTorchCompileSettings,
|
||||
"LoadWanVideoT5TextEncoder": LoadWanVideoT5TextEncoder,
|
||||
"LoadWanVideoClipTextEncoder": LoadWanVideoClipTextEncoder,
|
||||
"WanVideoSetAttentionModeOverride": WanVideoSetAttentionModeOverride,
|
||||
"WanVideoUltraVicoSettings": WanVideoUltraVicoSettings,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -2148,6 +2043,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoTorchCompileSettings": "WanVideo Torch Compile Settings",
|
||||
"LoadWanVideoT5TextEncoder": "WanVideo T5 Text Encoder Loader",
|
||||
"LoadWanVideoClipTextEncoder": "WanVideo CLIP Text Encoder Loader",
|
||||
"WanVideoSetAttentionModeOverride": "WanVideo Set Attention Mode Override",
|
||||
"WanVideoUltraVicoSettings": "WanVideo UltraVico Settings"
|
||||
}
|
||||
|
||||
+846
-212
File diff suppressed because it is too large
Load Diff
+4
-4
@@ -9,7 +9,7 @@ from einops import rearrange
|
||||
|
||||
try:
|
||||
from server import PromptServer
|
||||
except Exception:
|
||||
except:
|
||||
PromptServer = None
|
||||
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
@@ -256,7 +256,7 @@ class CreateCFGScheduleFloatList:
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
@@ -319,7 +319,7 @@ class CreateScheduleFloatList:
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
@@ -454,7 +454,7 @@ class NormalizeAudioLoudness:
|
||||
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
|
||||
try:
|
||||
import pyloudnorm
|
||||
except Exception:
|
||||
except:
|
||||
raise ImportError("pyloudnorm package is not installed")
|
||||
meter = pyloudnorm.Meter(sr)
|
||||
loudness = meter.integrated_loudness(audio_array)
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "ComfyUI-WanVideoWrapper"
|
||||
description = "ComfyUI wrapper nodes for WanVideo"
|
||||
version = "1.4.7"
|
||||
version = "1.4.4"
|
||||
license = {file = "LICENSE"}
|
||||
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()
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
#region main loop start
|
||||
@@ -615,7 +615,7 @@ class WanVideoDiffusionForcingSampler:
|
||||
try:
|
||||
print_memory(device)
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
return ({
|
||||
|
||||
@@ -38,7 +38,7 @@ def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
|
||||
|
||||
qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale
|
||||
|
||||
window_th = frame_tokens * window_width / 2
|
||||
window_th = 1560 * 21 / 2
|
||||
dist2 = tl.abs(m - n).to(tl.int32)
|
||||
dist_mask = dist2 <= window_th
|
||||
|
||||
@@ -46,7 +46,7 @@ def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
|
||||
|
||||
qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor)
|
||||
|
||||
window3 = (m <= frame_tokens) & (n > window_width*frame_tokens)
|
||||
window3 = (m <= frame_tokens) & (n > 21*frame_tokens)
|
||||
qk = tl.where(window3, -1e4, qk)
|
||||
|
||||
|
||||
|
||||
+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:
|
||||
try:
|
||||
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
|
||||
except Exception:
|
||||
except:
|
||||
raise ValueError("No pose detected in reference image")
|
||||
prev_pose = None
|
||||
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)
|
||||
if handle_not_detected == "repeat":
|
||||
prev_pose = pose
|
||||
except Exception:
|
||||
except:
|
||||
if prev_pose is not None:
|
||||
pose = prev_pose
|
||||
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_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
|
||||
result = torch.from_numpy(dwpose_woface)
|
||||
#except Exception:
|
||||
#except:
|
||||
# result = torch.zeros((height, width, 3), dtype=torch.uint8)
|
||||
dwpose_woface_list.append(result)
|
||||
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
|
||||
|
||||
@@ -4,121 +4,17 @@ import logging
|
||||
import math
|
||||
from tqdm import tqdm
|
||||
from pathlib import Path
|
||||
import gc
|
||||
import os
|
||||
import types, collections
|
||||
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
|
||||
from comfy.model_patcher import get_key_weight
|
||||
from comfy.model_patcher import get_key_weight, string_to_seed
|
||||
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.model_management import cast_to_device
|
||||
from comfy.float import stochastic_rounding
|
||||
from .custom_linear import remove_lora_from_module
|
||||
import folder_paths
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
import comfy.model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
try:
|
||||
from .gguf.gguf import GGUFParameter
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
COLOR_CODES = {
|
||||
"reset": "\033[0m",
|
||||
"red": "\033[31m",
|
||||
"green": "\033[32m",
|
||||
"yellow": "\033[33m",
|
||||
"blue": "\033[34m",
|
||||
"magenta": "\033[35m",
|
||||
"cyan": "\033[36m",
|
||||
"white": "\033[37m",
|
||||
}
|
||||
|
||||
def color_text(text, color):
|
||||
try:
|
||||
return f"{COLOR_CODES.get(color, COLOR_CODES['reset'])}{text}{COLOR_CODES['reset']}"
|
||||
except Exception:
|
||||
return text
|
||||
|
||||
class MetaParameter(torch.nn.Parameter):
|
||||
def __new__(cls, dtype, quant_type=None):
|
||||
data = torch.empty(0, dtype=dtype)
|
||||
self = torch.nn.Parameter(data, requires_grad=False)
|
||||
self.quant_type = quant_type
|
||||
return self
|
||||
|
||||
def offload_transformer(transformer, remove_lora=True):
|
||||
transformer.teacache_state.clear_all()
|
||||
transformer.magcache_state.clear_all()
|
||||
transformer.easycache_state.clear_all()
|
||||
|
||||
if transformer.patched_linear:
|
||||
for name, param in transformer.named_parameters():
|
||||
if "loras" in name or "controlnet" in name:
|
||||
continue
|
||||
module = transformer
|
||||
subnames = name.split('.')
|
||||
for subname in subnames[:-1]:
|
||||
module = getattr(module, subname)
|
||||
attr_name = subnames[-1]
|
||||
if param.data.is_floating_point():
|
||||
meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False)
|
||||
setattr(module, attr_name, meta_param)
|
||||
elif isinstance(param.data, GGUFParameter):
|
||||
quant_type = getattr(param, 'quant_type', None)
|
||||
setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type))
|
||||
else:
|
||||
pass
|
||||
if remove_lora:
|
||||
remove_lora_from_module(transformer)
|
||||
else:
|
||||
transformer.to(offload_device)
|
||||
|
||||
for block in transformer.blocks:
|
||||
block.kv_cache = None
|
||||
if transformer.audio_model is not None and hasattr(block, 'audio_block'):
|
||||
block.audio_block = None
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
|
||||
def init_blockswap(transformer, block_swap_args, model):
|
||||
if not transformer.patched_linear:
|
||||
if block_swap_args is not None:
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name or "control_adapter" in name or "face" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
)
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
for block in transformer.blocks:
|
||||
block.modulation = torch.nn.Parameter(block.modulation.to(device))
|
||||
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
|
||||
else:
|
||||
transformer.to(device)
|
||||
|
||||
def check_device_same(first_device, second_device):
|
||||
if first_device.type != second_device.type:
|
||||
return False
|
||||
@@ -195,9 +91,9 @@ def set_module_tensor_to_device(module, tensor_name, device, value=None, dtype=N
|
||||
device = device_quantization
|
||||
if is_buffer:
|
||||
module._buffers[tensor_name] = new_value
|
||||
elif value is not None or not check_device_same(device, module._parameters[tensor_name].device):
|
||||
elif value is not None or not check_device_same(torch.device(device), module._parameters[tensor_name].device):
|
||||
param_cls = type(module._parameters[tensor_name])
|
||||
new_value = param_cls(new_value, requires_grad=False)
|
||||
new_value = param_cls(new_value, requires_grad=False).to(device)
|
||||
module._parameters[tensor_name] = new_value
|
||||
|
||||
#if device != "cpu":
|
||||
@@ -213,8 +109,10 @@ def check_diffusers_version():
|
||||
raise AssertionError("diffusers is not installed.")
|
||||
|
||||
def print_memory(device, process="Sampling"):
|
||||
memory = torch.cuda.memory_allocated(device) / 1024**3
|
||||
max_memory = torch.cuda.max_memory_allocated(device) / 1024**3
|
||||
max_reserved = torch.cuda.max_memory_reserved(device) / 1024**3
|
||||
log.info(f"[{process}] Allocated memory: {memory=:.3f} GB")
|
||||
log.info(f"[{process}] Max allocated memory: {max_memory=:.3f} GB")
|
||||
log.info(f"[{process}] Max reserved memory: {max_reserved=:.3f} GB")
|
||||
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
|
||||
@@ -227,18 +125,6 @@ def get_module_memory_mb(module):
|
||||
memory += param.nelement() * param.element_size()
|
||||
return memory / (1024 * 1024) # Convert to MB
|
||||
|
||||
def get_module_memory_mb_per_device(module):
|
||||
memory_per_device = {}
|
||||
memory = 0
|
||||
for param in module.parameters():
|
||||
if param.data is not None:
|
||||
device = str(param.device)
|
||||
memory += param.nelement() * param.element_size()
|
||||
memory_per_device[device] = memory_per_device.get(device, 0) + memory
|
||||
|
||||
memory_per_device = {dev: mem / (1024 * 1024) for dev, mem in memory_per_device.items()}
|
||||
return memory_per_device
|
||||
|
||||
def get_tensor_memory(tensor):
|
||||
memory_bytes = tensor.element_size() * tensor.nelement()
|
||||
return f"{memory_bytes / (1024 * 1024):.2f} MB"
|
||||
@@ -254,7 +140,7 @@ def patch_weight_to_device(self, key, device_to=None, inplace_update=False, back
|
||||
self.backup[key] = collections.namedtuple('Dimension', ['weight', 'inplace_update'])(weight.to(device=self.offload_device, copy=inplace_update), inplace_update)
|
||||
|
||||
if device_to is not None:
|
||||
temp_weight = mm.cast_to_device(weight, device_to, torch.float32, copy=True)
|
||||
temp_weight = cast_to_device(weight, device_to, torch.float32, copy=True)
|
||||
else:
|
||||
temp_weight = weight.to(torch.float32, copy=True)
|
||||
if convert_func is not None:
|
||||
@@ -309,7 +195,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
key = f"{name.replace('diffusion_model.', '')}.{param}"
|
||||
try:
|
||||
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key])
|
||||
except Exception:
|
||||
except:
|
||||
continue
|
||||
key = f"{name}.{param}"
|
||||
if scale_weights is not None:
|
||||
@@ -323,7 +209,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
if low_mem_load:
|
||||
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])
|
||||
except Exception:
|
||||
except:
|
||||
continue
|
||||
m.comfy_patched_weights = True
|
||||
cnt += 1
|
||||
@@ -352,7 +238,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
|
||||
dtype_to_use = torch.float32
|
||||
try:
|
||||
set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name])
|
||||
except Exception:
|
||||
except:
|
||||
continue
|
||||
return model
|
||||
|
||||
@@ -698,18 +584,17 @@ def check_duplicate_nodes():
|
||||
"""Check ComfyUI custom_nodes directory for duplicate installations"""
|
||||
custom_nodes_dir = Path(folder_paths.folder_names_and_paths["custom_nodes"][0][0])
|
||||
current_path = Path(__file__).parent
|
||||
|
||||
|
||||
wanvideo_dirs = []
|
||||
|
||||
|
||||
# Check all directories in custom_nodes
|
||||
for path in custom_nodes_dir.iterdir():
|
||||
if (path.is_dir() and
|
||||
if (path.is_dir() and
|
||||
path != current_path and
|
||||
not path.name.endswith('.disabled') and
|
||||
'wanvideo' in path.name.lower() and
|
||||
'wrapper' in path.name.lower()):
|
||||
wanvideo_dirs.append(str(path))
|
||||
|
||||
|
||||
return wanvideo_dirs
|
||||
|
||||
#https://github.com/temporalscorerescaling/TSR/
|
||||
@@ -724,55 +609,3 @@ def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.
|
||||
if not t == 1.0:
|
||||
model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t)
|
||||
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,36 +65,36 @@ try:
|
||||
# Return tensor with same shape as q
|
||||
return q.clone()
|
||||
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
|
||||
except Exception:
|
||||
except:
|
||||
sageattn_varlen_func = attention_func_error
|
||||
|
||||
# sage3
|
||||
try:
|
||||
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
|
||||
except Exception:
|
||||
except:
|
||||
try:
|
||||
from sageattn import sageattn_blackwell
|
||||
except Exception:
|
||||
except:
|
||||
sageattn_blackwell = attention_func_error
|
||||
|
||||
try:
|
||||
from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico
|
||||
@torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=())
|
||||
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9, frame_tokens: int = 1536
|
||||
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9
|
||||
) -> torch.Tensor:
|
||||
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor, frame_tokens=frame_tokens)
|
||||
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor)
|
||||
|
||||
@sageattn_func_ultravico.register_fake
|
||||
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
|
||||
return torch.empty_like(qkv[0]).contiguous()
|
||||
sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico
|
||||
except Exception:
|
||||
except:
|
||||
sageattn_func_ultravico = attention_func_error
|
||||
|
||||
|
||||
def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0.,
|
||||
softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16,
|
||||
attention_mode='sdpa', attn_mask=None, transformer_options={}, frame_tokens=1536, heads=128):
|
||||
attention_mode='sdpa', attn_mask=None, multi_factor=0.9, heads=128):
|
||||
if "flash" in attention_mode:
|
||||
return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale,
|
||||
q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3,
|
||||
@@ -108,7 +108,7 @@ def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k
|
||||
elif attention_mode == 'sageattn':
|
||||
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
|
||||
elif attention_mode == 'sageattn_ultravico':
|
||||
return sageattn_func_ultravico([q, k, v], multi_factor=transformer_options.get("ultravico_alpha", 0.9), frame_tokens=frame_tokens).contiguous()
|
||||
return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor).contiguous()
|
||||
elif attention_mode == 'comfy':
|
||||
return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True)
|
||||
else: # sdpa
|
||||
|
||||
+73
-188
@@ -10,7 +10,7 @@ from contextlib import nullcontext
|
||||
|
||||
try:
|
||||
from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense, MaskMap
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
from .attention import attention
|
||||
@@ -467,7 +467,7 @@ class WanSelfAttention(nn.Module):
|
||||
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
def forward(self, q, k, v, seq_lens, transformer_options={}, attention_mode_override=None, lynx_ref_feature=None, lynx_ref_scale=1.0, onetoall_ref=None, onetoall_ref_scale=1.0, frame_tokens=1536):
|
||||
def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None, onetoall_ref=None, onetoall_ref_scale=1.0):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||
@@ -482,7 +482,7 @@ class WanSelfAttention(nn.Module):
|
||||
if self.ref_adapter is not None and lynx_ref_feature is not None:
|
||||
ref_x = self.ref_adapter(self, q, lynx_ref_feature)
|
||||
|
||||
x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads, frame_tokens=frame_tokens, transformer_options=transformer_options)
|
||||
x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads)
|
||||
|
||||
if self.ref_adapter is not None and lynx_ref_feature is not None:
|
||||
x = x.add(ref_x, alpha=lynx_ref_scale)
|
||||
@@ -497,7 +497,7 @@ class WanSelfAttention(nn.Module):
|
||||
attention_mode = self.attention_mode
|
||||
if attention_mode_override is not None:
|
||||
attention_mode = attention_mode_override
|
||||
|
||||
|
||||
# Concatenate main and IP keys/values for main attention
|
||||
full_k = torch.cat([k, k_ip], dim=1)
|
||||
full_v = torch.cat([v, v_ip], dim=1)
|
||||
@@ -558,53 +558,39 @@ class WanSelfAttention(nn.Module):
|
||||
# output
|
||||
return self.o(x.flatten(2))
|
||||
|
||||
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={}):
|
||||
def normalized_attention_guidance(self, b, n, d, q, context, nag_context=None, nag_params={}):
|
||||
# NAG text attention
|
||||
context_positive = context
|
||||
context_negative = nag_context
|
||||
nag_scale = nag_params['nag_scale']
|
||||
nag_alpha = nag_params['nag_alpha']
|
||||
nag_tau = nag_params['nag_tau']
|
||||
inplace = nag_params.get('inplace', True)
|
||||
|
||||
if inplace:
|
||||
nag_guidance = x_negative.mul_(nag_scale - 1).neg_().add_(x_positive, alpha=nag_scale)
|
||||
else:
|
||||
nag_guidance = x_positive * nag_scale - x_negative * (nag_scale - 1)
|
||||
del x_negative
|
||||
k_positive = self.norm_k(self.k(context_positive).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
|
||||
v_positive = self.v(context_positive).view(b, -1, n, d)
|
||||
k_negative = self.norm_k(self.k(context_negative).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype)
|
||||
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)
|
||||
|
||||
norm_positive = torch.norm(x_positive, p=1, dim=-1, keepdim=True)
|
||||
norm_guidance = torch.norm(nag_guidance, p=1, dim=-1, keepdim=True)
|
||||
|
||||
|
||||
scale = norm_guidance / norm_positive
|
||||
torch.nan_to_num_(scale, nan=10.0)
|
||||
scale = torch.nan_to_num(scale, nan=10.0)
|
||||
|
||||
mask = scale > nag_tau
|
||||
del scale
|
||||
|
||||
adjustment = (norm_positive * nag_tau) / (norm_guidance + 1e-7)
|
||||
del norm_positive, norm_guidance
|
||||
|
||||
nag_guidance.mul_(torch.where(mask, adjustment, 1.0))
|
||||
nag_guidance = torch.where(mask, nag_guidance * adjustment, nag_guidance)
|
||||
del mask, adjustment
|
||||
|
||||
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
|
||||
|
||||
return nag_guidance * nag_alpha + x_positive * (1 - nag_alpha)
|
||||
|
||||
class LoRALinearLayer(nn.Module):
|
||||
def __init__(
|
||||
@@ -647,7 +633,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
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",
|
||||
inner_t=None, inner_c=None, cross_freqs=None,
|
||||
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):
|
||||
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):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
s = x.size(1)
|
||||
# compute query
|
||||
@@ -662,10 +648,7 @@ 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)
|
||||
|
||||
if nag_context is not None:
|
||||
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
|
||||
x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
||||
else:
|
||||
if is_longcat:
|
||||
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype).view(b, -1, n, d)).to(x.dtype)
|
||||
@@ -702,7 +685,7 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
# FantasyPortrait adapter attention
|
||||
if adapter_proj is not None:
|
||||
if len(adapter_proj.shape) == 4:
|
||||
q_in = q[:, :orig_seq_len]
|
||||
q_in = q[:, :orig_seq_len]
|
||||
adapter_q = q_in.view(b * num_latent_frames, -1, n, d)
|
||||
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
|
||||
@@ -745,7 +728,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
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",
|
||||
adapter_proj=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -757,22 +740,21 @@ 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)
|
||||
|
||||
if nag_context is not None:
|
||||
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
|
||||
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
||||
else:
|
||||
# text attention
|
||||
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)
|
||||
x = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
|
||||
del k, v
|
||||
x_text = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
|
||||
|
||||
#img attention
|
||||
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)
|
||||
v_img = self.v_img(clip_embed).view(b, -1, n, d)
|
||||
x.add_(attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2))
|
||||
del k_img, v_img
|
||||
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
|
||||
x = x_text + img_x
|
||||
else:
|
||||
x = x_text
|
||||
|
||||
# FantasyTalking audio attention
|
||||
if audio_proj is not None:
|
||||
@@ -805,7 +787,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads)
|
||||
adapter_x = adapter_x.flatten(2)
|
||||
x = x + adapter_x * ip_scale
|
||||
del q
|
||||
|
||||
return self.o(x)
|
||||
|
||||
class WanHuMoCrossAttention(WanSelfAttention):
|
||||
@@ -1024,7 +1006,6 @@ class WanAttentionBlock(nn.Module):
|
||||
longcat_num_cond_latents=0, longcat_avatar_options=None, #longcat image cond amount
|
||||
x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all
|
||||
e_tr=None, tr_num=0, tr_start=0, #token replacement
|
||||
attention_mode_override=None, frame_tokens=None, transformer_options={}
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
@@ -1169,10 +1150,6 @@ class WanAttentionBlock(nn.Module):
|
||||
if enhance_enabled:
|
||||
feta_scores = get_feta_scores(q, k)
|
||||
|
||||
if self.attention_mode == "sageattn_3" and attention_mode_override is None:
|
||||
if current_step != 0 and not last_step:
|
||||
attention_mode_override = "sageattn"
|
||||
|
||||
#self-attention
|
||||
split_attn = (context is not None
|
||||
and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1))
|
||||
@@ -1184,14 +1161,19 @@ class WanAttentionBlock(nn.Module):
|
||||
y = self.self_attn.forward_split(q, k, v, seq_lens, grid_sizes, seq_chunks)
|
||||
elif ref_target_masks is not None: #multi/infinite talk
|
||||
y, x_ref_attn_map = self.self_attn.forward_multitalk(q, k, v, seq_lens, grid_sizes, ref_target_masks)
|
||||
elif self.attention_mode == "radial_sage_attention" or attention_mode_override is not None and attention_mode_override == "radial_sage_attention":
|
||||
elif self.attention_mode == "radial_sage_attention":
|
||||
if self.dense_block or self.dense_timesteps is not None and current_step < self.dense_timesteps:
|
||||
if self.dense_attention_mode == "sparse_sage_attn":
|
||||
y = self.self_attn.forward_radial(q, k, v, dense_step=True)
|
||||
else:
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override=attention_mode_override)
|
||||
y = self.self_attn.forward(q, k, v, seq_lens)
|
||||
else:
|
||||
y = self.self_attn.forward_radial(q, k, v, dense_step=False)
|
||||
elif self.attention_mode == "sageattn_3":
|
||||
if current_step != 0 and not last_step:
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override="sageattn_3")
|
||||
else:
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, attention_mode_override="sageattn")
|
||||
elif x_ip is not None and self.kv_cache is None: #stand-in
|
||||
# First pass: cache IP keys/values and compute attention
|
||||
self.kv_cache = {"k_ip": k_ip.detach(), "v_ip": v_ip.detach()}
|
||||
@@ -1202,18 +1184,18 @@ class WanAttentionBlock(nn.Module):
|
||||
v_ip = self.kv_cache["v_ip"]
|
||||
full_k = torch.cat([k, k_ip], dim=1)
|
||||
full_v = torch.cat([v, v_ip], dim=1)
|
||||
y = self.self_attn.forward(q, full_k, full_v, seq_lens, attention_mode_override=attention_mode_override)
|
||||
y = self.self_attn.forward(q, full_k, full_v, seq_lens)
|
||||
elif is_longcat and longcat_num_cond_latents > 0:
|
||||
if longcat_num_cond_latents == 1:
|
||||
num_cond_latents_thw = longcat_num_cond_latents * (N // num_latent_frames)
|
||||
# process the noise tokens
|
||||
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
|
||||
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
|
||||
# process the condition tokens
|
||||
x_cond = self.self_attn.forward(
|
||||
q[:, :num_cond_latents_thw].contiguous(),
|
||||
k[:, :num_cond_latents_thw].contiguous(),
|
||||
v[:, :num_cond_latents_thw].contiguous(),
|
||||
seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
|
||||
seq_lens)
|
||||
# merge x_cond and x_noise
|
||||
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
|
||||
elif longcat_num_cond_latents > 1: # video continuation
|
||||
@@ -1242,12 +1224,12 @@ class WanAttentionBlock(nn.Module):
|
||||
k_non_ref = k[:, num_ref_latents_thw:].contiguous()
|
||||
v_non_ref = v[:, num_ref_latents_thw:].contiguous()
|
||||
|
||||
x_noise_front = self.self_attn.forward(q_noise_front, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_front has attention with ref + cond + noisy
|
||||
x_noise_back = self.self_attn.forward(q_noise_back, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_back has attention with ref + cond + noisy
|
||||
x_noise_maskref = self.self_attn.forward(q_noise_maskref, k_non_ref, v_non_ref, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_mask has attention with cond+noisy
|
||||
x_noise_front = self.self_attn.forward(q_noise_front, k, v, seq_lens) # q_front has attention with ref + cond + noisy
|
||||
x_noise_back = self.self_attn.forward(q_noise_back, k, v, seq_lens) # q_back has attention with ref + cond + noisy
|
||||
x_noise_maskref = self.self_attn.forward(q_noise_maskref, k_non_ref, v_non_ref, seq_lens) # q_mask has attention with cond+noisy
|
||||
x_noise = torch.cat([x_noise_front, x_noise_maskref, x_noise_back], dim=1).contiguous()
|
||||
else:
|
||||
x_noise = self.self_attn.forward(q_noise, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
|
||||
x_noise = self.self_attn.forward(q_noise, k, v, seq_lens)
|
||||
# process the condition tokens
|
||||
q_ref = q[:, :num_ref_latents_thw].contiguous()
|
||||
k_ref = k[:, :num_ref_latents_thw].contiguous()
|
||||
@@ -1255,14 +1237,13 @@ class WanAttentionBlock(nn.Module):
|
||||
q_cond = q[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
|
||||
k_cond = k[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
|
||||
v_cond = v[:, num_ref_latents_thw:num_cond_latents_thw].contiguous()
|
||||
x_ref = self.self_attn.forward(q_ref, k_ref, v_ref, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
|
||||
x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options)
|
||||
x_ref = self.self_attn.forward(q_ref, k_ref, v_ref, seq_lens)
|
||||
x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens)
|
||||
|
||||
# merge x_cond and x_noise
|
||||
y = torch.cat([x_ref, x_cond, x_noise], dim=1).contiguous()
|
||||
else:
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale,
|
||||
onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale, attention_mode_override=attention_mode_override, transformer_options=transformer_options, frame_tokens=frame_tokens)
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale, onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale)
|
||||
|
||||
del q, k, v
|
||||
|
||||
@@ -1298,7 +1279,7 @@ class WanAttentionBlock(nn.Module):
|
||||
y[:, tr_end:] * gate_msa
|
||||
], dim=1).to(input_dtype)
|
||||
else:
|
||||
x.addcmul_(y, gate_msa)
|
||||
x = x.addcmul(y, gate_msa)
|
||||
del y, gate_msa
|
||||
|
||||
# cross-attention & ffn function
|
||||
@@ -1328,10 +1309,11 @@ class WanAttentionBlock(nn.Module):
|
||||
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
|
||||
else:
|
||||
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 = 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,
|
||||
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).to(input_dtype)
|
||||
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)
|
||||
x = x.to(input_dtype)
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
|
||||
@@ -1345,8 +1327,7 @@ class WanAttentionBlock(nn.Module):
|
||||
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,
|
||||
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
x.add_(x_audio, alpha=audio_scale)
|
||||
del x_audio
|
||||
x = x.add(x_audio, alpha=audio_scale)
|
||||
|
||||
# MTV-Crafter Motion Attention
|
||||
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
|
||||
@@ -2037,24 +2018,24 @@ class WanModel(torch.nn.Module):
|
||||
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False):
|
||||
# Clamp blocks_to_swap to valid range
|
||||
blocks_to_swap = max(0, min(blocks_to_swap, len(self.blocks)))
|
||||
|
||||
|
||||
log.info(f"Swapping {blocks_to_swap} transformer blocks")
|
||||
self.blocks_to_swap = blocks_to_swap
|
||||
self.prefetch_blocks = prefetch_blocks
|
||||
self.block_swap_debug = block_swap_debug
|
||||
|
||||
|
||||
self.offload_img_emb = offload_img_emb
|
||||
self.offload_txt_emb = offload_txt_emb
|
||||
|
||||
total_offload_memory = 0
|
||||
total_main_memory = 0
|
||||
|
||||
|
||||
# Calculate the index where swapping starts
|
||||
swap_start_idx = len(self.blocks) - blocks_to_swap
|
||||
|
||||
|
||||
for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
|
||||
if b < swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
@@ -2069,13 +2050,13 @@ class WanModel(torch.nn.Module):
|
||||
# Clamp vace_blocks_to_swap to valid range
|
||||
vace_blocks_to_swap = max(0, min(vace_blocks_to_swap, len(self.vace_blocks)))
|
||||
self.vace_blocks_to_swap = vace_blocks_to_swap
|
||||
|
||||
|
||||
# Calculate the index where VACE swapping starts
|
||||
vace_swap_start_idx = len(self.vace_blocks) - vace_blocks_to_swap
|
||||
|
||||
for b, block in tqdm(enumerate(self.vace_blocks), total=len(self.vace_blocks), desc="Initializing vace block swap"):
|
||||
block_memory = get_module_memory_mb(block)
|
||||
|
||||
|
||||
if b < vace_swap_start_idx:
|
||||
block.to(self.main_device)
|
||||
total_main_memory += block_memory
|
||||
@@ -2086,13 +2067,13 @@ class WanModel(torch.nn.Module):
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
log.info("-" * 25)
|
||||
log.info("Block swap memory summary:")
|
||||
log.info("----------------------")
|
||||
log.info(f"Block swap memory summary:")
|
||||
log.info(f"Transformer blocks on {self.offload_device}: {total_offload_memory:.2f}MB")
|
||||
log.info(f"Transformer blocks on {self.main_device}: {total_main_memory:.2f}MB")
|
||||
log.info(f"Total memory used by transformer blocks: {(total_offload_memory + total_main_memory):.2f}MB")
|
||||
log.info(f"Non-blocking memory transfer: {self.use_non_blocking}")
|
||||
log.info("-" * 25)
|
||||
log.info("----------------------")
|
||||
|
||||
def forward_vace(
|
||||
self,
|
||||
@@ -2206,7 +2187,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,
|
||||
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, num_memory_frames=3, rope_negative_offset=0):
|
||||
ref_frame_index=10, longcat_num_ref_latents=None):
|
||||
|
||||
patch_size = self.patch_size
|
||||
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
||||
@@ -2230,15 +2211,6 @@ class WanModel(torch.nn.Module):
|
||||
torch.arange(0, steps_t - longcat_num_ref_latents, dtype=dtype, device=device)
|
||||
], dim=0)
|
||||
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:
|
||||
# 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)
|
||||
@@ -2346,10 +2318,6 @@ class WanModel(torch.nn.Module):
|
||||
sdancer_input=None, # SteadyDancer
|
||||
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
|
||||
scail_input=None, # SCAIL pose
|
||||
dual_control_input=None, # LongVie2 dual controlnet
|
||||
transformer_options={},
|
||||
rope_negative_offset=0,
|
||||
num_memory_frames=0,
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -2576,16 +2544,6 @@ class WanModel(torch.nn.Module):
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
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
|
||||
if scail_input is not None:
|
||||
scail_pose_latents = scail_input.get("pose_latent", None)
|
||||
@@ -2594,7 +2552,6 @@ class WanModel(torch.nn.Module):
|
||||
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)]
|
||||
seq_len += scail_x[0].shape[1]
|
||||
del scail_x
|
||||
pose_frame_shape = scail_pose_latents.shape
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
|
||||
@@ -2673,8 +2630,6 @@ class WanModel(torch.nn.Module):
|
||||
self.rope_embedder.k,
|
||||
tuple(ntk_alphas),
|
||||
longcat_num_ref_latents,
|
||||
rope_negative_offset,
|
||||
num_memory_frames,
|
||||
)
|
||||
|
||||
# Check cache using key comparison
|
||||
@@ -2690,12 +2645,10 @@ class WanModel(torch.nn.Module):
|
||||
ref_frame_shape=ref_frame_shape,
|
||||
pose_frame_shape=pose_frame_shape,
|
||||
longcat_num_ref_latents=longcat_num_ref_latents,
|
||||
rope_negative_offset=rope_negative_offset,
|
||||
num_memory_frames=num_memory_frames,
|
||||
device=x.device,
|
||||
dtype=x.dtype
|
||||
)
|
||||
tqdm.write("Generated new RoPE frequencies")
|
||||
log.info("Generated new RoPE frequencies")
|
||||
|
||||
if s2v_ref_latent is not None:
|
||||
freqs_ref = self.rope_encode_comfy(
|
||||
@@ -2881,44 +2834,6 @@ class WanModel(torch.nn.Module):
|
||||
chunked_self_attention = False
|
||||
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
|
||||
if multitalk_audio is not None:
|
||||
self.multitalk_audio_proj.to(self.main_device)
|
||||
@@ -3124,7 +3039,6 @@ class WanModel(torch.nn.Module):
|
||||
camera_embed=camera_embed,
|
||||
audio_proj=audio_proj,
|
||||
num_latent_frames = F,
|
||||
frame_tokens=x.shape[1] // F,
|
||||
original_seq_len=self.original_seq_len,
|
||||
enhance_enabled=enhance_enabled,
|
||||
audio_scale=audio_scale,
|
||||
@@ -3152,7 +3066,6 @@ class WanModel(torch.nn.Module):
|
||||
e_tr=e0_token_replace if use_token_replace else None,
|
||||
tr_start=token_replace_start,
|
||||
tr_num=replace_token_num,
|
||||
transformer_options=transformer_options
|
||||
)
|
||||
if self.audio_model is not None:
|
||||
kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)
|
||||
@@ -3212,22 +3125,8 @@ class WanModel(torch.nn.Module):
|
||||
if lynx_ref_buffer is None and lynx_ref_feature_extractor:
|
||||
lynx_ref_buffer = {}
|
||||
|
||||
attn_override_blocks = attention_mode = None
|
||||
attention_mode_override_active = False
|
||||
attention_mode_override = transformer_options.get("attention_mode_override", None)
|
||||
if attention_mode_override is not None:
|
||||
attn_override_blocks = attention_mode_override.get("blocks", range(len(self.blocks)))
|
||||
if attention_mode_override["start_step"] <= current_step < attention_mode_override["end_step"]:
|
||||
attention_mode_override_active = True
|
||||
if attention_mode_override["verbose"]:
|
||||
tqdm.write(f"Applying attention mode override: {attention_mode_override['mode']} at step {current_step} on blocks: {attn_override_blocks if attn_override_blocks is not None else 'all'}")
|
||||
|
||||
for b, block in enumerate(self.blocks):
|
||||
mm.throw_exception_if_processing_interrupted()
|
||||
if attention_mode_override_active and b in attn_override_blocks:
|
||||
attention_mode = attention_mode_override['mode']
|
||||
else:
|
||||
attention_mode = None
|
||||
block_idx = f"{b:02d}"
|
||||
if lynx_ref_buffer is not None and not lynx_ref_feature_extractor:
|
||||
lynx_ref_feature = lynx_ref_buffer.get(block_idx, None)
|
||||
@@ -3271,21 +3170,9 @@ class WanModel(torch.nn.Module):
|
||||
x_onetoall_ref = onetoall_ref_block_samples[b // interval_ref]
|
||||
|
||||
# ---run block----#
|
||||
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, **kwargs)
|
||||
# ---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:
|
||||
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:
|
||||
@@ -3388,10 +3275,8 @@ class WanModel(torch.nn.Module):
|
||||
# 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)
|
||||
|
||||
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,
|
||||
e_tr=e_token_replace.to(x.device) if use_token_replace else None, tr_start=token_replace_start, tr_num=replace_token_num)
|
||||
|
||||
@@ -4,15 +4,15 @@ import torch
|
||||
try:
|
||||
from spas_sage_attn import block_sparse_sage2_attn_cuda
|
||||
sparse_attn_func = block_sparse_sage2_attn_cuda
|
||||
except Exception:
|
||||
except:
|
||||
try:
|
||||
from sparse_sageattn import sparse_sageattn
|
||||
sparse_attn_func = sparse_sageattn
|
||||
except Exception:
|
||||
except:
|
||||
try:
|
||||
from .sparse_sage.core import sparse_sageattn
|
||||
sparse_attn_func = sparse_sageattn
|
||||
except Exception:
|
||||
except:
|
||||
sparse_sageattn = None
|
||||
raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.")
|
||||
|
||||
|
||||
@@ -42,7 +42,7 @@ def _apply_custom_sigmas(sample_scheduler, sigmas, device):
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, denoise_strength=1.0, sigmas=None, log_timesteps=False, enhance_hf=False, **kwargs):
|
||||
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False, enhance_hf=False, **kwargs):
|
||||
timesteps = None
|
||||
if sigmas is not None:
|
||||
steps = len(sigmas) - 1
|
||||
|
||||
@@ -34,7 +34,7 @@ class ERSDEScheduler():
|
||||
sigmas.append(0.0)
|
||||
self.sigmas = torch.FloatTensor(sigmas)
|
||||
self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
|
||||
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
self.step_index = 0
|
||||
self.old_denoised = None
|
||||
self.old_denoised_d = None
|
||||
|
||||
@@ -35,7 +35,9 @@ class FlowMatchSchedulerResMultistep():
|
||||
self.sigmas = torch.FloatTensor(sigmas)
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
#print(f"Timesteps: {self.timesteps}, Sigmas: {self.sigmas}")
|
||||
|
||||
|
||||
def step(self, model_output, timestep, sample):
|
||||
if timestep.ndim == 2:
|
||||
@@ -46,14 +48,14 @@ class FlowMatchSchedulerResMultistep():
|
||||
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
|
||||
else:
|
||||
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
|
||||
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma
|
||||
if (timestep_id + 1 >= len(self.sigmas)).any():
|
||||
sigma_next = torch.tensor(0)
|
||||
else:
|
||||
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
|
||||
|
||||
|
||||
x0_pred = (sample - sigma * model_output)
|
||||
|
||||
if sigma_next == 0 or self.prev_model_output is None:
|
||||
@@ -71,7 +73,7 @@ class FlowMatchSchedulerResMultistep():
|
||||
self.old_sigma_next = sigma_next
|
||||
self.prev_model_output = x0_pred
|
||||
return x
|
||||
|
||||
|
||||
|
||||
def add_noise(self, original_samples, noise, timestep):
|
||||
"""
|
||||
|
||||
+37
-46
@@ -989,8 +989,7 @@ class VideoVAE_(nn.Module):
|
||||
mean=None,
|
||||
inv_std=None,
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False,
|
||||
verbose=False):
|
||||
cpu_cache=False):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -1001,7 +1000,6 @@ class VideoVAE_(nn.Module):
|
||||
self.temperal_upsample = temperal_downsample[::-1]
|
||||
self.mean = mean
|
||||
self.inv_std = inv_std
|
||||
self.verbose = verbose
|
||||
|
||||
# modules
|
||||
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
|
||||
@@ -1061,7 +1059,7 @@ class VideoVAE_(nn.Module):
|
||||
pbar = ProgressBar(iter_)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
|
||||
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
|
||||
@@ -1087,13 +1085,12 @@ class VideoVAE_(nn.Module):
|
||||
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
|
||||
eps = torch.randn_like(std)
|
||||
return mu + std * eps
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE encode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE encode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return mu
|
||||
|
||||
|
||||
@@ -1137,10 +1134,10 @@ class VideoVAE_(nn.Module):
|
||||
pbar = ProgressBar(iter_)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
x = self.conv2(z)
|
||||
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
|
||||
for i in range(iter_):
|
||||
self._conv_idx = [0]
|
||||
if i == 0:
|
||||
out = self.decoder(x[:, :, i:i + 1, :, :],
|
||||
@@ -1157,13 +1154,12 @@ class VideoVAE_(nn.Module):
|
||||
if pbar:
|
||||
pbar.update_absolute(0)
|
||||
self.clear_cache()
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return out
|
||||
|
||||
def reparameterize(self, mu, log_var):
|
||||
@@ -1190,12 +1186,12 @@ class VideoVAE_(nn.Module):
|
||||
|
||||
class WanVideoVAE(nn.Module):
|
||||
|
||||
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False, verbose=False):
|
||||
def __init__(self, z_dim=16, dtype=torch.float32, pruning_rate=0.0, cpu_cache=False):
|
||||
super().__init__()
|
||||
|
||||
self.dtype = dtype
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
|
||||
mean = [
|
||||
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
|
||||
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
|
||||
@@ -1209,7 +1205,7 @@ class WanVideoVAE(nn.Module):
|
||||
self.z_dim = z_dim
|
||||
|
||||
# init model
|
||||
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache, verbose=self.verbose).eval().requires_grad_(False)
|
||||
self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=self.cpu_cache).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 8
|
||||
|
||||
|
||||
@@ -1435,8 +1431,7 @@ class VideoVAE38_(VideoVAE_):
|
||||
mean=None,
|
||||
inv_std=None,
|
||||
pruning_rate=0.0,
|
||||
cpu_cache=False,
|
||||
verbose=False):
|
||||
cpu_cache=False):
|
||||
super(VideoVAE_, self).__init__()
|
||||
self.dim = dim
|
||||
self.z_dim = z_dim
|
||||
@@ -1449,7 +1444,6 @@ class VideoVAE38_(VideoVAE_):
|
||||
self.mean = mean
|
||||
self.inv_std = inv_std
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
|
||||
# modules
|
||||
self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks,
|
||||
@@ -1464,7 +1458,7 @@ class VideoVAE38_(VideoVAE_):
|
||||
self.clear_cache()
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
x = patchify(x, patch_size=2)
|
||||
t = x.shape[2]
|
||||
@@ -1487,13 +1481,12 @@ class VideoVAE38_(VideoVAE_):
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
self.clear_cache()
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return mu
|
||||
|
||||
|
||||
@@ -1502,7 +1495,7 @@ class VideoVAE38_(VideoVAE_):
|
||||
input_shape = z.shape
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
except:
|
||||
pass
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
|
||||
@@ -1526,19 +1519,18 @@ class VideoVAE38_(VideoVAE_):
|
||||
pbar.update(1)
|
||||
out = unpatchify(out, patch_size=2)
|
||||
self.clear_cache()
|
||||
if self.verbose:
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||
print_memory(device, process="WanVAE decode")
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
pass
|
||||
return out
|
||||
|
||||
|
||||
class WanVideoVAE38(WanVideoVAE):
|
||||
|
||||
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False, verbose=False):
|
||||
def __init__(self, z_dim=48, dim=160, dtype=torch.bfloat16, pruning_rate=0.0, cpu_cache=False):
|
||||
super(WanVideoVAE, self).__init__()
|
||||
|
||||
mean = [
|
||||
@@ -1562,8 +1554,7 @@ class WanVideoVAE38(WanVideoVAE):
|
||||
self.dtype = dtype
|
||||
self.z_dim = z_dim
|
||||
self.cpu_cache = cpu_cache
|
||||
self.verbose = verbose
|
||||
|
||||
# init model
|
||||
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache, verbose=verbose).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 16
|
||||
self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std, pruning_rate=pruning_rate, cpu_cache=cpu_cache).eval().requires_grad_(False)
|
||||
self.upsampling_factor = 16
|
||||
Reference in New Issue
Block a user