Compare commits
16
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
088128b224 | ||
|
|
126819826a | ||
|
|
8f1804bf72 | ||
|
|
5437b016e3 | ||
|
|
d18cdb1859 | ||
|
|
0d78230336 | ||
|
|
df8f3e49da | ||
|
|
86ad93d616 | ||
|
|
309491b269 | ||
|
|
06122f1e9d | ||
|
|
5d36631795 | ||
|
|
b99eac73da | ||
|
|
0fbcbed06a | ||
|
|
6f9832ed47 | ||
|
|
58c1bcb7ce | ||
|
|
64cbd28e00 |
@@ -1 +0,0 @@
|
|||||||
github: [kijai]
|
|
||||||
+191
-3
@@ -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
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -167,7 +167,7 @@ class FantasyTalkingWav2VecEmbeds:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
audio_segment = audio_input[start_sample:end_sample]
|
audio_segment = audio_input[start_sample:end_sample]
|
||||||
except:
|
except Exception:
|
||||||
audio_segment = audio_input
|
audio_segment = audio_input
|
||||||
|
|
||||||
print("audio_segment.shape", audio_segment.shape)
|
print("audio_segment.shape", audio_segment.shape)
|
||||||
|
|||||||
+1
-1
@@ -85,7 +85,7 @@ def get_previewer(device, latent_format):
|
|||||||
taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device)
|
taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device)
|
||||||
previewer = TAESDPreviewerImpl(taesd)
|
previewer = TAESDPreviewerImpl(taesd)
|
||||||
previewer = WrappedPreviewer(previewer, rate=16)
|
previewer = WrappedPreviewer(previewer, rate=16)
|
||||||
except:
|
except Exception:
|
||||||
log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors")
|
log.info("Could not find TAEW model file 'taew2_1.safetensors' from models/vae_approx. You can download it from https://huggingface.co/Kijai/WanVideo_comfy/blob/main/taew2_1.safetensors")
|
||||||
log.info("Using Latent2RGB previewer instead.")
|
log.info("Using Latent2RGB previewer instead.")
|
||||||
method = LatentPreviewMethod.Latent2RGB
|
method = LatentPreviewMethod.Latent2RGB
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ from .utils import(log, clip_encode_image_tiled, add_noise_to_reference_video, s
|
|||||||
from .taehv import TAEHV
|
from .taehv import TAEHV
|
||||||
|
|
||||||
from comfy import model_management as mm
|
from comfy import model_management as mm
|
||||||
from comfy_api.latest import io
|
|
||||||
from comfy.utils import ProgressBar, common_upscale
|
from comfy.utils import ProgressBar, common_upscale
|
||||||
from comfy.clip_vision import clip_preprocess, ClipVisionModel
|
from comfy.clip_vision import clip_preprocess, ClipVisionModel
|
||||||
import folder_paths
|
import folder_paths
|
||||||
@@ -328,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"]
|
||||||
@@ -503,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"]
|
||||||
@@ -1208,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"}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1218,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
|
||||||
@@ -1229,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
|
||||||
@@ -1327,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])
|
||||||
|
|
||||||
@@ -1346,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,
|
||||||
@@ -2070,76 +2077,6 @@ class WanVideoAddTTMLatents:
|
|||||||
|
|
||||||
return (updated,)
|
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
|
#region VideoDecode
|
||||||
class WanVideoDecode:
|
class WanVideoDecode:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -2394,7 +2331,6 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
|
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
|
||||||
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
|
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
|
||||||
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
|
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
|
||||||
"WanVideoSelfRefineVideo": WanVideoSelfRefineVideo,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
|||||||
+26
-10
@@ -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
|
||||||
|
|||||||
+738
-830
File diff suppressed because it is too large
Load Diff
+4
-4
@@ -9,7 +9,7 @@ from einops import rearrange
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
from server import PromptServer
|
from server import PromptServer
|
||||||
except:
|
except Exception:
|
||||||
PromptServer = None
|
PromptServer = None
|
||||||
|
|
||||||
VAE_STRIDE = (4, 8, 8)
|
VAE_STRIDE = (4, 8, 8)
|
||||||
@@ -256,7 +256,7 @@ class CreateCFGScheduleFloatList:
|
|||||||
f"{cfg_list}",
|
f"{cfg_list}",
|
||||||
unique_id
|
unique_id
|
||||||
)
|
)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
return (cfg_list,)
|
return (cfg_list,)
|
||||||
@@ -319,7 +319,7 @@ class CreateScheduleFloatList:
|
|||||||
f"{cfg_list}",
|
f"{cfg_list}",
|
||||||
unique_id
|
unique_id
|
||||||
)
|
)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
return (cfg_list,)
|
return (cfg_list,)
|
||||||
@@ -454,7 +454,7 @@ class NormalizeAudioLoudness:
|
|||||||
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
|
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
|
||||||
try:
|
try:
|
||||||
import pyloudnorm
|
import pyloudnorm
|
||||||
except:
|
except Exception:
|
||||||
raise ImportError("pyloudnorm package is not installed")
|
raise ImportError("pyloudnorm package is not installed")
|
||||||
meter = pyloudnorm.Meter(sr)
|
meter = pyloudnorm.Meter(sr)
|
||||||
loudness = meter.integrated_loudness(audio_array)
|
loudness = meter.integrated_loudness(audio_array)
|
||||||
|
|||||||
+2
-2
@@ -548,7 +548,7 @@ class WanVideoDiffusionForcingSampler:
|
|||||||
gc.collect()
|
gc.collect()
|
||||||
try:
|
try:
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
#region main loop start
|
#region main loop start
|
||||||
@@ -615,7 +615,7 @@ class WanVideoDiffusionForcingSampler:
|
|||||||
try:
|
try:
|
||||||
print_memory(device)
|
print_memory(device)
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
return ({
|
return ({
|
||||||
|
|||||||
+3
-3
@@ -200,7 +200,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
|||||||
if ref_image is not None:
|
if ref_image is not None:
|
||||||
try:
|
try:
|
||||||
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
|
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
|
||||||
except:
|
except Exception:
|
||||||
raise ValueError("No pose detected in reference image")
|
raise ValueError("No pose detected in reference image")
|
||||||
prev_pose = None
|
prev_pose = None
|
||||||
for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)):
|
for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)):
|
||||||
@@ -208,7 +208,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
|||||||
pose = dwpose_model(img, score_threshold=score_threshold)
|
pose = dwpose_model(img, score_threshold=score_threshold)
|
||||||
if handle_not_detected == "repeat":
|
if handle_not_detected == "repeat":
|
||||||
prev_pose = pose
|
prev_pose = pose
|
||||||
except:
|
except Exception:
|
||||||
if prev_pose is not None:
|
if prev_pose is not None:
|
||||||
pose = prev_pose
|
pose = prev_pose
|
||||||
else:
|
else:
|
||||||
@@ -675,7 +675,7 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre
|
|||||||
draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size,
|
draw_body=draw_body, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size,
|
||||||
draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
|
draw_feet=draw_feet, body_keypoint_size=body_keypoint_size, draw_head=draw_head)
|
||||||
result = torch.from_numpy(dwpose_woface)
|
result = torch.from_numpy(dwpose_woface)
|
||||||
#except:
|
#except Exception:
|
||||||
# result = torch.zeros((height, width, 3), dtype=torch.uint8)
|
# result = torch.zeros((height, width, 3), dtype=torch.uint8)
|
||||||
dwpose_woface_list.append(result)
|
dwpose_woface_list.append(result)
|
||||||
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
|
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
@@ -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
@@ -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])
|
||||||
|
|||||||
@@ -4,15 +4,15 @@ import torch
|
|||||||
try:
|
try:
|
||||||
from spas_sage_attn import block_sparse_sage2_attn_cuda
|
from spas_sage_attn import block_sparse_sage2_attn_cuda
|
||||||
sparse_attn_func = block_sparse_sage2_attn_cuda
|
sparse_attn_func = block_sparse_sage2_attn_cuda
|
||||||
except:
|
except Exception:
|
||||||
try:
|
try:
|
||||||
from sparse_sageattn import sparse_sageattn
|
from sparse_sageattn import sparse_sageattn
|
||||||
sparse_attn_func = sparse_sageattn
|
sparse_attn_func = sparse_sageattn
|
||||||
except:
|
except Exception:
|
||||||
try:
|
try:
|
||||||
from .sparse_sage.core import sparse_sageattn
|
from .sparse_sage.core import sparse_sageattn
|
||||||
sparse_attn_func = sparse_sageattn
|
sparse_attn_func = sparse_sageattn
|
||||||
except:
|
except Exception:
|
||||||
sparse_sageattn = None
|
sparse_sageattn = None
|
||||||
raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.")
|
raise ImportError("sparse_sageattn is not available. Please install the sparse_sageattn package or check your import path.")
|
||||||
|
|
||||||
|
|||||||
@@ -1061,7 +1061,7 @@ class VideoVAE_(nn.Module):
|
|||||||
pbar = ProgressBar(iter_)
|
pbar = ProgressBar(iter_)
|
||||||
try:
|
try:
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
|
for i in tqdm(range(iter_), desc="WanVAE encoding frames", disable=not pbar):
|
||||||
@@ -1092,7 +1092,7 @@ class VideoVAE_(nn.Module):
|
|||||||
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
|
log.info(f"WanVAE encoded input:{input_shape} to {out.shape}")
|
||||||
print_memory(device, process="WanVAE encode")
|
print_memory(device, process="WanVAE encode")
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
return mu
|
return mu
|
||||||
|
|
||||||
@@ -1137,7 +1137,7 @@ class VideoVAE_(nn.Module):
|
|||||||
pbar = ProgressBar(iter_)
|
pbar = ProgressBar(iter_)
|
||||||
try:
|
try:
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
x = self.conv2(z)
|
x = self.conv2(z)
|
||||||
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
|
for i in tqdm(range(iter_), desc="WanVAE decoding frames", disable=not pbar):
|
||||||
@@ -1162,7 +1162,7 @@ class VideoVAE_(nn.Module):
|
|||||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||||
print_memory(device, process="WanVAE decode")
|
print_memory(device, process="WanVAE decode")
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
return out
|
return out
|
||||||
|
|
||||||
@@ -1464,7 +1464,7 @@ class VideoVAE38_(VideoVAE_):
|
|||||||
self.clear_cache()
|
self.clear_cache()
|
||||||
try:
|
try:
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
x = patchify(x, patch_size=2)
|
x = patchify(x, patch_size=2)
|
||||||
t = x.shape[2]
|
t = x.shape[2]
|
||||||
@@ -1492,7 +1492,7 @@ class VideoVAE38_(VideoVAE_):
|
|||||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||||
print_memory(device, process="WanVAE decode")
|
print_memory(device, process="WanVAE decode")
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
return mu
|
return mu
|
||||||
|
|
||||||
@@ -1502,7 +1502,7 @@ class VideoVAE38_(VideoVAE_):
|
|||||||
input_shape = z.shape
|
input_shape = z.shape
|
||||||
try:
|
try:
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||||
|
|
||||||
@@ -1531,7 +1531,7 @@ class VideoVAE38_(VideoVAE_):
|
|||||||
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
log.info(f"WanVAE decoded input:{input_shape} to {out.shape}")
|
||||||
print_memory(device, process="WanVAE decode")
|
print_memory(device, process="WanVAE decode")
|
||||||
torch.cuda.reset_peak_memory_stats(device)
|
torch.cuda.reset_peak_memory_stats(device)
|
||||||
except:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user