From 8a151b5402a2dca818e32e086333dc6e8f818b93 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 24 Aug 2025 11:39:18 +0300 Subject: [PATCH] Multi/InfiniteTalk sampling loop cleanup and optimizations, support FantasyPortrait within the loop --- nodes.py | 191 +++++++++++++------------------- wanvideo/schedulers/__init__.py | 33 +++++- 2 files changed, 109 insertions(+), 115 deletions(-) diff --git a/nodes.py b/nodes.py index d458982..0de48c7 100644 --- a/nodes.py +++ b/nodes.py @@ -1665,7 +1665,7 @@ class WanVideoSampler: #region Scheduler sample_scheduler = None if scheduler != "multitalk": - sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g) log.info(f"sigmas: {sample_scheduler.sigmas}") else: timesteps = torch.tensor([1000, 750, 500, 250], device=device) @@ -1689,27 +1689,6 @@ class WanVideoSampler: steps = len(cfg) else: cfg = [cfg] * (steps + 1) - - if end_step != -1: - timesteps = timesteps[:end_step] - sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1] - log.info(f"Sampling until step {end_step}, timestep: {timesteps[-1]}") - if start_step > 0: - timesteps = timesteps[start_step:] - sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:] - log.info(f"Skipping first {start_step} steps, starting from timestep {timesteps[0]}") - - log.info(f"timesteps: {timesteps}") - - if sample_scheduler is not None: - if hasattr(sample_scheduler, 'timesteps'): - sample_scheduler.timesteps = timesteps - - scheduler_step_args = {"generator": seed_g} - step_sig = inspect.signature(sample_scheduler.step) - for arg in list(scheduler_step_args.keys()): - if arg not in step_sig.parameters: - scheduler_step_args.pop(arg) control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = None vace_data = vace_context = vace_scale = None @@ -2686,7 +2665,7 @@ class WanVideoSampler: # FreeInit noise reinitialization (after first iteration) if freeinit_args is not None and iter_idx > 0: # restart scheduler for each iteration - sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g) # Re-apply start_step and end_step logic to timesteps and sigmas if end_step != -1: @@ -3066,9 +3045,10 @@ class WanVideoSampler: audio_end_idx = audio_start_idx + clip_length indices = (torch.arange(4 + 1) - 2) * 1 current_condframe_index = 0 - - if multitalk_embeds is not None: - total_frames = len(multitalk_audio_embedding[0]) + + audio_embedding = multitalk_audio_embedding + human_num = len(audio_embedding) + audio_embs = None pcd_data = pcd_data_input = None if uni3c_embeds is not None: @@ -3082,13 +3062,10 @@ class WanVideoSampler: "end": uni3c_embeds["end"], } + total_frames = len(audio_embedding[0]) estimated_iterations = total_frames // (frame_num - motion_frame) + 1 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 @@ -3105,29 +3082,13 @@ class WanVideoSampler: audio_embs.append(audio_emb) audio_embs = torch.concat(audio_embs, dim=0).to(dtype) - if uni3c_embeds is not None: - vae.to(device) - # 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]) 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]) + latent_frame_num = (frame_num - 1) // 4 + 1 noise = torch.randn( - 16, (frame_num - 1) // 4 + 1, + 16, latent_frame_num, 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 @@ -3180,52 +3141,27 @@ class WanVideoSampler: thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device) masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > 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 - else: - msk[:, 1:] = 0 - msk = torch.concat([ - torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:] - ], dim=1) - msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w) - msk = msk.transpose(1, 2).to(dtype) # B 4 T H W - - mm.soft_empty_cache() - - # zero padding and vae encode + # zero padding and vae encode for img cond 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, pbar=False).to(dtype) + y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0] + if mode == "multitalk": - latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W + latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W else: cond_ = cond_image if is_first_clip else cond_frame latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0] vae.model.clear_cache() vae.to(offload_device) - y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W + + motion_frame_index = cur_motion_frames_num if mode == "multitalk" else 1 + msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype) + msk[:, :motion_frame_index] = 1 + y = torch.cat([msk, y]) # 4+C T H W mm.soft_empty_cache() else: y = None @@ -3237,33 +3173,8 @@ class WanVideoSampler: timesteps = [torch.tensor([t], device=device) for t in timesteps] timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps] 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: - 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) - transformed_timesteps.append(t_tensor) - - transformed_timesteps.append(torch.tensor([0.], device=device)) - timesteps = transformed_timesteps + sample_scheduler, timesteps, scheduler_step_args = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas, seed_g=seed_g) + timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)] # sample videos latent = noise @@ -3316,6 +3227,43 @@ class WanVideoSampler: else: positive = text_embeds["prompt_embeds"] + 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"] + # }) + + # uni3c slices + if uni3c_embeds is not None: + vae.to(device) + # 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) + vae.model.clear_cache() + vae.to(offload_device) + pcd_data['render_latent'] = render_latent + + # unianimate slices partial_unianim_data = None if unianim_data is not None: partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx] @@ -3328,6 +3276,22 @@ class WanVideoSampler: "end_percent": unianimate_poses["end_percent"] } + # fantasy portrait slices + partial_fantasy_portrait_input = None + if fantasy_portrait_input is not None: + adapter_proj = fantasy_portrait_input["adapter_proj"] + if latent_end_idx > adapter_proj.shape[1]: + pad_len = latent_end_idx - adapter_proj.shape[1] + last_frame = adapter_proj[:, -1:, :, :].repeat(1, pad_len, 1, 1) + padded_proj = torch.cat([adapter_proj, last_frame], dim=1) + else: + padded_proj = adapter_proj + partial_fantasy_portrait_input = fantasy_portrait_input.copy() + partial_fantasy_portrait_input["adapter_proj"] = padded_proj[:, latent_start_idx:latent_end_idx] + + mm.soft_empty_cache() + gc.collect() + # sampling loop 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] @@ -3338,7 +3302,7 @@ class WanVideoSampler: noise_pred, self.cache_state = predict_with_cfg( latent_model_input, cfg[i], positive, text_embeds["negative_prompt_embeds"], timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond, - cache_state=self.cache_state, multitalk_audio_embeds=audio_embs) + cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input) 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) @@ -3373,9 +3337,13 @@ class WanVideoSampler: else: latent[:, :cur_motion_frames_latent_num] = latent_motion_frames - del noise, y, msk, latent_motion_frames + del noise, latent_motion_frames if offload: transformer.to(offload_device) + + mm.soft_empty_cache() + gc.collect() + vae.to(device) videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() vae.model.clear_cache() @@ -3419,7 +3387,6 @@ class WanVideoSampler: cond_image = cond_ del videos, latent - mm.soft_empty_cache() # Repeat audio emb if multitalk_embeds is not None: diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index 5217225..c820865 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -6,7 +6,7 @@ from .flowmatch_pusa import FlowMatchSchedulerPusa from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep from .scheduling_flow_match_lcm import FlowMatchLCMScheduler from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler -import numpy as np +import inspect from ...utils import log scheduler_list = [ @@ -23,7 +23,7 @@ scheduler_list = [ "multitalk" ] -def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None): +def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None, seed_g=None): timesteps = None if 'unipc' in scheduler: sample_scheduler = FlowUniPCMultistepScheduler(shift=shift) @@ -99,4 +99,31 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) if timesteps is None: timesteps = sample_scheduler.timesteps - return sample_scheduler, timesteps \ No newline at end of file + + 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: + 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:] + + log.info(f"timesteps: {timesteps}") + + if hasattr(sample_scheduler, 'timesteps'): + sample_scheduler.timesteps = timesteps + + if seed_g is not None: + scheduler_step_args = {"generator": seed_g} + step_sig = inspect.signature(sample_scheduler.step) + for arg in list(scheduler_step_args.keys()): + if arg not in step_sig.parameters: + scheduler_step_args.pop(arg) + + return sample_scheduler, timesteps, scheduler_step_args \ No newline at end of file