Add experimental rCM scheduler
Based on the original code, works but doesn't feel better than dpm++_sde so far
This commit is contained in:
@@ -1297,6 +1297,12 @@ class WanVideoSampler:
|
||||
extra_channel_latents_input = extra_channel_latents.to(z)
|
||||
z = torch.cat([z, extra_channel_latents_input])
|
||||
|
||||
if scheduler.lower() == "rcm":
|
||||
c_in = 1 / (torch.cos(timestep) + torch.sin(timestep))
|
||||
c_noise = (torch.sin(timestep) / (torch.cos(timestep) + torch.sin(timestep))) * 1000
|
||||
z = z * c_in
|
||||
timestep = c_noise
|
||||
|
||||
base_params = {
|
||||
'x': [z], # latent
|
||||
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
|
||||
@@ -2956,6 +2962,8 @@ class WanVideoSampler:
|
||||
# callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
|
||||
elif humo_reference_count > 0:
|
||||
callback_latent = (latent_model_input[:,:-humo_reference_count].to(device) - noise_pred[:,:-humo_reference_count].to(device) * t.to(device) / 1000).detach()
|
||||
elif scheduler == "rcm":
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device)).detach()
|
||||
else:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach()
|
||||
callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps))
|
||||
|
||||
@@ -6,6 +6,7 @@ from .flowmatch_pusa import FlowMatchSchedulerPusa
|
||||
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 ...utils import log
|
||||
|
||||
try:
|
||||
@@ -26,7 +27,8 @@ scheduler_list = [
|
||||
"flowmatch_distill",
|
||||
"flowmatch_pusa",
|
||||
"multitalk",
|
||||
"sa_ode_stable"
|
||||
"sa_ode_stable",
|
||||
"rcm"
|
||||
]
|
||||
|
||||
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):
|
||||
@@ -105,6 +107,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
elif "sa_ode_stable" in scheduler:
|
||||
sample_scheduler = FlowMatchSAODEStableScheduler(shift=shift, **kwargs)
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
elif 'rcm' in scheduler:
|
||||
sample_scheduler = rCMFlowMatchScheduler()
|
||||
sample_scheduler.set_timesteps(steps, sigma_max=120)
|
||||
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
# based on https://github.com/NVlabs/rcm
|
||||
import torch
|
||||
import math
|
||||
|
||||
class rCMFlowMatchScheduler():
|
||||
|
||||
def __init__(self, num_inference_steps=4, num_train_timesteps=1000, sigma_max=120):
|
||||
self.num_train_timesteps = num_train_timesteps
|
||||
self.sigma_max = sigma_max
|
||||
self.step_index = 0
|
||||
|
||||
self.set_timesteps(num_inference_steps, sigma_max=sigma_max)
|
||||
|
||||
def set_timesteps(self, num_inference_steps=4, sigma_max=120):
|
||||
mid_t = [1.5, 1.4, 1.0][:num_inference_steps - 1]
|
||||
self.timesteps = torch.tensor([math.atan(sigma_max), *mid_t], dtype=torch.float32)
|
||||
self.sigmas = torch.tensor([math.atan(sigma_max), *mid_t, 0], dtype=torch.float32)
|
||||
self.step_index = 0
|
||||
|
||||
def step(self, model_output, timestep, sample, generator):
|
||||
|
||||
c_skip = 1 / (torch.cos(timestep) + torch.sin(timestep))
|
||||
c_out = -1 * torch.sin(timestep) / (torch.cos(timestep) + torch.sin(timestep))
|
||||
|
||||
# Get next timestep
|
||||
if self.step_index + 1 < len(self.sigmas):
|
||||
t_next = self.sigmas[self.step_index + 1]
|
||||
else:
|
||||
t_next = torch.tensor(0.0)
|
||||
|
||||
x = c_skip * sample + c_out * model_output
|
||||
if t_next > 1e-5:
|
||||
x = torch.cos(t_next) * x + torch.sin(t_next) * torch.randn(
|
||||
*x.shape,
|
||||
dtype=torch.float32,
|
||||
device=torch.device("cpu"),
|
||||
generator=generator,
|
||||
).to(x)
|
||||
self.step_index += 1
|
||||
return x
|
||||
Reference in New Issue
Block a user