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.nn.functional as F
|
||||
from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
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.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"),
|
||||
@@ -32,7 +35,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) -> 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()
|
||||
|
||||
@@ -55,7 +58,20 @@ class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
|
||||
prev_samples = prev_latents["samples"].clone()
|
||||
if overlap != 0:
|
||||
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
|
||||
if ref_latent is not None:
|
||||
@@ -65,7 +81,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 = 2
|
||||
audio_stride = new_audio_embed.get("audio_stride", 2)
|
||||
indices = torch.arange(2 * 2 + 1) - 2
|
||||
|
||||
if frames_processed == 0:
|
||||
@@ -112,9 +128,181 @@ 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
@@ -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:
|
||||
except Exception:
|
||||
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
@@ -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:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
from .utils import log
|
||||
|
||||
@@ -167,7 +167,7 @@ class FantasyTalkingWav2VecEmbeds:
|
||||
|
||||
try:
|
||||
audio_segment = audio_input[start_sample:end_sample]
|
||||
except:
|
||||
except Exception:
|
||||
audio_segment = audio_input
|
||||
|
||||
print("audio_segment.shape", audio_segment.shape)
|
||||
|
||||
+1
-1
@@ -85,7 +85,7 @@ def get_previewer(device, latent_format):
|
||||
taesd = TAEHV(comfy.utils.load_torch_file(taehv_path)).to(device)
|
||||
previewer = TAESDPreviewerImpl(taesd)
|
||||
previewer = WrappedPreviewer(previewer, rate=16)
|
||||
except:
|
||||
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("Using Latent2RGB previewer instead.")
|
||||
method = LatentPreviewMethod.Latent2RGB
|
||||
|
||||
@@ -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:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
pass
|
||||
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):
|
||||
try:
|
||||
import pyloudnorm
|
||||
except:
|
||||
except Exception:
|
||||
raise ImportError("pyloudnorm package is not installed")
|
||||
meter = pyloudnorm.Meter(sr)
|
||||
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 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
|
||||
@@ -328,7 +327,7 @@ class WanVideoTextEncode:
|
||||
try:
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
except:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
encoder = t5["model"]
|
||||
@@ -503,7 +502,7 @@ class WanVideoTextEncodeSingle:
|
||||
log.info(f"Moving video model to {offload_device}")
|
||||
model_to_offload.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
except:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
encoder = t5["model"]
|
||||
@@ -1208,6 +1207,7 @@ 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):
|
||||
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
|
||||
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
|
||||
looping = num_frames > frame_window_size or start_ref_image is not None
|
||||
|
||||
if num_frames < frame_window_size:
|
||||
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.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])
|
||||
|
||||
@@ -1346,6 +1352,7 @@ 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,
|
||||
@@ -2070,76 +2077,6 @@ 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
|
||||
@@ -2394,7 +2331,6 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddTTMLatents": WanVideoAddTTMLatents,
|
||||
"WanVideoAddStoryMemLatents": WanVideoAddStoryMemLatents,
|
||||
"WanVideoSVIProEmbeds": WanVideoSVIProEmbeds,
|
||||
"WanVideoSelfRefineVideo": WanVideoSelfRefineVideo,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
|
||||
+26
-10
@@ -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:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@@ -1512,12 +1512,24 @@ class WanVideoModelLoader:
|
||||
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
|
||||
# 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...")
|
||||
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():
|
||||
@@ -1534,7 +1546,7 @@ class WanVideoModelLoader:
|
||||
class_interval=4,
|
||||
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
|
||||
# SkyreelsV3
|
||||
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:
|
||||
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()
|
||||
# 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})")
|
||||
|
||||
patcher.model["base_dtype"] = base_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:
|
||||
from server import PromptServer
|
||||
except:
|
||||
except Exception:
|
||||
PromptServer = None
|
||||
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
@@ -256,7 +256,7 @@ class CreateCFGScheduleFloatList:
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
)
|
||||
except:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
@@ -319,7 +319,7 @@ class CreateScheduleFloatList:
|
||||
f"{cfg_list}",
|
||||
unique_id
|
||||
)
|
||||
except:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return (cfg_list,)
|
||||
@@ -454,7 +454,7 @@ class NormalizeAudioLoudness:
|
||||
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
|
||||
try:
|
||||
import pyloudnorm
|
||||
except:
|
||||
except Exception:
|
||||
raise ImportError("pyloudnorm package is not installed")
|
||||
meter = pyloudnorm.Meter(sr)
|
||||
loudness = meter.integrated_loudness(audio_array)
|
||||
|
||||
+2
-2
@@ -548,7 +548,7 @@ class WanVideoDiffusionForcingSampler:
|
||||
gc.collect()
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
#region main loop start
|
||||
@@ -615,7 +615,7 @@ class WanVideoDiffusionForcingSampler:
|
||||
try:
|
||||
print_memory(device)
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
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:
|
||||
try:
|
||||
pose_ref = dwpose_model(ref_image.squeeze(0), score_threshold=score_threshold)
|
||||
except:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
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:
|
||||
#except Exception:
|
||||
# result = torch.zeros((height, width, 3), dtype=torch.uint8)
|
||||
dwpose_woface_list.append(result)
|
||||
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
|
||||
|
||||
@@ -12,7 +12,7 @@ from comfy.lora import calculate_weight
|
||||
|
||||
try:
|
||||
from comfy.utils import string_to_seed
|
||||
except:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
continue
|
||||
return model
|
||||
|
||||
@@ -703,8 +703,9 @@ 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))
|
||||
|
||||
@@ -65,16 +65,16 @@ try:
|
||||
# Return tensor with same shape as q
|
||||
return q.clone()
|
||||
sageattn_varlen_func = torch.ops.wanvideo.sageattn_varlen
|
||||
except:
|
||||
except Exception:
|
||||
sageattn_varlen_func = attention_func_error
|
||||
|
||||
# sage3
|
||||
try:
|
||||
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
|
||||
except:
|
||||
except Exception:
|
||||
try:
|
||||
from sageattn import sageattn_blackwell
|
||||
except:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
sageattn_func_ultravico = attention_func_error
|
||||
|
||||
|
||||
|
||||
+13
-13
@@ -10,7 +10,7 @@ from contextlib import nullcontext
|
||||
|
||||
try:
|
||||
from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense, MaskMap
|
||||
except:
|
||||
except Exception:
|
||||
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, 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
|
||||
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, 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"""
|
||||
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_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:
|
||||
# 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_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
|
||||
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)
|
||||
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
|
||||
x.add_(attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2))
|
||||
del k_img, v_img
|
||||
|
||||
# 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=5):
|
||||
ref_frame_index=10, longcat_num_ref_latents=0, num_memory_frames=3, rope_negative_offset=0):
|
||||
|
||||
patch_size = self.patch_size
|
||||
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
||||
|
||||
@@ -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:
|
||||
except Exception:
|
||||
try:
|
||||
from sparse_sageattn import sparse_sageattn
|
||||
sparse_attn_func = sparse_sageattn
|
||||
except:
|
||||
except Exception:
|
||||
try:
|
||||
from .sparse_sage.core import sparse_sageattn
|
||||
sparse_attn_func = sparse_sageattn
|
||||
except:
|
||||
except Exception:
|
||||
sparse_sageattn = None
|
||||
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_)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
pass
|
||||
return mu
|
||||
|
||||
@@ -1137,7 +1137,7 @@ class VideoVAE_(nn.Module):
|
||||
pbar = ProgressBar(iter_)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
pass
|
||||
return out
|
||||
|
||||
@@ -1464,7 +1464,7 @@ class VideoVAE38_(VideoVAE_):
|
||||
self.clear_cache()
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
pass
|
||||
return mu
|
||||
|
||||
@@ -1502,7 +1502,7 @@ class VideoVAE38_(VideoVAE_):
|
||||
input_shape = z.shape
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
except:
|
||||
except Exception:
|
||||
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:
|
||||
except Exception:
|
||||
pass
|
||||
return out
|
||||
|
||||
|
||||
Reference in New Issue
Block a user