Author SHA1 Message Date
kijai 40dfb93e11 Merge branch 'main' into self-refine-video 2026-02-02 19:53:03 +02:00
kijai 2dc29b7d30 init 2026-01-27 20:34:21 +02:00
19 changed files with 970 additions and 1018 deletions
+1
View File
@@ -0,0 +1 @@
github: [kijai]
+3 -191
View File
@@ -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,7 +32,7 @@ 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()
@@ -58,20 +55,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 +65,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 +112,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 -1
View File
@@ -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.")
+1 -1
View File
@@ -6,7 +6,7 @@ try:
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:
except:
pass
from .utils import log
+1 -1
View File
@@ -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
View File
@@ -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
+2 -2
View File
@@ -113,7 +113,7 @@ def multitalk_loop(self, **kwargs):
try:
silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except Exception:
except:
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0])
@@ -564,6 +564,6 @@ def multitalk_loop(self, **kwargs):
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
except Exception:
except:
pass
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):
try:
import pyloudnorm
except Exception:
except:
raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
+76 -12
View File
@@ -8,6 +8,7 @@ from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, s
from .taehv import TAEHV
from comfy import model_management as mm
from comfy_api.latest import io
from comfy.utils import ProgressBar, common_upscale
from comfy.clip_vision import clip_preprocess, ClipVisionModel
import folder_paths
@@ -327,7 +328,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"]
@@ -502,7 +503,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"]
@@ -1207,7 +1208,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 +1218,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 +1229,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 +1327,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 +1346,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,
@@ -2077,6 +2070,76 @@ class WanVideoAddTTMLatents:
return (updated,)
#region self-refine-video
class WanVideoSelfRefineVideo(io.ComfyNode):
@classmethod
def define_schema(cls):
# Default values for each range
default_ranges = [
(2, 5, 3), # Range 1
(6, 14, 1), # Range 2
(6, 14, 1), # Range 3
(6, 14, 1), # Range 4
(6, 14, 1), # Range 5
]
options = []
for num_ranges in range(1, 6): # 1 to 5 ranges
range_inputs = []
for i in range(1, num_ranges + 1):
start_default, end_default, steps_default = default_ranges[i - 1]
range_inputs.extend([
io.Int.Input(f"start_step{i}", default=start_default, min=0, max=999, step=1, tooltip=f"Start step for range {i}"),
io.Int.Input(f"end_step{i}", default=end_default, min=0, max=999, step=1, tooltip=f"End step for range {i}"),
io.Int.Input(f"steps_{i}", default=steps_default, min=1, max=100, step=1, tooltip=f"Number of P&P steps for range {i}"),
])
options.append(io.DynamicCombo.Option(
key=str(num_ranges),
inputs=range_inputs
))
return io.Schema(
node_id="WanVideoSelfRefineVideo",
category="WanVideoWrapper",
description="https://github.com/agwmon/self-refine-video - Configure stochastic plan for Perturb-and-Project sampling",
inputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Input("embeds", tooltip="Image embeddings to update"),
io.Float.Input(
"uncertainty_threshold",
default=0.25, min=0.0, max=1.0, step=0.01,
tooltip="Lower values make it harder for regions to be considered \"certain\", meaning more pixels will continue being refined. Higher values make it easier to lock in pixels early."
),
io.Float.Input("certain_percentage", default=0.999, min=0.0, max=1.0, step=0.001, tooltip="Higher values = stricter requirement = fewer early stops = more iterations"),
io.DynamicCombo.Input("num_ranges", options=options, display_name="Number of Ranges", tooltip="Number of step ranges to configure for the stochastic plan"),
],
outputs=[
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Updated image embeddings with self-refine parameters"),
],
)
@classmethod
def execute(cls, embeds, uncertainty_threshold, certain_percentage, num_ranges) -> io.NodeOutput:
updated = dict(embeds)
updated["self_refine_uncertainty_threshold"] = uncertainty_threshold
updated["self_refine_certain_percentage"] = certain_percentage
# Build stochastic plan from the dynamic inputs in list format: [(start, end, steps), ...]
stochastic_plan = []
range_keys = sorted([k for k in num_ranges.keys() if k.startswith('start_step')])
for start_key in range_keys:
i = start_key.replace('start_step', '')
start = num_ranges.get(f"start_step{i}")
end = num_ranges.get(f"end_step{i}")
steps = num_ranges.get(f"steps_{i}")
if start is not None and end is not None and steps is not None:
stochastic_plan.append((start, end, steps))
updated["stochastic_plan"] = stochastic_plan
return io.NodeOutput(updated)
#region VideoDecode
class WanVideoDecode:
@classmethod
@@ -2331,6 +2394,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
"WanVideoSelfRefineVideo": WanVideoSelfRefineVideo,
}
NODE_DISPLAY_NAME_MAPPINGS = {
+10 -26
View File
@@ -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,7 +33,7 @@ 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",
@@ -414,7 +414,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
@@ -1151,7 +1151,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
@@ -1512,24 +1512,12 @@ class WanVideoModelLoader:
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):
if "multitalk_audio_proj.proj1.weight" in sd and "blocks.0.audio_cross_attn.q_norm.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
# 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():
@@ -1546,7 +1534,7 @@ class WanVideoModelLoader:
class_interval=4,
attention_mode=attention_mode,
)
multitalk_proj_model = AudioProjModel(blocks=audio_proj_blocks, channels=audio_proj_channels)
multitalk_proj_model = AudioProjModel()
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:
@@ -1807,14 +1795,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
+830 -738
View File
File diff suppressed because it is too large Load Diff
+4 -4
View File
@@ -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)
+2 -2
View File
@@ -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 ({
+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:
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)
+6 -7
View File
@@ -12,7 +12,7 @@ from comfy.lora import calculate_weight
try:
from comfy.utils import string_to_seed
except Exception:
except:
from comfy.model_patcher import string_to_seed
from comfy.float import stochastic_rounding
@@ -27,7 +27,7 @@ offload_device = mm.unet_offload_device()
try:
from .gguf.gguf import GGUFParameter
except Exception:
except:
pass
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}"
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 +323,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 +352,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
@@ -703,9 +703,8 @@ def check_duplicate_nodes():
# 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))
+4 -4
View File
@@ -65,16 +65,16 @@ 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:
@@ -88,7 +88,7 @@ try:
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
+13 -13
View File
@@ -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
@@ -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,
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
@@ -702,7 +702,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 +745,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 +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)
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_text.add_(img_x)
x = x_text
else:
x = x_text
# FantasyTalking audio attention
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 = adapter_x.flatten(2)
x = x + adapter_x * ip_scale
del q
return self.o(x)
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,
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=0, num_memory_frames=3, rope_negative_offset=5):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
+3 -3
View File
@@ -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.")
+8 -8
View File
@@ -1061,7 +1061,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):
@@ -1092,7 +1092,7 @@ class VideoVAE_(nn.Module):
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:
except:
pass
return mu
@@ -1137,7 +1137,7 @@ 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):
@@ -1162,7 +1162,7 @@ class VideoVAE_(nn.Module):
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:
except:
pass
return out
@@ -1464,7 +1464,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]
@@ -1492,7 +1492,7 @@ class VideoVAE38_(VideoVAE_):
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:
except:
pass
return mu
@@ -1502,7 +1502,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)
@@ -1531,7 +1531,7 @@ class VideoVAE38_(VideoVAE_):
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:
except:
pass
return out