From 56bd2ec5f6ad9be0e2e9814968be7ab0ad42e6bb Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 13:19:07 +0300 Subject: [PATCH] Update WanVideoScheduler -node --- nodes.py | 16 ++++++++++------ wanvideo/schedulers/__init__.py | 7 ++++--- 2 files changed, 14 insertions(+), 9 deletions(-) diff --git a/nodes.py b/nodes.py index 20812ca..765de2d 100644 --- a/nodes.py +++ b/nodes.py @@ -1608,7 +1608,7 @@ class WanVideoScheduler: #WIP EXPERIMENTAL = True def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None): - sample_scheduler, timesteps = get_scheduler( + sample_scheduler, timesteps, start_idx, end_idx = get_scheduler( scheduler, steps, start_step, end_step, shift, @@ -1642,8 +1642,12 @@ class WanVideoScheduler: #WIP ax.tick_params(axis='x', colors='white') # X tick color ax.tick_params(axis='y', colors='white') # Y tick color # Add split point if end_step is defined - if end_step != -1 and 0 <= end_step < len(sigmas_np): - ax.axvline(end_step, color='red', linestyle='--', linewidth=2, label='end_step split') + if end_idx != -1 and 0 <= end_idx < len(sigmas_np): + ax.axvline(end_idx, color='red', linestyle='--', linewidth=2, label='end_step split') + # Add split point if start_step is defined + if start_idx > 0 and 0 <= start_idx < len(sigmas_np): + ax.axvline(start_idx, color='green', linestyle='--', linewidth=2, label='start_step split') + if (end_idx != -1 and 0 <= end_idx < len(sigmas_np)) or (start_idx > 0 and 0 <= start_idx < len(sigmas_np)): ax.legend() plt.tight_layout() plt.savefig(buf, format='png') @@ -1815,7 +1819,7 @@ class WanVideoSampler: sample_scheduler = scheduler["sample_scheduler"] timesteps = scheduler["timesteps"] elif scheduler != "multitalk": - sample_scheduler, timesteps = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) log.info(f"sigmas: {sample_scheduler.sigmas}") else: timesteps = torch.tensor([1000, 750, 500, 250], device=device) @@ -2868,7 +2872,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, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) # Re-apply start_step and end_step logic to timesteps and sigmas if end_step != -1: @@ -3385,7 +3389,7 @@ 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, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)] # sample videos diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index 407e107..d81aa98 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -6,7 +6,6 @@ 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 inspect from ...utils import log scheduler_list = [ @@ -112,6 +111,8 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo start_idx = 0 end_idx = len(timesteps) - 1 + log.info(f"Total timesteps: {timesteps}") + if isinstance(start_step, float): idxs = (sample_scheduler.sigmas <= start_step).nonzero(as_tuple=True)[0] if len(idxs) > 0: @@ -134,9 +135,9 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer - log.info(f"timesteps: {timesteps}") + log.info(f"Using timesteps: {timesteps}") if hasattr(sample_scheduler, 'timesteps'): sample_scheduler.timesteps = timesteps - return sample_scheduler, timesteps \ No newline at end of file + return sample_scheduler, timesteps, start_idx, end_idx \ No newline at end of file