diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index dcb1a0e..6ee36b2 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -7,6 +7,7 @@ from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep from .scheduling_flow_match_lcm import FlowMatchLCMScheduler from .fm_sa_ode import FlowMatchSAODEStableScheduler from .fm_rcm import rCMFlowMatchScheduler +from .vitb_unipc import ViBTScheduler from ...utils import log try: @@ -29,12 +30,17 @@ scheduler_list = [ "flowmatch_pusa", "multitalk", "sa_ode_stable", - "rcm" + "rcm", + "vitb_unipc", ] def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False, **kwargs): timesteps = None - if 'unipc' in scheduler: + if scheduler == 'vitb_unipc': + sample_scheduler = ViBTScheduler() + sample_scheduler.set_parameters(shift=shift) + sample_scheduler.set_timesteps(steps, device=device) + elif 'unipc' in scheduler: sample_scheduler = FlowUniPCMultistepScheduler(shift=shift) if sigmas is None: sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler)) @@ -42,7 +48,6 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo sample_scheduler.sigmas = sigmas.to(device) sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device) sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps) - elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']: if 'longcat' in scheduler: num_distill_sample_steps = 50 diff --git a/wanvideo/schedulers/vitb_unipc.py b/wanvideo/schedulers/vitb_unipc.py new file mode 100644 index 0000000..60a2e2e --- /dev/null +++ b/wanvideo/schedulers/vitb_unipc.py @@ -0,0 +1,41 @@ +from diffusers.schedulers import UniPCMultistepScheduler +import torch + + +class ViBTScheduler(UniPCMultistepScheduler): + def __init__(self, **kwargs): + super().__init__(**{**kwargs, "use_flow_sigmas": True}) + self.set_parameters() + + def set_parameters(self, noise_scale=1.0, shift=5.0, seed=None): + self.noise_scale = noise_scale + self.config.flow_shift = shift + + def step(self, model_output, timestep, sample, generator, **kwargs): + delta_t = ( + max(self.timesteps[self.timesteps < timestep]) - timestep + if any(self.timesteps < timestep) + else -timestep - 1 + ) / 1000 + + current_t = (timestep + 1) / 1000.0 + eta = (-delta_t * (current_t + delta_t) / current_t) ** 0.5 + + noise = torch.randn( + sample.shape, + generator=generator, + device=torch.device("cpu"), + dtype=sample.dtype, + ).to(sample.device) + latents = sample + delta_t * model_output + eta * self.noise_scale * noise + + return (latents,) + + @classmethod + def from_scheduler( + cls, scheduler: UniPCMultistepScheduler, noise_scale=1.0, shift_gamma=5.0 + ): + obj = cls.__new__(cls) + obj.__dict__ = scheduler.__dict__.copy() + obj.set_parameters(noise_scale, shift_gamma) + return obj \ No newline at end of file