16 Commits
Author SHA1 Message Date
kijai 088128b224 Don't count .disabled as duplicate 2026-05-24 16:07:12 +03:00
Jukka Seppänen 126819826a Merge pull request #1948 from haosenwang1018/fix/bare-excepts
fix: replace 47 bare excepts with except Exception
2026-05-24 16:02:36 +03:00
kijai 8f1804bf72 Fix multitalk_audio_stride init 2026-05-24 15:59:40 +03:00
kijai 5437b016e3 Initial LongCatAvatar 1.5 support 2026-05-23 23:13:28 +03:00
kijai d18cdb1859 Fix offload on interrupt 2026-05-05 14:00:56 +03:00
haosenwang1018 0d78230336 fix: replace 47 bare excepts with except Exception
Bare except catches KeyboardInterrupt and SystemExit, masking real errors.
2026-02-25 06:04:52 +00:00
Jukka Seppänen df8f3e49da Delete .github/FUNDING.yml 2026-02-22 15:18:05 +02:00
Jukka Seppänen 86ad93d616 Merge pull request #1840 from little6neko/main
feat(wanvideo): add manual start reference support for WanAnimate loop
2026-02-16 15:44:15 +02:00
Jukka Seppänen 309491b269 Merge pull request #1908 from jjdejong/patch-1
Implement conditional offloading for diffusion model
2026-02-16 15:35:45 +02:00
kijai 06122f1e9d Rope offset should be disabled by default 2026-02-16 15:32:54 +02:00
kijai 5d36631795 I2V cross_attn fixes 2026-02-16 15:31:20 +02:00
Jean J. de Jong b99eac73da Implement conditional offloading for diffusion model
Add condition to skip offloading if load_device is main_device.
2026-01-21 17:33:58 +01:00
小六妞儿 0fbcbed06a Merge branch 'kijai:main' into main 2026-01-16 15:52:00 +08:00
小六妞儿 6f9832ed47 Merge branch 'kijai:main' into main 2026-01-15 14:05:52 +08:00
小六妞儿 58c1bcb7ce Merge branch 'kijai:main' into main 2025-12-28 12:59:47 +08:00
小六妞儿 64cbd28e00 feat(wanvideo): add manual start reference support for WanAnimate loop 2025-12-27 16:04:00 +08:00
19 changed files with 306 additions and 85 deletions
-1
View File
@@ -1 +0,0 @@
github: [kijai]
+191 -3
View File
@@ -1,4 +1,5 @@
import torch import torch
import torch.nn.functional as F
from ..utils import log from ..utils import log
import comfy.model_management as mm import comfy.model_management as mm
from comfy_api.latest import io from comfy_api.latest import io
@@ -24,6 +25,8 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
io.Int.Input("ref_mask_frame_range", default=3, min=0, max=20, step=1, tooltip="Larger range can further help mitigate repeated actions, but excessively large values may introduce artifacts"), io.Int.Input("ref_mask_frame_range", default=3, min=0, max=20, step=1, tooltip="Larger range can further help mitigate repeated actions, but excessively large values may introduce artifacts"),
io.Latent.Input("ref_latent", optional=True, tooltip="Reference latent used for consistency, generally should be either the init image, or first latent from first generation"), io.Latent.Input("ref_latent", optional=True, tooltip="Reference latent used for consistency, generally should be either the init image, or first latent from first generation"),
io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"), io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"),
io.Custom("IMAGE").Input("prev_images", optional=True, tooltip="LongCat-Avatar-1.5: decoded frames from the previous segment. When provided together with `vae`, the trailing `overlap` frames are re-encoded through the VAE and used as the overlap conditioning (matches v1.5's use_vcond=False behavior). Leave disconnected for v1.0."),
io.Custom("WANVAE").Input("vae", optional=True, tooltip="LongCat-Avatar-1.5: VAE used to re-encode `prev_images` for the overlap region. Only used when `prev_images` is also provided."),
], ],
outputs=[ outputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"), io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
@@ -32,7 +35,7 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
) )
@classmethod @classmethod
def execute(cls, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed, ref_frame_index, ref_mask_frame_range, ref_latent=None, samples=None) -> io.NodeOutput: def execute(cls, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed, ref_frame_index, ref_mask_frame_range, ref_latent=None, samples=None, prev_images=None, vae=None) -> io.NodeOutput:
new_audio_embed = audio_embeds.copy() new_audio_embed = audio_embeds.copy()
@@ -55,7 +58,20 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
prev_samples = prev_latents["samples"].clone() prev_samples = prev_latents["samples"].clone()
if overlap != 0: if overlap != 0:
latent_overlap = (overlap - 1) // 4 + 1 latent_overlap = (overlap - 1) // 4 + 1
prev_samples = prev_samples[:, :, -latent_overlap:] 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:]
ref_sample = None ref_sample = None
if ref_latent is not None: if ref_latent is not None:
@@ -65,7 +81,7 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
new_latent_frames = (num_frames - 1) // 4 + 1 new_latent_frames = (num_frames - 1) // 4 + 1
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1]) target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
audio_stride = 2 audio_stride = new_audio_embed.get("audio_stride", 2)
indices = torch.arange(2 * 2 + 1) - 2 indices = torch.arange(2 * 2 + 1) - 2
if frames_processed == 0: if frames_processed == 0:
@@ -112,9 +128,181 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
return io.NodeOutput(embeds, samples_slice) return io.NodeOutput(embeds, samples_slice)
class LongCatAvatarWhisperEmbeds:
"""Audio embeds for LongCat-Video-Avatar-1.5 (Whisper-large-v3).
Produces a MULTITALK_EMBEDS dict whose audio_features are shaped [T, 5, 1280]
(5 grouped Whisper layers, 1280-d hidden state), matching the audio stream
the v1.5 AudioProjModel expects. audio_stride is set to 1 to signal v1.5
timing to the consumer nodes (vs. 2 for the v1.0 wav2vec2 path).
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"whisper_model": ("WHISPERMODEL",),
"audio_1": ("AUDIO",),
"normalize_loudness": ("BOOLEAN", {"default": True, "tooltip": "Normalize audio loudness to -23 LUFS before encoding (matches the v1.5 reference pipeline)"}),
"num_frames": ("INT", {"default": 93, "min": 1, "max": 10000, "step": 1, "tooltip": "Total frame count to generate; bounds how much audio is consumed"}),
"fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 60.0, "step": 0.1, "tooltip": "Target video fps. LongCat-Video-Avatar-1.5 is trained at 25 fps."}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done"}),
"multi_audio_type": (["para", "add"], {"default": "para", "tooltip": "'para' overlays speakers in parallel (equal length); 'add' concatenates speakers sequentially with silence padding"}),
},
"optional": {
"audio_2": ("AUDIO",),
"audio_3": ("AUDIO",),
"audio_4": ("AUDIO",),
"ref_target_masks": ("MASK", {"tooltip": "Per-speaker semantic mask(s) in pixel space, one per speaker"}),
},
}
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", "INT",)
RETURN_NAMES = ("multitalk_embeds", "audio", "num_frames",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, whisper_model, audio_1, normalize_loudness, num_frames, fps,
audio_scale, audio_cfg_scale, multi_audio_type,
audio_2=None, audio_3=None, audio_4=None, ref_target_masks=None):
import torchaudio
import numpy as np
from ..multitalk.nodes import loudness_norm
model = whisper_model["model"]
feature_extractor = whisper_model["feature_extractor"]
dtype = whisper_model["dtype"]
sr = 16000
MEL_CHUNK = 750 * 640 # 480000 samples = 30s at 16kHz; matches Whisper's chunk_length
ENC_CHUNK = 3000 # encoder window in mel frames
ENC_FPS = 50 # whisper encoder output frames per second
def linear_interp(features, output_len):
features = features.transpose(1, 2) # [B, D, T]
out = F.interpolate(features, size=output_len, align_corners=True, mode='linear')
return out.transpose(1, 2)
audio_inputs = [a for a in [audio_1, audio_2, audio_3, audio_4] if a is not None]
audio_features_list = []
seq_lengths = []
audio_outputs = []
end_time = num_frames / float(fps)
end_sample = int(end_time * sr)
for audio in audio_inputs:
audio_input = audio["waveform"]
sample_rate = audio["sample_rate"]
if sample_rate != sr:
audio_input = torchaudio.functional.resample(audio_input, sample_rate, sr)
audio_input = audio_input[0][0]
audio_segment = audio_input[:end_sample].cpu().numpy().astype(np.float32)
if normalize_loudness:
audio_segment = loudness_norm(audio_segment, sr=sr)
audio_duration = len(audio_segment) / sr
video_length = int(audio_duration * fps)
if video_length < 1:
continue
mel_chunks = []
for i in range(0, len(audio_segment), MEL_CHUNK):
mel = feature_extractor(audio_segment[i:i + MEL_CHUNK], sampling_rate=sr,
return_tensors="pt").input_features
mel_chunks.append(mel)
mel_features = torch.cat(mel_chunks, dim=-1).to(device=device, dtype=dtype)
model.to(device)
enc_chunks = []
with torch.no_grad():
for i in range(0, mel_features.shape[-1], ENC_CHUNK):
chunk = mel_features[:, :, i:i + ENC_CHUNK]
chunk_hs = model.encoder(chunk, output_hidden_states=True).hidden_states
enc_chunks.append(torch.stack(chunk_hs, dim=2)) # [1, T_enc, n_layers+1, D]
model.to(offload_device)
audio_prompts = torch.cat(enc_chunks, dim=1)
audio_prompts = audio_prompts[:, :video_length * 2]
feat0 = linear_interp(audio_prompts[:, :, 0:8].mean(dim=2), video_length)
feat1 = linear_interp(audio_prompts[:, :, 8:16].mean(dim=2), video_length)
feat2 = linear_interp(audio_prompts[:, :, 16:24].mean(dim=2), video_length)
feat3 = linear_interp(audio_prompts[:, :, 24:32].mean(dim=2), video_length)
feat4 = linear_interp(audio_prompts[:, :, 32], video_length)
audio_emb = torch.stack([feat0, feat1, feat2, feat3, feat4], dim=2)[0] # [T, 5, 1280]
audio_features_list.append(audio_emb.cpu().detach())
seq_lengths.append(audio_emb.shape[0])
waveform_tensor = torch.from_numpy(audio_segment).float().unsqueeze(0).unsqueeze(0)
audio_outputs.append({"waveform": waveform_tensor, "sample_rate": sr})
if len(audio_features_list) == 0:
raise RuntimeError("No valid Whisper audio embeddings extracted, please check inputs")
if len(audio_features_list) > 1:
if multi_audio_type == "para":
max_len = max(seq_lengths)
padded = []
for emb in audio_features_list:
if emb.shape[0] < max_len:
pad = torch.zeros(max_len - emb.shape[0], *emb.shape[1:], dtype=emb.dtype)
emb = torch.cat([emb, pad], dim=0)
padded.append(emb)
audio_features_list = padded
else: # "add"
total_len = sum(seq_lengths)
full_list = []
offset = 0
for emb, length in zip(audio_features_list, seq_lengths):
full = torch.zeros(total_len, *emb.shape[1:], dtype=emb.dtype)
full[offset:offset + length] = emb
full_list.append(full)
offset += length
audio_features_list = full_list
multitalk_embeds = {
"audio_features": audio_features_list,
"audio_scale": audio_scale,
"audio_cfg_scale": audio_cfg_scale,
"ref_target_masks": ref_target_masks,
"audio_stride": 1,
"audio_encoder_type": "whisper",
}
if len(audio_outputs) == 1:
out_audio = audio_outputs[0]
elif multi_audio_type == "para":
max_len = max(a["waveform"].shape[-1] for a in audio_outputs)
mixed = torch.zeros(1, 1, max_len, dtype=audio_outputs[0]["waveform"].dtype)
for a in audio_outputs:
w = a["waveform"]
if w.shape[-1] < max_len:
w = F.pad(w, (0, max_len - w.shape[-1]))
mixed += w
out_audio = {"waveform": mixed, "sample_rate": sr}
else:
total_len = sum(a["waveform"].shape[-1] for a in audio_outputs)
mixed = torch.zeros(1, 1, total_len, dtype=audio_outputs[0]["waveform"].dtype)
offset = 0
for a in audio_outputs:
w = a["waveform"]
mixed[:, :, offset:offset + w.shape[-1]] += w
offset += w.shape[-1]
out_audio = {"waveform": mixed, "sample_rate": sr}
return (multitalk_embeds, out_audio, num_frames)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds, "WanVideoLongCatAvatarExtendEmbeds": WanVideoLongCatAvatarExtendEmbeds,
"LongCatAvatarWhisperEmbeds": LongCatAvatarWhisperEmbeds,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds", "WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
"LongCatAvatarWhisperEmbeds": "LongCat Avatar Whisper Embeds (v1.5)",
} }
+1 -1
View File
@@ -31,7 +31,7 @@ def check_jit_script_function():
f" Qualified name: {qualname}\n" f" Qualified name: {qualname}\n"
f" Defined in: {code_file}:{code_line}\n" f" Defined in: {code_file}:{code_line}\n"
f"This may cause issues with the NLF model.") f"This may cause issues with the NLF model.")
except: except Exception:
log.warning("--------------------------------") log.warning("--------------------------------")
log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, " log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, "
f"this has been modified by another custom node. This may cause issues with the NLF model.") f"this has been modified by another custom node. This may cause issues with the NLF model.")
+1 -1
View File
@@ -6,7 +6,7 @@ try:
for dir_path in duplicate_dirs: for dir_path in duplicate_dirs:
warning_msg += f" - {color_text(dir_path, 'yellow')}\n" warning_msg += f" - {color_text(dir_path, 'yellow')}\n"
log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red")) log.warning(color_text(warning_msg + "Please remove duplicates to avoid possible conflicts.", "red"))
except: except Exception:
pass pass
from .utils import log from .utils import log
+1 -1
View File
@@ -167,7 +167,7 @@ class FantasyTalkingWav2VecEmbeds:
try: try:
audio_segment = audio_input[start_sample:end_sample] audio_segment = audio_input[start_sample:end_sample]
except: except Exception:
audio_segment = audio_input audio_segment = audio_input
print("audio_segment.shape", audio_segment.shape) print("audio_segment.shape", audio_segment.shape)
+1 -1
View File
@@ -85,7 +85,7 @@ def get_previewer(device, latent_format):
taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device) taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device)
previewer = TAESDPreviewerImpl(taesd) previewer = TAESDPreviewerImpl(taesd)
previewer = WrappedPreviewer(previewer, rate=16) previewer = WrappedPreviewer(previewer, rate=16)
except: except Exception:
log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors") log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors")
log.info("Using Latent2RGB previewer instead.") log.info("Using Latent2RGB previewer instead.")
method = LatentPreviewMethod.Latent2RGB method = LatentPreviewMethod.Latent2RGB
+2 -2
View File
@@ -113,7 +113,7 @@ def multitalk_loop(self, **kwargs):
try: try:
silence_path = os.path.join(script_directory, "encoded_silence.safetensors") silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype) encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except: except Exception:
log.warning("No encoded silence file found, padding with end of audio embedding instead.") log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0]) total_frames = len(audio_embedding[0])
@@ -564,6 +564,6 @@ def multitalk_loop(self, **kwargs):
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path}, return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
+1 -1
View File
@@ -128,7 +128,7 @@ class MultiTalkModelLoader:
def loudness_norm(audio_array, sr=16000, lufs=-23): def loudness_norm(audio_array, sr=16000, lufs=-23):
try: try:
import pyloudnorm import pyloudnorm
except: except Exception:
raise ImportError("pyloudnorm package is not installed") raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr) meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array) loudness = meter.integrated_loudness(audio_array)
+12 -4
View File
@@ -327,7 +327,7 @@ class WanVideoTextEncode:
try: try:
log.info(f"Moving video model to {offload_device}") log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device) model_to_offload.model.to(offload_device)
except: except Exception:
pass pass
encoder = t5["model"] encoder = t5["model"]
@@ -502,7 +502,7 @@ class WanVideoTextEncodeSingle:
log.info(f"Moving video model to {offload_device}") log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device) model_to_offload.model.to(offload_device)
mm.soft_empty_cache() mm.soft_empty_cache()
except: except Exception:
pass pass
encoder = t5["model"] encoder = t5["model"]
@@ -1207,6 +1207,7 @@ class WanVideoAnimateEmbeds:
"face_images": ("IMAGE", {"tooltip": "end frame"}), "face_images": ("IMAGE", {"tooltip": "end frame"}),
"bg_images": ("IMAGE", {"tooltip": "background images"}), "bg_images": ("IMAGE", {"tooltip": "background images"}),
"mask": ("MASK", {"tooltip": "mask"}), "mask": ("MASK", {"tooltip": "mask"}),
"start_ref_image": ("IMAGE", {"tooltip": "start ref image"}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
} }
} }
@@ -1217,7 +1218,7 @@ class WanVideoAnimateEmbeds:
CATEGORY = "WanVideoWrapper" CATEGORY = "WanVideoWrapper"
def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength, def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength,
ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None): ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None, start_ref_image=None):
W = (width // 16) * 16 W = (width // 16) * 16
H = (height // 16) * 16 H = (height // 16) * 16
@@ -1228,7 +1229,7 @@ class WanVideoAnimateEmbeds:
num_refs = ref_images.shape[0] if ref_images is not None else 0 num_refs = ref_images.shape[0] if ref_images is not None else 0
num_frames = ((num_frames - 1) // 4) * 4 + 1 num_frames = ((num_frames - 1) // 4) * 4 + 1
looping = num_frames > frame_window_size looping = num_frames > frame_window_size or start_ref_image is not None
if num_frames < frame_window_size: if num_frames < frame_window_size:
frame_window_size = num_frames frame_window_size = num_frames
@@ -1326,6 +1327,12 @@ class WanVideoAnimateEmbeds:
resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0) resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0)
resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype) resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype)
if start_ref_image is not None:
if start_ref_image.shape[1] != H or start_ref_image.shape[2] != W:
resized_start_ref_image = common_upscale(start_ref_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
else:
resized_start_ref_image = start_ref_image.permute(3, 0, 1, 2) # C, T, H, W
resized_start_ref_image = resized_start_ref_image[:3] * 2 - 1
seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1]) seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1])
@@ -1345,6 +1352,7 @@ class WanVideoAnimateEmbeds:
"is_masked": mask is not None, "is_masked": mask is not None,
"ref_latent": ref_latent, "ref_latent": ref_latent,
"ref_image": resized_ref_images if ref_images is not None else None, "ref_image": resized_ref_images if ref_images is not None else None,
"start_ref_image": resized_start_ref_image if start_ref_image is not None else None,
"face_pixels": resized_face_images if face_images is not None else None, "face_pixels": resized_face_images if face_images is not None else None,
"num_frames": num_frames, "num_frames": num_frames,
"target_shape": target_shape, "target_shape": target_shape,
+26 -10
View File
@@ -23,7 +23,7 @@ from comfy.sd import load_lora_for_models
try: try:
from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter
from gguf import GGMLQuantizationType from gguf import GGMLQuantizationType
except: except Exception:
pass pass
script_directory = os.path.dirname(os.path.abspath(__file__)) script_directory = os.path.dirname(os.path.abspath(__file__))
@@ -33,7 +33,7 @@ offload_device = mm.unet_offload_device()
try: try:
from server import PromptServer from server import PromptServer
except: except Exception:
PromptServer = None PromptServer = None
attention_modes = ["sdpa", "flash_attn_2", "flash_attn_3", "sageattn", "sageattn_3", "radial_sage_attention", "sageattn_compiled", attention_modes = ["sdpa", "flash_attn_2", "flash_attn_3", "sageattn", "sageattn_3", "radial_sage_attention", "sageattn_compiled",
@@ -414,7 +414,7 @@ class WanVideoLoraSelect:
try: try:
lora_path = folder_paths.get_full_path_or_raise("loras", lora) lora_path = folder_paths.get_full_path_or_raise("loras", lora)
except: except Exception:
lora_path = lora lora_path = lora
# Load metadata from the safetensors file # Load metadata from the safetensors file
@@ -1151,7 +1151,7 @@ class WanVideoModelLoader:
try: try:
if hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation"): if hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation"):
torch.backends.cuda.matmul.allow_fp16_accumulation = False torch.backends.cuda.matmul.allow_fp16_accumulation = False
except: except Exception:
pass pass
@@ -1512,12 +1512,24 @@ class WanVideoModelLoader:
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False) block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
# LongCat Avatar # LongCat Avatar
if "multitalk_audio_proj.proj1.weight" in sd and "blocks.0.audio_cross_attn.q_norm.weight" in sd: 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...") log.info("MultiTalk/InfiniteTalk model detected, patching model...")
from .multitalk.multitalk import AudioProjModel from .multitalk.multitalk import AudioProjModel
from .wanvideo.modules.model import WanLayerNorm from .wanvideo.modules.model import WanLayerNorm
from .LongCat.layers import SingleStreamAttention 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: for block in transformer.blocks:
with init_empty_weights(): with init_empty_weights():
@@ -1534,7 +1546,7 @@ class WanVideoModelLoader:
class_interval=4, class_interval=4,
attention_mode=attention_mode, attention_mode=attention_mode,
) )
multitalk_proj_model = AudioProjModel() multitalk_proj_model = AudioProjModel(blocks=audio_proj_blocks, channels=audio_proj_channels)
transformer.multitalk_audio_proj = multitalk_proj_model transformer.multitalk_audio_proj = multitalk_proj_model
# SkyreelsV3 # SkyreelsV3
elif "blocks.1.audio_cross_attn.kv_linear.weight" in sd and "audio_proj.proj1.weight" in sd: elif "blocks.1.audio_cross_attn.kv_linear.weight" in sd and "audio_proj.proj1.weight" in sd:
@@ -1795,10 +1807,14 @@ class WanVideoModelLoader:
) )
if merge_loras and lora is not None: if merge_loras and lora is not None:
log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}") # Skip offloading if load_device is main_device (for unified memory systems like AMD Strix Halo)
patcher.model.diffusion_model.to(offload_device) if load_device != "main_device":
gc.collect() log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}")
mm.soft_empty_cache() 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})")
patcher.model["base_dtype"] = base_dtype patcher.model["base_dtype"] = base_dtype
patcher.model["weight_dtype"] = weight_dtype patcher.model["weight_dtype"] = weight_dtype
+26 -17
View File
@@ -564,6 +564,7 @@ class WanVideoSampler:
# MultiTalk # MultiTalk
multitalk_audio_embeds = audio_emb_slice = audio_features_in = None multitalk_audio_embeds = audio_emb_slice = audio_features_in = None
multitalk_audio_stride = None
multitalk_embeds = image_embeds.get("multitalk_embeds", multitalk_embeds) multitalk_embeds = image_embeds.get("multitalk_embeds", multitalk_embeds)
if multitalk_embeds is not None: if multitalk_embeds is not None:
@@ -584,6 +585,7 @@ class WanVideoSampler:
audio_scale = multitalk_embeds.get("audio_scale", 1.0) audio_scale = multitalk_embeds.get("audio_scale", 1.0)
audio_cfg_scale = multitalk_embeds.get("audio_cfg_scale", 1.0) audio_cfg_scale = multitalk_embeds.get("audio_cfg_scale", 1.0)
ref_target_masks = multitalk_embeds.get("ref_target_masks", None) ref_target_masks = multitalk_embeds.get("ref_target_masks", None)
multitalk_audio_stride = multitalk_embeds.get("audio_stride", None)
if not isinstance(audio_cfg_scale, list): if not isinstance(audio_cfg_scale, list):
audio_cfg_scale = [audio_cfg_scale] * (steps + 1) audio_cfg_scale = [audio_cfg_scale] * (steps + 1)
@@ -817,7 +819,11 @@ class WanVideoSampler:
latent_video_length += insert_len latent_video_length += insert_len
longcat_num_cond_latents = len(clean_latent_indices) longcat_num_cond_latents = len(clean_latent_indices)
log.info(f"LongCat num_cond_latents: {longcat_num_cond_latents} num_ref_latents: {longcat_num_ref_latents}") log.info(f"LongCat num_cond_latents: {longcat_num_cond_latents} num_ref_latents: {longcat_num_ref_latents}")
audio_stride = 2 if transformer.is_longcat else 1 # v1.5 (Whisper) embeds set audio_stride=1; v1.0 (wav2vec2) uses 2 for LongCat
if multitalk_audio_stride is not None:
audio_stride = multitalk_audio_stride
else:
audio_stride = 2 if transformer.is_longcat else 1
#controlnet #controlnet
controlnet_latents = controlnet = None controlnet_latents = controlnet = None
@@ -1730,7 +1736,7 @@ class WanVideoSampler:
gc.collect() gc.collect()
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
# Main sampling loop with FreeInit iterations # Main sampling loop with FreeInit iterations
@@ -2182,7 +2188,7 @@ class WanVideoSampler:
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return {"video": gen_video_samples}, return {"video": gen_video_samples},
# region wananimate loop # region wananimate loop
@@ -2206,7 +2212,11 @@ class WanVideoSampler:
bg_images = image_embeds.get("bg_images", None) bg_images = image_embeds.get("bg_images", None)
pose_images = image_embeds.get("pose_images", None) pose_images = image_embeds.get("pose_images", None)
current_ref_images = face_images = face_images_in = None current_ref_images = image_embeds.get("start_ref_image", None)
if current_ref_images is not None:
log.info(
"WanAnimate: Detected manual start reference image, enabling continuous generation across windows.")
face_images = face_images_in = None
if wananim_face_pixels is not None: if wananim_face_pixels is not None:
face_images = tensor_pingpong_pad(wananim_face_pixels, target_len) face_images = tensor_pingpong_pad(wananim_face_pixels, target_len)
@@ -2245,7 +2255,10 @@ class WanVideoSampler:
mm.soft_empty_cache() mm.soft_empty_cache()
mask_reft_len = 0 if start == 0 else refert_num if current_ref_images is not None:
mask_reft_len = refert_num
else:
mask_reft_len = 0 if start == 0 else refert_num
self.cache_state = [None, None] self.cache_state = [None, None]
@@ -2424,7 +2437,7 @@ class WanVideoSampler:
videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
del latent del latent
if start != 0: if start != 0 or current_ref_images is not None:
videos = videos[:, refert_num:] videos = videos[:, refert_num:]
sampling_pbar.close() sampling_pbar.close()
@@ -2476,7 +2489,7 @@ class WanVideoSampler:
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path}, return {"video": gen_video_samples.permute(1, 2, 3, 0), "output_path": output_path},
@@ -2586,10 +2599,10 @@ class WanVideoSampler:
except Exception as e: except Exception as e:
log.error(f"Error during sampling: {e}") log.error(f"Error during sampling: {e}")
if force_offload: raise
if not model["auto_cpu_offload"]: finally:
offload_transformer(transformer) if force_offload and not model["auto_cpu_offload"]:
raise e offload_transformer(transformer)
if phantom_latents is not None: if phantom_latents is not None:
latent = latent[:,:-phantom_latents.shape[1]] latent = latent[:,:-phantom_latents.shape[1]]
@@ -2613,14 +2626,10 @@ class WanVideoSampler:
"magcache_state": transformer.magcache_state, "magcache_state": transformer.magcache_state,
} }
if force_offload:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return ({ return ({
"samples": latent.unsqueeze(0).cpu(), "samples": latent.unsqueeze(0).cpu(),
@@ -2764,7 +2773,7 @@ class WanVideoScheduler:
import io import io
import base64 import base64
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
except: except Exception:
PromptServer = None PromptServer = None
if unique_id and PromptServer is not None: if unique_id and PromptServer is not None:
try: try:
+4 -4
View File
@@ -9,7 +9,7 @@ from einops import rearrange
try: try:
from server import PromptServer from server import PromptServer
except: except Exception:
PromptServer = None PromptServer = None
VAE_STRIDE = (4, 8, 8) VAE_STRIDE = (4, 8, 8)
@@ -256,7 +256,7 @@ class CreateCFGScheduleFloatList:
f"{cfg_list}", f"{cfg_list}",
unique_id unique_id
) )
except: except Exception:
pass pass
return (cfg_list,) return (cfg_list,)
@@ -319,7 +319,7 @@ class CreateScheduleFloatList:
f"{cfg_list}", f"{cfg_list}",
unique_id unique_id
) )
except: except Exception:
pass pass
return (cfg_list,) return (cfg_list,)
@@ -454,7 +454,7 @@ class NormalizeAudioLoudness:
def loudness_norm(self, audio_array, sr=16000, lufs=-23): def loudness_norm(self, audio_array, sr=16000, lufs=-23):
try: try:
import pyloudnorm import pyloudnorm
except: except Exception:
raise ImportError("pyloudnorm package is not installed") raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr) meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array) loudness = meter.integrated_loudness(audio_array)
+2 -2
View File
@@ -548,7 +548,7 @@ class WanVideoDiffusionForcingSampler:
gc.collect() gc.collect()
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
#region main loop start #region main loop start
@@ -615,7 +615,7 @@ class WanVideoDiffusionForcingSampler:
try: try:
print_memory(device) print_memory(device)
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return ({ return ({
+3 -3
View File
@@ -200,7 +200,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
if ref_image is not None: if ref_image is not None:
try: try:
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold) pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
except: except Exception:
raise ValueError("No pose detected in reference image") raise ValueError("No pose detected in reference image")
prev_pose = None prev_pose = None
for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)): for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)):
@@ -208,7 +208,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
pose = dwpose_model(img, score_threshold=score_threshold) pose = dwpose_model(img, score_threshold=score_threshold)
if handle_not_detected == "repeat": if handle_not_detected == "repeat":
prev_pose = pose prev_pose = pose
except: except Exception:
if prev_pose is not None: if prev_pose is not None:
pose = prev_pose pose = prev_pose
else: else:
@@ -675,7 +675,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size,
draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head) draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
result = torch.from_numpy(dwpose_woface) result = torch.from_numpy(dwpose_woface)
#except: #except Exception:
# result = torch.zeros((height, width, 3), dtype=torch.uint8) # result = torch.zeros((height, width, 3), dtype=torch.uint8)
dwpose_woface_list.append(result) dwpose_woface_list.append(result)
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0) dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
+7 -6
View File
@@ -12,7 +12,7 @@ from comfy.lora import calculate_weight
try: try:
from comfy.utils import string_to_seed from comfy.utils import string_to_seed
except: except Exception:
from comfy.model_patcher import string_to_seed from comfy.model_patcher import string_to_seed
from comfy.float import stochastic_rounding from comfy.float import stochastic_rounding
@@ -27,7 +27,7 @@ offload_device = mm.unet_offload_device()
try: try:
from .gguf.gguf import GGUFParameter from .gguf.gguf import GGUFParameter
except: except Exception:
pass pass
COLOR_CODES = { COLOR_CODES = {
@@ -309,7 +309,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
key = f"{name.replace('diffusion_model.', '')}.{param}" key = f"{name.replace('diffusion_model.', '')}.{param}"
try: try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key]) set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[key])
except: except Exception:
continue continue
key = f"{name}.{param}" key = f"{name}.{param}"
if scale_weights is not None: if scale_weights is not None:
@@ -323,7 +323,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
if low_mem_load: if low_mem_load:
try: try:
set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key]) set_module_tensor_to_device(model.model.diffusion_model, key, device=transformer_load_device, dtype=dtype_to_use, value=model.model.diffusion_model.state_dict()[key])
except: except Exception:
continue continue
m.comfy_patched_weights = True m.comfy_patched_weights = True
cnt += 1 cnt += 1
@@ -352,7 +352,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d
dtype_to_use = torch.float32 dtype_to_use = torch.float32
try: try:
set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name]) set_module_tensor_to_device(model.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=state_dict[name])
except: except Exception:
continue continue
return model return model
@@ -703,8 +703,9 @@ def check_duplicate_nodes():
# Check all directories in custom_nodes # Check all directories in custom_nodes
for path in custom_nodes_dir.iterdir(): for path in custom_nodes_dir.iterdir():
if (path.is_dir() and if (path.is_dir() and
path != current_path and path != current_path and
not path.name.endswith('.disabled') and
'wanvideo' in path.name.lower() and 'wanvideo' in path.name.lower() and
'wrapper' in path.name.lower()): 'wrapper' in path.name.lower()):
wanvideo_dirs.append(str(path)) wanvideo_dirs.append(str(path))
+4 -4
View File
@@ -65,16 +65,16 @@ try:
# Return tensor with same shape as q # Return tensor with same shape as q
return q.clone() return q.clone()
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
except: except Exception:
sageattn_varlen_func = attention_func_error sageattn_varlen_func = attention_func_error
# sage3 # sage3
try: try:
from sageattn3 import sageattn3_blackwell as sageattn_blackwell from sageattn3 import sageattn3_blackwell as sageattn_blackwell
except: except Exception:
try: try:
from sageattn import sageattn_blackwell from sageattn import sageattn_blackwell
except: except Exception:
sageattn_blackwell = attention_func_error sageattn_blackwell = attention_func_error
try: try:
@@ -88,7 +88,7 @@ try:
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9): def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
return torch.empty_like(qkv[0]).contiguous() return torch.empty_like(qkv[0]).contiguous()
sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico sageattn_func_ultravico = torch.ops.wanvideo.sageattn_ultravico
except: except Exception:
sageattn_func_ultravico = attention_func_error sageattn_func_ultravico = attention_func_error
+13 -13
View File
@@ -10,7 +10,7 @@ from contextlib import nullcontext
try: try:
from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense, MaskMap from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense, MaskMap
except: except Exception:
pass pass
from .attention import attention from .attention import attention
@@ -647,7 +647,7 @@ class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0, def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy", num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy",
inner_t=None, inner_c=None, cross_freqs=None, inner_t=None, inner_c=None, cross_freqs=None,
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, longcat_num_cond_latents=None, **kwargs): adapter_proj=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, longcat_num_cond_latents=None, **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim b, n, d = x.size(0), self.num_heads, self.head_dim
s = x.size(1) s = x.size(1)
# compute query # compute query
@@ -702,7 +702,7 @@ class WanT2VCrossAttention(WanSelfAttention):
# FantasyPortrait adapter attention # FantasyPortrait adapter attention
if adapter_proj is not None: if adapter_proj is not None:
if len(adapter_proj.shape) == 4: 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) 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_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) ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
@@ -745,7 +745,7 @@ class WanI2VCrossAttention(WanSelfAttention):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy", audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy",
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs): adapter_proj=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
r""" r"""
Args: Args:
x(Tensor): Shape [B, L1, C] x(Tensor): Shape [B, L1, C]
@@ -757,22 +757,22 @@ class WanI2VCrossAttention(WanSelfAttention):
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype) q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype)
if nag_context is not None: if nag_context is not None:
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) x_positive, x_negative = self.nag_attention(b, n, d, q, context, nag_context)
x = self.normalized_attention_guidance(x_positive, x_negative, nag_params)
del x_positive, x_negative
else: else:
# text attention # text attention
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype) k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype)
v = self.v(context).view(b, -1, n, d) v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) x = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
del k, v
#img attention #img attention
if clip_embed is not None: if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype) k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype)
v_img = self.v_img(clip_embed).view(b, -1, n, d) v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) x.add_(attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2))
x_text.add_(img_x) del k_img, v_img
x = x_text
else:
x = x_text
# FantasyTalking audio attention # FantasyTalking audio attention
if audio_proj is not None: if audio_proj is not None:
@@ -805,7 +805,7 @@ class WanI2VCrossAttention(WanSelfAttention):
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads) adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads)
adapter_x = adapter_x.flatten(2) adapter_x = adapter_x.flatten(2)
x = x + adapter_x * ip_scale x = x + adapter_x * ip_scale
del q
return self.o(x) return self.o(x)
class WanHuMoCrossAttention(WanSelfAttention): class WanHuMoCrossAttention(WanSelfAttention):
@@ -2206,7 +2206,7 @@ class WanModel(torch.nn.Module):
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, ref_frame_shape=None, pose_frame_shape=None, def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, ref_frame_shape=None, pose_frame_shape=None,
steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None,
ref_frame_index=10, longcat_num_ref_latents=0, num_memory_frames=3, rope_negative_offset=5): ref_frame_index=10, longcat_num_ref_latents=0, num_memory_frames=3, rope_negative_offset=0):
patch_size = self.patch_size patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0]) t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
+3 -3
View File
@@ -4,15 +4,15 @@ import torch
try: try:
from spas_sage_attn import block_sparse_sage2_attn_cuda from spas_sage_attn import block_sparse_sage2_attn_cuda
sparse_attn_func = block_sparse_sage2_attn_cuda sparse_attn_func = block_sparse_sage2_attn_cuda
except: except Exception:
try: try:
from sparse_sageattn import sparse_sageattn from sparse_sageattn import sparse_sageattn
sparse_attn_func = sparse_sageattn sparse_attn_func = sparse_sageattn
except: except Exception:
try: try:
from .sparse_sage.core import sparse_sageattn from .sparse_sage.core import sparse_sageattn
sparse_attn_func = sparse_sageattn sparse_attn_func = sparse_sageattn
except: except Exception:
sparse_sageattn = None sparse_sageattn = None
raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.") raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.")
+8 -8
View File
@@ -1061,7 +1061,7 @@ class VideoVAE_(nn.Module):
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar): for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
@@ -1092,7 +1092,7 @@ class VideoVAE_(nn.Module):
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}") log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE encode") print_memory(device, process="WanVAE encode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return mu return mu
@@ -1137,7 +1137,7 @@ class VideoVAE_(nn.Module):
pbar = ProgressBar(iter_) pbar = ProgressBar(iter_)
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
x = self.conv2(z) x = self.conv2(z)
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar): for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
@@ -1162,7 +1162,7 @@ class VideoVAE_(nn.Module):
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}") log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode") print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return out return out
@@ -1464,7 +1464,7 @@ class VideoVAE38_(VideoVAE_):
self.clear_cache() self.clear_cache()
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
x = patchify(x, patch_size=2) x = patchify(x, patch_size=2)
t = x.shape[2] t = x.shape[2]
@@ -1492,7 +1492,7 @@ class VideoVAE38_(VideoVAE_):
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}") log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode") print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return mu return mu
@@ -1502,7 +1502,7 @@ class VideoVAE38_(VideoVAE_):
input_shape = z.shape input_shape = z.shape
try: try:
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
z = z / self.inv_std.to(z) + self.mean.to(z) z = z / self.inv_std.to(z) + self.mean.to(z)
@@ -1531,7 +1531,7 @@ class VideoVAE38_(VideoVAE_):
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}") log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
print_memory(device, process="WanVAE decode") print_memory(device, process="WanVAE decode")
torch.cuda.reset_peak_memory_stats(device) torch.cuda.reset_peak_memory_stats(device)
except: except Exception:
pass pass
return out return out