Merge branch 'main' into dev

This commit is contained in:
kijai
2025-08-21 18:42:47 +03:00
9 changed files with 5645 additions and 165 deletions
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -20,8 +20,8 @@ class DownloadAndLoadWav2VecModel:
"required": {
"model": (
[
"facebook/wav2vec2-base-960h",
"TencentGameMate/chinese-wav2vec2-base"
"TencentGameMate/chinese-wav2vec2-base",
"facebook/wav2vec2-base-960h"
],
),
+39 -18
View File
@@ -52,7 +52,8 @@ class MultiTalkModelLoader:
multitalk = {
"proj_model": multitalk_proj_model,
"sd": sd,
"is_gguf": model_path.endswith(".gguf")
"is_gguf": model_path.endswith(".gguf"),
"model_type": "InfiniteTalk" if "infinite" in model.lower() else "MultiTalk",
}
return (multitalk,)
@@ -75,8 +76,8 @@ class MultiTalkWav2VecEmbeds:
return {"required": {
"wav2vec_model": ("WAV2VECMODEL",),
"audio_1": ("AUDIO",),
"normalize_loudness": ("BOOLEAN", {"default": True}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 1}),
"normalize_loudness": ("BOOLEAN", {"default": True, "tooltip": "Normalize the audio loudness to -23 LUFS"}),
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 1, "tooltip": "The total frame count to generate."}),
"fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 60.0, "step": 0.1}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Strength of the audio conditioning"}),
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}),
@@ -90,8 +91,8 @@ class MultiTalkWav2VecEmbeds:
}
}
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", )
RETURN_NAMES = ("multitalk_embeds", "audio", )
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", "INT", )
RETURN_NAMES = ("multitalk_embeds", "audio", "num_frames", )
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
@@ -230,12 +231,30 @@ class MultiTalkWav2VecEmbeds:
offset += w.shape[-1]
out_audio = {"waveform": mixed, "sample_rate": sr}
# Calculate actual frames based on audio duration
actual_num_frames = num_frames
if len(audio_outputs) > 0:
if multi_audio_type == "para":
# For parallel mode, use the longest audio duration
max_audio_duration = max([ao["waveform"].shape[-1] / sr for ao in audio_outputs])
actual_frames_from_audio = int(max_audio_duration * fps)
else: # "add"
# For sequential mode, use the total audio duration
total_audio_duration = sum([ao["waveform"].shape[-1] / sr for ao in audio_outputs])
actual_frames_from_audio = int(total_audio_duration * fps)
# Use the smaller of requested frames or actual audio frames
actual_num_frames = min(num_frames, actual_frames_from_audio)
if actual_frames_from_audio < num_frames:
log.info(f"[MultiTalk] Audio duration ({actual_frames_from_audio} frames) is shorter than requested ({num_frames} frames). Using {actual_num_frames} frames.")
# Debug: log final mixed audio length and mode
total_samples_raw = sum([ao["waveform"].shape[-1] for ao in audio_outputs])
log.info(f"[MultiTalk] total raw duration = {total_samples_raw/sr:.3f}s")
log.info(f"[MultiTalk] multi_audio_type={multi_audio_type} | final waveform shape={out_audio['waveform'].shape} | length={out_audio['waveform'].shape[-1]} samples | seconds={out_audio['waveform'].shape[-1]/sr:.3f}s (expected {'sum' if multi_audio_type=='add' else 'max'} of raw)")
return (multitalk_embeds, out_audio)
return (multitalk_embeds, out_audio, actual_num_frames)
class WanVideoImageToVideoMultiTalk:
@@ -243,11 +262,11 @@ class WanVideoImageToVideoMultiTalk:
def INPUT_TYPES(s):
return {"required": {
"vae": ("WANVAE",),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
"frame_window_size": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
"motion_frame": ("INT", {"default": 25, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation."}),
"force_offload": ("BOOLEAN", {"default": True}),
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the generation"}),
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the generation"}),
"frame_window_size": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "The number of frames to process at once, should be a value the model is generally good at."}),
"motion_frame": ("INT", {"default": 25, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation. Basically the overlap length."}),
"force_offload": ("BOOLEAN", {"default": False, "tooltip": "Whether to force offload the model within the loop for VAE operations, enable if you encounter memory issues."}),
"colormatch": (
[
'disabled',
@@ -258,17 +277,18 @@ class WanVideoImageToVideoMultiTalk:
'hm-mvgd-hm',
'hm-mkl-hm',
], {
"default": 'disabled'
}),
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
},),
},
"optional": {
"start_image": ("IMAGE", {"tooltip": "Image to encode"}),
"start_image": ("IMAGE", {"tooltip": "Images to encode"}),
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
"mode": ([
"auto",
"multitalk",
"infinitetalk"
], {"default": "multitalk", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
], {"default": "auto", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
}
}
@@ -276,6 +296,7 @@ class WanVideoImageToVideoMultiTalk:
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Enables Multi/InfiniteTalk long video generation sampling method, the video is created in windows with overlapping frames. Not compatible or necessary to be used with context windows and many other features besides Multi/InfiniteTalk."
def process(self, vae, width, height, frame_window_size, motion_frame, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None, mode="multitalk"):
@@ -320,7 +341,7 @@ NODE_CLASS_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MultiTalkModelLoader": "MultiTalk Model Loader",
"MultiTalkWav2VecEmbeds": "MultiTalk Wav2Vec Embeds",
"WanVideoImageToVideoMultiTalk": "WanVideo Image To Video MultiTalk"
"MultiTalkModelLoader": "Multi/InfiniteTalk Model Loader",
"MultiTalkWav2VecEmbeds": "Multi/InfiniteTalk Wav2Vec Embeds",
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk"
}
+219 -94
View File
@@ -874,9 +874,9 @@ class WanVideoImageToVideoEncode:
H = height
W = width
lat_h = H // 8
lat_w = W // 8
lat_h = H // vae.upsampling_factor
lat_w = W // vae.upsampling_factor
num_frames = ((num_frames - 1) // 4) * 4 + 1
two_ref_images = start_image is not None and end_image is not None
@@ -1280,8 +1280,6 @@ class WanVideoVACEEncode:
CATEGORY = "WanVideoWrapper"
def process(self, vae, width, height, num_frames, strength, vace_start_percent, vace_end_percent, input_frames=None, ref_images=None, input_masks=None, prev_vace_embeds=None, tiled_vae=False):
vae = vae.to(device)
width = (width // 16) * 16
height = (height // 16) * 16
@@ -1292,8 +1290,8 @@ class WanVideoVACEEncode:
if input_frames is None:
input_frames = torch.zeros((1, 3, num_frames, height, width), device=device, dtype=vae.dtype)
else:
input_frames = input_frames[:num_frames]
input_frames = common_upscale(input_frames.clone().movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
input_frames = input_frames.clone()[:num_frames, :, :, :3]
input_frames = common_upscale(input_frames.movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
input_frames = input_frames.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W
input_frames = input_frames * 2 - 1
if input_masks is None:
@@ -1306,10 +1304,11 @@ class WanVideoVACEEncode:
input_masks = input_masks.unsqueeze(-1).unsqueeze(0).permute(0, 4, 1, 2, 3).repeat(1, 3, 1, 1, 1) # B, C, T, H, W
if ref_images is not None:
ref_images = ref_images.clone()[..., :3]
# Create padded image
if ref_images.shape[0] > 1:
ref_images = torch.cat([ref_images[i] for i in range(ref_images.shape[0])], dim=1).unsqueeze(0)
B, H, W, C = ref_images.shape
current_aspect = W / H
target_aspect = width / height
@@ -1331,12 +1330,12 @@ class WanVideoVACEEncode:
ref_images = ref_images.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3).unsqueeze(0)
ref_images = ref_images * 2 - 1
vae = vae.to(device)
z0 = self.vace_encode_frames(vae, input_frames, ref_images, masks=input_masks, tiled_vae=tiled_vae)
vae.model.clear_cache()
m0 = self.vace_encode_masks(input_masks, ref_images)
z = self.vace_latent(z0, m0)
vae.to(offload_device)
vace_input = {
@@ -1357,6 +1356,7 @@ class WanVideoVACEEncode:
vace_input["additional_vace_inputs"].append(prev_vace_embeds)
return (vace_input,)
def vace_encode_frames(self, vae, frames, ref_images, masks=None, tiled_vae=False):
if ref_images is None:
ref_images = [None] * len(frames)
@@ -1675,6 +1675,10 @@ class WanVideoSampler:
transformer = compile_model(transformer, model["compile_args"])
multitalk_sampling = image_embeds.get("multitalk_sampling", False)
if multitalk_sampling and context_options is not None:
raise Exception("context_options are not compatible or necessary with 'WanVideoImageToVideoMultiTalk' node, since it's already an alternative method that creates the video in a loop.")
if not multitalk_sampling and scheduler == "multitalk":
raise Exception("multitalk scheduler is only for multitalk sampling when using ImagetoVideoMultiTalk -node")
@@ -1803,7 +1807,7 @@ class WanVideoSampler:
control_embeds = image_embeds.get("control_embeds", None)
if control_embeds is not None:
if transformer.in_dim not in [52, 48, 36, 32]:
if transformer.in_dim not in [148, 52, 48, 36, 32]:
raise ValueError("Control signal only works with Fun-Control model")
control_latents = control_embeds.get("control_images", None)
@@ -1900,10 +1904,10 @@ class WanVideoSampler:
patcher = apply_lora(patcher, device, device, low_mem_load=False, control_lora=True)
patcher.model.is_patched = True
else:
if transformer.in_dim not in [48, 36, 32, 52]:
if transformer.in_dim not in [148, 48, 36, 32, 52]:
raise ValueError("Control signal only works with Fun-Control model")
image_cond = torch.zeros_like(noise).to(device) #fun control
if transformer.in_dim == 52 or transformer.control_adapter is not None: #fun 2.2 control
if transformer.in_dim in [148, 52] or transformer.control_adapter is not None: #fun 2.2 control
mask_latents = torch.tile(
torch.zeros_like(noise[:1]), [4, 1, 1, 1]
)
@@ -1917,7 +1921,7 @@ class WanVideoSampler:
control_start_percent = control_embeds.get("start_percent", 0.0)
control_end_percent = control_embeds.get("end_percent", 1.0)
else:
if transformer.in_dim == 36: #fun inp
if transformer.in_dim in [148, 52]: #fun inp
mask_latents = torch.tile(
torch.zeros_like(noise[:1]), [4, 1, 1, 1]
)
@@ -2113,7 +2117,7 @@ class WanVideoSampler:
mtv_freqs = mtv_freqs.to(device, dtype)
# vid2vid
if samples is not None:
if samples is not None and not multitalk_sampling:
saved_generator_state = samples.get("generator_state", None)
if saved_generator_state is not None:
seed_g.set_state(saved_generator_state)
@@ -2127,29 +2131,29 @@ class WanVideoSampler:
else:
noise = input_samples
mask = samples.get("noise_mask", None)
if mask is not None:
log.info(f"Latent mask shape: {mask.shape}")
noise_mask = samples.get("noise_mask", None)
if noise_mask is not None:
log.info(f"Latent noise_mask shape: {noise_mask.shape}")
original_image = input_samples.to(device)
if len(mask.shape) == 4:
mask = mask.squeeze(1)
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
mask = torch.nn.functional.interpolate(
mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
noise_mask = torch.nn.functional.interpolate(
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
mode='trilinear',
align_corners=False
).squeeze(0) # Remove batch dim, keep channel dim
# Add batch & channel dims for final output
mask = mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1)
noise_mask = noise_mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1)
if mask.shape[2] != noise.shape[1]:
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
if noise_mask.shape[2] != noise.shape[1]:
noise_mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - noise_mask.shape[2], noise.shape[2], noise.shape[3]), noise_mask], dim=2)
# extra latents (Pusa) and 5b
latents_to_insert = add_index = None
if (extra_latents := image_embeds.get("extra_latents", None)) is not None:
if (extra_latents := image_embeds.get("extra_latents", None)) is not None and transformer.multitalk_model_type.lower() != "infinitetalk":
all_indices = []
for entry in extra_latents:
add_index = entry["index"]
@@ -2178,7 +2182,7 @@ class WanVideoSampler:
if uni3c_embeds is not None:
transformer.controlnet = uni3c_embeds["controlnet"]
pcd_data = {
"render_latent": uni3c_embeds["render_latent"].to(dtype),
"render_latent": uni3c_embeds["render_latent"],
"render_mask": uni3c_embeds["render_mask"],
"camera_embedding": uni3c_embeds["camera_embedding"],
"controlnet_weight": uni3c_embeds["controlnet_weight"],
@@ -2701,18 +2705,19 @@ class WanVideoSampler:
from .latent_preview import prepare_callback #custom for tiny VAE previews
callback = prepare_callback(patcher, len(timesteps))
log.info(f"Input sequence length: {seq_len}")
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
if not multitalk_sampling:
log.info(f"Input sequence length: {seq_len}")
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
intermediate_device = device
# diff diff prep
masks = None
if samples is not None and mask is not None:
mask = 1 - mask
if not multitalk_sampling and samples is not None and noise_mask is not None:
noise_mask = 1 - noise_mask
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
masks = mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = masks > thresholds
latent_shift_loop = False
@@ -2789,27 +2794,24 @@ class WanVideoSampler:
try:
pbar = ProgressBar(len(timesteps))
#region main loop start
for idx, t in enumerate(tqdm(timesteps)):
for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling)):
if flowedit_args is not None:
if idx < skip_steps:
continue
# diff diff
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image, torch.tensor([noise_timestep]), noise.to(device)
)
mask = masks[idx]
mask = mask.to(latent)
latent = image_latent * mask + latent * (1-mask)
# end diff diff
if bidirectional_sampling:
latent_flipped = torch.flip(latent, dims=[1])
latent_model_input_flipped = latent_flipped.to(device)
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
latent_model_input = latent.to(device)
current_step_percentage = idx / len(timesteps)
@@ -3097,8 +3099,10 @@ class WanVideoSampler:
#region multitalk
elif multitalk_sampling:
mode = image_embeds.get("multitalk_mode", "multitalk")
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
log.info(f"Multitalk mode: {mode}")
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
cond_frame = None
offload = image_embeds.get("force_offload", False)
tiled_vae = image_embeds.get("tiled_vae", False)
frame_num = clip_length = image_embeds.get("num_frames", 81)
@@ -3110,23 +3114,20 @@ class WanVideoSampler:
motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None)
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
face_scale = 0.1
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
background_mask = torch.zeros([target_h, target_w])
background_mask = torch.zeros([target_h, target_w])
human_mask1 = torch.zeros([target_h, target_w])
human_mask2 = torch.zeros([target_h, target_w])
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
human_mask2[x_min:x_max, righty_min:righty_max] = 1
background_mask += human_mask1
background_mask += human_mask2
human_masks = [human_mask1, human_mask2]
background_mask = torch.where(background_mask > 0, torch.tensor(0), torch.tensor(1))
human_masks.append(background_mask)
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
human_masks = [human_mask1, human_mask2, background_mask]
ref_target_masks = torch.stack(human_masks, dim=0)
multitalk_embeds['ref_target_masks'] = ref_target_masks
@@ -3134,7 +3135,7 @@ class WanVideoSampler:
is_first_clip = True
arrive_last_frame = False
cur_motion_frames_num = 1
audio_start_idx = iteration_count = 0
audio_start_idx = iteration_count = step_iteration_count= 0
audio_end_idx = audio_start_idx + clip_length
indices = (torch.arange(4 + 1) - 2) * 1
current_condframe_index = 0
@@ -3146,7 +3147,7 @@ class WanVideoSampler:
if uni3c_embeds is not None:
transformer.controlnet = uni3c_embeds["controlnet"]
pcd_data = {
"render_latent": uni3c_embeds["render_latent"].to(dtype),
"render_latent": uni3c_embeds["render_latent"],
"render_mask": uni3c_embeds["render_mask"],
"camera_embedding": uni3c_embeds["camera_embedding"],
"controlnet_weight": uni3c_embeds["controlnet_weight"],
@@ -3155,19 +3156,19 @@ class WanVideoSampler:
}
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
log.info(f"Total frames: {total_frames}, frame_num: {frame_num}, motion_frame: {motion_frame}")
log.info(f"Estimated iterations: {estimated_iterations}")
loop_pbar = tqdm(total=estimated_iterations, desc="Generating video clips")
loop_pbar = tqdm(total=estimated_iterations, desc="Total progress", position=1, leave=True)
callback = prepare_callback(patcher, estimated_iterations)
audio_embedding = multitalk_audio_embedding
human_num = len(audio_embedding)
audio_embs = None
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
while True: # start video generation iteratively
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk":
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1]
log.info(f"current_condframe_index: {current_condframe_index}")
log.info(f"audio_start_idx: {audio_start_idx}")
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
if multitalk_embeds is not None:
audio_embs = []
# split audio with window size
@@ -3180,18 +3181,83 @@ class WanVideoSampler:
if uni3c_embeds is not None:
vae.to(device)
render_latent = vae.encode(original_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
# Pad original_images if needed
num_frames = original_images.shape[2]
required_frames = audio_end_idx - audio_start_idx
if audio_end_idx > num_frames:
pad_len = audio_end_idx - num_frames
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
padded_images = torch.cat([original_images, last_frame], dim=2)
else:
padded_images = original_images
render_latent = vae.encode(
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
device=device, tiled=tiled_vae
).to(dtype)
pcd_data['render_latent'] = render_latent
h, w = cond_image.shape[-2], cond_image.shape[-1]
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
seq_len = ((frame_num - 1) // VAE_STRIDE[0] + 1) * lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
noise = torch.randn(
16, (frame_num - 1) // 4 + 1,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
if samples is not None:
input_samples = samples["samples"].squeeze(0).to(noise)
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
pad_length = latent_end_idx - input_samples.shape[1]
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
input_samples = torch.cat([input_samples, last_frame], dim=1)
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
# get mask
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
if add_noise_to_samples:
latent_timestep = timesteps[0]
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
else:
noise = input_samples
# diff diff prep
masks = None
if noise_mask is not None:
noise_mask = 1 - noise_mask
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = masks > thresholds
window_vace_data = None
if vace_data is not None:
window_vace_data = []
for vace_entry in vace_data:
partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
if has_ref:
partial_context[:, 0] = vace_entry["context"][0][:, 0]
window_vace_data.append({
"context": [partial_context],
"scale": vace_entry["scale"],
"start": vace_entry["start"],
"end": vace_entry["end"],
"seq_len": vace_entry["seq_len"]
})
# get image cond mask
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
if mode == "multitalk":
msk[:, cur_motion_frames_num:] = 0
@@ -3206,26 +3272,28 @@ class WanVideoSampler:
mm.soft_empty_cache()
# zero padding and vae encode
video_frames = torch.zeros(1, cond_image.shape[1], frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
if cond_image is not None:
video_frames = torch.zeros(1, 3, frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae).to(dtype)
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
# encode
vae.to(device)
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)
if mode == "multitalk":
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
else:
if is_first_clip:
latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
if mode == "multitalk":
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
else:
latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
latent_motion_frames = latent_motion_frames[0]
vae.to(offload_device)
y = torch.concat([msk, y], dim=1) # B 4+C T H W
mm.soft_empty_cache()
if is_first_clip:
latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
else:
latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
latent_motion_frames = latent_motion_frames[0]
vae.to(offload_device)
y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W
mm.soft_empty_cache()
else:
y = None
latent_motion_frames = noise[:, :1]
if scheduler == "multitalk":
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
@@ -3235,6 +3303,24 @@ class WanVideoSampler:
else:
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
steps = len(timesteps)
if end_step != -1 and start_step >= end_step:
raise ValueError("start_step must be less than end_step")
if denoise_strength < 1.0:
if start_step != 0:
raise ValueError("start_step must be 0 when denoise_strength is used")
start_step = steps - int(steps * denoise_strength) - 1
if (end_step != -1 or end_step >= steps):
timesteps = timesteps[:end_step]
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
if start_step > 0:
timesteps = timesteps[start_step:]
sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:]
if sample_scheduler is not None:
if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps
transformed_timesteps = []
for t in timesteps:
t_tensor = torch.tensor([t.item()], device=device)
@@ -3287,8 +3373,16 @@ class WanVideoSampler:
else:
transformer.to(device)
comfy_pbar = ProgressBar(len(timesteps)-1)
for i in tqdm(range(len(timesteps)-1)):
# Use the appropriate prompt for this section
if len(text_embeds["prompt_embeds"]) > 1:
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
positive = [text_embeds["prompt_embeds"][prompt_index]]
log.info(f"Using prompt index: {prompt_index}")
else:
positive = text_embeds["prompt_embeds"]
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
for i in range(len(timesteps)-1):
timestep = timesteps[i]
latent_model_input = latent.to(device)
if mode == "infinitetalk":
@@ -3297,14 +3391,18 @@ class WanVideoSampler:
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input,
cfg[idx],
text_embeds["prompt_embeds"],
positive,
text_embeds["negative_prompt_embeds"],
timestep, idx, y.squeeze(0), clip_embeds, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
timestep, idx, y, clip_embeds, control_latents, window_vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs)
sampling_pbar.update(1)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
callback(iteration_count, callback_latent, None, estimated_iterations)
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
step_iteration_count += 1
# update latent
if scheduler == "multitalk":
@@ -3314,13 +3412,22 @@ class WanVideoSampler:
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = latent.to(intermediate_device)
temp_x0 = sample_scheduler.step(
noise_pred.unsqueeze(0),
timestep,
latent.unsqueeze(0),
**scheduler_step_args)[0]
latent = temp_x0.squeeze(0)
# differential diffusion inpaint
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image, torch.tensor([noise_timestep]), noise.to(device)
)
mask = masks[idx].to(latent)
latent = image_latent * mask + latent * (1-mask)
# injecting motion frames
if not is_first_clip and mode == "multitalk":
@@ -3334,13 +3441,14 @@ class WanVideoSampler:
x0 = latent.to(device)
del latent_model_input, timestep
comfy_pbar.update(1)
if offload:
offload_transformer(transformer)
vae.to(device)
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae)
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae, pbar=False)
vae.to(offload_device)
sampling_pbar.close()
# cache generated samples
videos = torch.stack(videos).cpu() # B C T H W
@@ -3411,9 +3519,7 @@ class WanVideoSampler:
del noise, latent
if force_offload:
if not model["auto_cpu_offload"]:
transformer.to(offload_device)
mm.soft_empty_cache()
gc.collect()
offload_transformer(transformer)
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
@@ -3486,6 +3592,25 @@ class WanVideoSampler:
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
# differential diffusion inpaint
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image, torch.tensor([noise_timestep]), noise.to(device)
)
mask = masks[idx].to(latent)
latent = image_latent * mask + latent * (1-mask)
if freeinit_args is not None:
current_latent = latent.clone()
+7 -11
View File
@@ -503,7 +503,7 @@ class WanVideoVACEModelSelect:
def INPUT_TYPES(s):
return {
"required": {
"vace_model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' VACE model to use when not using model that has it included"}),
"vace_model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' VACE model to use when not using model that has it included"}),
},
}
@@ -1021,15 +1021,9 @@ class WanVideoModelLoader:
if extra_model is not None:
if gguf:
if not extra_model["path"].endswith(".gguf"):
raise ValueError("With GGUF main model the extra model must also be a GGUF quantized, if the main model already has extra included, you can disconnect the extra model loader")
from diffusers.models.model_loading_utils import load_gguf_checkpoint
extra_sd = load_gguf_checkpoint(extra_model["path"])
new_keys = {}
for k, v in extra_sd.items():
if "vace" in k:
new_keys[k] = v
extra_sd = new_keys
if not vace_model["path"].endswith(".gguf"):
raise ValueError("With GGUF main model the VACE module must also be a GGUF quantized, if the main model already has VACE included, you can disconnect the VACE module loader")
vace_sd = load_gguf_checkpoint(vace_model["path"])
else:
extra_sd = load_torch_file(extra_model["path"], device=transformer_load_device, safe_load=True)
sd.update(extra_sd)
@@ -1209,6 +1203,7 @@ class WanVideoModelLoader:
if multitalk_model is not None:
if multitalk_model["is_gguf"] and not gguf:
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
# init audio module
from .multitalk.multitalk import SingleStreamMultiAttention
from .wanvideo.modules.model import WanRMSNorm, WanLayerNorm
@@ -1228,8 +1223,9 @@ class WanVideoModelLoader:
attention_mode=attention_mode,
)
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity()
log.info("MultiTalk model detected, patching model...")
log.info(f"{multitalk_model_type} detected, patching model...")
transformer.audio_proj = multitalk_model["proj_model"]
transformer.multitalk_model_type = multitalk_model_type
sd.update(multitalk_model["sd"])
# Additional cond latents
+10 -10
View File
@@ -21,7 +21,7 @@ class WanVideoUni3C_ControlnetLoader:
"model": (folder_paths.get_filename_list("controlnet"), {"tooltip": "These models are loaded from the 'ComfyUI/models/controlnet' -folder",}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn'], {"default": 'disabled', "tooltip": "optional quantization method"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
"attention_mode": ([
"sdpa",
@@ -138,12 +138,12 @@ class WanVideoUni3C_embeds:
def INPUT_TYPES(s):
return {"required": {
"controlnet": ("WANVIDEOCONTROLNET",),
"render_latent": ("LATENT",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply the controlnet"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply the controlnet"}),
},
"optional": {
"render_latent": ("LATENT",),
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
},
}
@@ -153,16 +153,16 @@ class WanVideoUni3C_embeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, controlnet, render_latent, strength, start_percent, end_percent, render_mask=None):
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None):
device = mm.get_torch_device()
latent_mask = None
latents = render_latent["samples"]
nframe = latents.shape[2] * 4
height = latents.shape[3] * 8
width = latents.shape[4] * 8
latent_mask = latents = None
if render_latent is not None:
latents = render_latent["samples"]
# nframe = latents.shape[2] * 4
# height = latents.shape[3] * 8
# width = latents.shape[4] * 8
if render_mask is not None:
raise NotImplementedError("render_mask is not implemented at this time")
@@ -224,7 +224,7 @@ class WanVideoUni3C_embeds:
"controlnet_weight": strength,
"start": start_percent,
"end": end_percent,
"render_latent": latents.to(device),
"render_latent": latents,
"render_mask": latent_mask,
"camera_embedding": None
}
+3 -1
View File
@@ -1305,6 +1305,8 @@ class WanModel(torch.nn.Module):
self.video_attention_split_steps = []
self.lora_scheduling_enabled = False
self.multitalk_model_type = "none"
# embeddings
self.patch_embedding = nn.Conv3d(
in_dim, dim, kernel_size=patch_size, stride=patch_size)
@@ -1666,7 +1668,7 @@ class WanModel(torch.nn.Module):
#uni3c controlnet
if pcd_data is not None:
hidden_states = x[0].unsqueeze(0).clone().float()
render_latent = torch.cat([hidden_states[:, :20], pcd_data["render_latent"]], dim=1)
render_latent = torch.cat([hidden_states[:, :20], pcd_data["render_latent"].to(x[0].dtype)], dim=1)
# embeddings
if control_lora_enabled:
+45 -29
View File
@@ -1019,12 +1019,13 @@ class VideoVAE_(nn.Module):
return mu
def encode(self, x):
def encode(self, x, pbar=True):
self.clear_cache()
## cache
pbar = ProgressBar(x.shape[2])
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
if pbar:
pbar = ProgressBar(iter_)
for i in range(iter_):
self._enc_conv_idx = [0]
@@ -1037,10 +1038,13 @@ class VideoVAE_(nn.Module):
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
out = torch.cat([out, out_], 2)
pbar.update(iter_)
if pbar:
pbar.update(1)
mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
if pbar:
pbar.update_absolute(0)
return mu
@@ -1076,14 +1080,13 @@ class VideoVAE_(nn.Module):
def decode(self, z):
def decode(self, z, pbar=True):
self.clear_cache()
# z: [b,c,t,h,w]
pbar = ProgressBar(z.shape[2])
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
if pbar:
pbar = ProgressBar(iter_)
x = self.conv2(z)
for i in range(iter_):
self._conv_idx = [0]
@@ -1096,7 +1099,11 @@ class VideoVAE_(nn.Module):
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2) # may add tensor offload
pbar.update(1)
if pbar:
pbar.update(1)
if pbar:
pbar.update_absolute(0)
return out
def reparameterize(self, mu, log_var):
@@ -1167,7 +1174,7 @@ class WanVideoVAE(nn.Module):
return mask
def tiled_decode(self, hidden_states, device, tile_size, tile_stride):
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True):
_, _, T, H, W = hidden_states.shape
size_h, size_w = tile_size
stride_h, stride_w = tile_stride
@@ -1187,8 +1194,8 @@ class WanVideoVAE(nn.Module):
out_T = T * 4 - 3
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
pbar = ProgressBar(len(tasks))
if pbar:
pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device)
@@ -1215,13 +1222,14 @@ class WanVideoVAE(nn.Module):
target_h: target_h + hidden_states_batch.shape[3],
target_w: target_w + hidden_states_batch.shape[4],
] += mask
pbar.update(1)
if pbar:
pbar.update(1)
values = values / weight
values = values.float().clamp_(-1, 1)
return values
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False):
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False, pbar=True):
_, _, T, H, W = video.shape
if tile_size is None and tile_stride is None:
@@ -1248,8 +1256,8 @@ class WanVideoVAE(nn.Module):
out_T += 1
weight = torch.zeros((1, 1, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
values = torch.zeros((1, self.z_dim, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
pbar = ProgressBar(len(tasks))
if pbar:
pbar = ProgressBar(len(tasks))
for h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"):
hidden_states_batch = video[:, :, :, h:h_, w:w_].to(computation_device)
if end_:
@@ -1279,21 +1287,22 @@ class WanVideoVAE(nn.Module):
target_h: target_h + hidden_states_batch.shape[3],
target_w: target_w + hidden_states_batch.shape[4],
] += mask
pbar.update(1)
if pbar:
pbar.update(1)
values = values / weight
values = values.float()
return values
def single_encode(self, video, device):
def single_encode(self, video, device, pbar=True):
video = video.to(device)
x = self.model.encode(video)
x = self.model.encode(video, pbar=pbar)
return x.float()
def single_decode(self, hidden_state, device):
def single_decode(self, hidden_state, device, pbar=True):
hidden_state = hidden_state.to(device)
video = self.model.decode(hidden_state)
video = self.model.decode(hidden_state, pbar=pbar)
return video
def double_encode(self, video, device):
@@ -1308,36 +1317,36 @@ class WanVideoVAE(nn.Module):
video = self.model.decode_2(hidden_state)
return video
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None):
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None, pbar=True):
videos = [video.to("cpu") for video in videos]
hidden_states = []
for video in videos:
video = video.unsqueeze(0)
if tiled:
hidden_state = self.tiled_encode(video, device, tile_size, tile_stride, end_=end_)
hidden_state = self.tiled_encode(video, device, tile_size, tile_stride, end_=end_, pbar=pbar)
else:
if end_:
hidden_state = self.double_encode(video, device)
else:
hidden_state = self.single_encode(video, device)
hidden_state = self.single_encode(video, device, pbar=pbar)
hidden_state = hidden_state.squeeze(0)
hidden_states.append(hidden_state)
hidden_states = torch.stack(hidden_states)
return hidden_states
def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16)):
def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16), pbar=True):
hidden_states = [hidden_state.to("cpu") for hidden_state in hidden_states]
videos = []
for hidden_state in hidden_states:
hidden_state = hidden_state.unsqueeze(0)
if tiled:
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride)
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar)
else:
if end_:
video = self.double_decode(hidden_state, device)
else:
video = self.single_decode(hidden_state, device)
video = self.single_decode(hidden_state, device, pbar=pbar)
video = video.squeeze(0)
videos.append(video)
return videos
@@ -1397,11 +1406,13 @@ class VideoVAE38_(VideoVAE_):
attn_scales, self.temperal_upsample, dropout)
def encode(self, x):
def encode(self, x, pbar=True):
self.clear_cache()
x = patchify(x, patch_size=2)
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
if pbar:
pbar = ProgressBar(iter_)
for i in range(iter_):
self._enc_conv_idx = [0]
if i == 0:
@@ -1413,6 +1424,8 @@ class VideoVAE38_(VideoVAE_):
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
out = torch.cat([out, out_], 2)
if pbar:
pbar.update(1)
mu = self.conv1(out).chunk(2, dim=1)[0]
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
@@ -1421,12 +1434,13 @@ class VideoVAE38_(VideoVAE_):
return mu
def decode(self, z):
def decode(self, z, pbar=True):
self.clear_cache()
z = z / self.inv_std.to(z) + self.mean.to(z)
iter_ = z.shape[2]
if pbar:
pbar = ProgressBar(iter_)
x = self.conv2(z)
for i in range(iter_):
self._conv_idx = [0]
@@ -1440,6 +1454,8 @@ class VideoVAE38_(VideoVAE_):
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2)
if pbar:
pbar.update(1)
out = unpatchify(out, patch_size=2)
self.clear_cache()
return out