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:
kijai
2025-10-19 19:34:17 +03:00
parent 7ecd55e92a
commit 9cd79d3d4a
3 changed files with 55 additions and 1 deletions
+8
View File
@@ -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))
+7 -1
View File
@@ -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
+40
View File
@@ -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