Add ViBTScheduler to use ViBT models
https://github.com/Yuanshi9815/ViBT/tree/main
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user