Update WanVideoScheduler -node
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
return sample_scheduler, timesteps, start_idx, end_idx
|
||||
Reference in New Issue
Block a user