From d20baf00247fd06553fdc9253e18732244e54172 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 26 Mar 2025 23:12:15 +0200 Subject: [PATCH] Add euler/beta, clean scheduler selection code --- nodes.py | 38 ++++++++++++++------------------------ 1 file changed, 14 insertions(+), 24 deletions(-) diff --git a/nodes.py b/nodes.py index a0e790c..d7e510c 100644 --- a/nodes.py +++ b/nodes.py @@ -1670,7 +1670,7 @@ class WanVideoSampler: "shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}), - "scheduler": (["unipc", "dpm++", "dpm++_sde", "euler"], + "scheduler": (["unipc", "dpm++", "dpm++_sde", "euler", "euler/beta"], { "default": 'unipc' }), @@ -1713,36 +1713,26 @@ class WanVideoSampler: steps = int(steps/denoise_strength) + scheduler_args = { + "num_train_timesteps": 1000, + "shift": shift, + "use_dynamic_shifting": False, + } + if scheduler == 'unipc': - sample_scheduler = FlowUniPCMultistepScheduler( - num_train_timesteps=1000, - shift=shift, - use_dynamic_shifting=False) - sample_scheduler.set_timesteps( - steps, device=device, shift=shift) - timesteps = sample_scheduler.timesteps - elif scheduler == 'euler': - sample_scheduler = FlowMatchEulerDiscreteScheduler( - num_train_timesteps=1000, - shift=shift, - use_dynamic_shifting=False) - sample_scheduler.set_timesteps(steps, device=device, mu=1) - timesteps = sample_scheduler.timesteps - + sample_scheduler = FlowUniPCMultistepScheduler(**scheduler_args) + sample_scheduler.set_timesteps(steps, device=device, shift=shift) + elif scheduler in ['euler/beta', 'euler']: + sample_scheduler = FlowMatchEulerDiscreteScheduler(**scheduler_args, use_beta_sigmas=(scheduler == 'euler/beta')) + sample_scheduler.set_timesteps(steps, device=device, mu=1) elif 'dpm++' in scheduler: if scheduler == 'dpm++_sde': algorithm_type = "sde-dpmsolver++" else: algorithm_type = "dpmsolver++" - sample_scheduler = FlowDPMSolverMultistepScheduler( - num_train_timesteps=1000, - shift=shift, - use_dynamic_shifting=False, - algorithm_type= algorithm_type) + sample_scheduler = FlowDPMSolverMultistepScheduler(**scheduler_args, algorithm_type= algorithm_type) sample_scheduler.set_timesteps(steps, device=device, mu=1) - timesteps = sample_scheduler.timesteps - else: - raise NotImplementedError("Unsupported solver.") + timesteps = sample_scheduler.timesteps if denoise_strength < 1.0: steps = int(steps * denoise_strength)