From 053f26ce828eabd0c7a67fde6138742db5442cd9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 20 Aug 2025 16:40:13 +0300 Subject: [PATCH] Multi/InfiniteTalk v2v --- nodes.py | 72 ++++++++++++++++++++++++++++++++------- nodes_model_loading.py | 1 + wanvideo/wan_video_vae.py | 61 +++++++++++++++++++-------------- 3 files changed, 96 insertions(+), 38 deletions(-) diff --git a/nodes.py b/nodes.py index 0eb3773..ec5b869 100644 --- a/nodes.py +++ b/nodes.py @@ -2033,7 +2033,7 @@ class WanVideoSampler: context = get_context_scheduler(context_schedule) # 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) @@ -2623,7 +2623,7 @@ class WanVideoSampler: # diff diff prep masks = None - if samples is not None and mask is not None: + if not multitalk_sampling and samples is not None and mask is not None: mask = 1 - mask thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps) thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device) @@ -2704,7 +2704,7 @@ 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 @@ -3071,17 +3071,16 @@ class WanVideoSampler: } estimated_iterations = total_frames // (frame_num - motion_frame) + 1 - 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 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_indices: {audio_start_idx}-{audio_end_idx}|{clip_length}") if multitalk_embeds is not None: audio_embs = [] # split audio with window size @@ -3094,7 +3093,6 @@ class WanVideoSampler: if uni3c_embeds is not None: vae.to(device) - print("original_images", original_images.shape) # Pad original_images if needed num_frames = original_images.shape[2] required_frames = audio_end_idx - audio_start_idx @@ -3117,6 +3115,34 @@ class WanVideoSampler: 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) + + if samples is not None: + input_samples = samples["samples"].squeeze(0).to(noise) + # Calculate the correct 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] + + # 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] + + 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 # get mask msk = torch.ones(1, frame_num, lat_h, lat_w, device=device) @@ -3138,8 +3164,7 @@ class WanVideoSampler: # 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) + 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 @@ -3162,6 +3187,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) @@ -3214,8 +3257,8 @@ class WanVideoSampler: elif model["manual_offloading"]: transformer.to(device) - comfy_pbar = ProgressBar(len(timesteps)-1) - for i in tqdm(range(len(timesteps)-1)): + 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": @@ -3229,6 +3272,8 @@ class WanVideoSampler: timestep, idx, y.squeeze(0), clip_embeds, control_latents, 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) @@ -3261,13 +3306,14 @@ class WanVideoSampler: x0 = latent.to(device) del latent_model_input, timestep - comfy_pbar.update(1) if offload: transformer.to(offload_device) 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 diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 117151b..1c9e16b 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1098,6 +1098,7 @@ class WanVideoModelLoader: #for name, param in transformer.named_parameters(): # print(name, param.dtype, param.device, param.shape) pbar.update_absolute(param_count) + pbar.update_absolute(0) comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device diff --git a/wanvideo/wan_video_vae.py b/wanvideo/wan_video_vae.py index 7ac9a05..472dad7 100644 --- a/wanvideo/wan_video_vae.py +++ b/wanvideo/wan_video_vae.py @@ -1019,10 +1019,11 @@ 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]) + if pbar: + pbar = ProgressBar(x.shape[2]) t = x.shape[2] iter_ = 1 + (t - 1) // 4 @@ -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(iter_) 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,11 +1080,12 @@ 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]) - + if pbar: + pbar = ProgressBar(z.shape[2]) + z = z / self.inv_std.to(z) + self.mean.to(z) iter_ = z.shape[2] @@ -1096,7 +1101,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 +1176,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 +1196,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 +1224,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 +1258,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 +1289,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 +1319,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