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 .scheduling_flow_match_lcm import FlowMatchLCMScheduler
|
||||||
from .fm_sa_ode import FlowMatchSAODEStableScheduler
|
from .fm_sa_ode import FlowMatchSAODEStableScheduler
|
||||||
from .fm_rcm import rCMFlowMatchScheduler
|
from .fm_rcm import rCMFlowMatchScheduler
|
||||||
|
from .vitb_unipc import ViBTScheduler
|
||||||
from ...utils import log
|
from ...utils import log
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -29,12 +30,17 @@ scheduler_list = [
|
|||||||
"flowmatch_pusa",
|
"flowmatch_pusa",
|
||||||
"multitalk",
|
"multitalk",
|
||||||
"sa_ode_stable",
|
"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):
|
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
|
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)
|
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
|
||||||
if sigmas is None:
|
if sigmas is None:
|
||||||
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
|
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.sigmas = sigmas.to(device)
|
||||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||||
|
|
||||||
elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']:
|
elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']:
|
||||||
if 'longcat' in scheduler:
|
if 'longcat' in scheduler:
|
||||||
num_distill_sample_steps = 50
|
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