Add ViBTScheduler to use ViBT models

https://github.com/Yuanshi9815/ViBT/tree/main
This commit is contained in:
kijai
2025-12-02 00:26:50 +02:00
parent c4ca252fea
commit c1fbc93521
2 changed files with 49 additions and 3 deletions
+8 -3
View File
@@ -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
+41
View File
@@ -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