Extend ER-SDE noise scaler by scaling h(t) (#15428)

This commit is contained in:
chaObserv 2026-08-09 10:06:12 +08:00 committed by GitHub
parent a683fa6e57
commit 40e46c7110
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 25 additions and 11 deletions

View File

@ -591,7 +591,7 @@ class SamplerER_SDE(io.ComfyNode):
inputs=[
io.Combo.Input("solver_type", options=["ER-SDE", "Reverse-time SDE", "ODE"]),
io.Int.Input("max_stage", default=3, min=1, max=3, advanced=True),
io.Float.Input("eta", default=1.0, min=0.0, max=100.0, step=0.01, round=False, tooltip="Stochastic strength of reverse-time SDE.\nWhen eta=0, it reduces to deterministic ODE. This setting doesn't apply to ER-SDE solver type.", advanced=True),
io.Float.Input("eta", default=1.0, min=0.0, max=10.0, step=0.01, round=False, tooltip="Stochastic strength of SDEs.\nWhen eta=0, they reduce to deterministic ODE.\nLarge eta may cause invalid outputs. If this occurs, try decreasing this value.", advanced=True),
io.Float.Input("s_noise", default=1.0, min=0.0, max=100.0, step=0.01, round=False, advanced=True),
],
outputs=[io.Sampler.Output()]
@ -599,21 +599,35 @@ class SamplerER_SDE(io.ComfyNode):
@classmethod
def execute(cls, solver_type, max_stage, eta, s_noise) -> io.NodeOutput:
if solver_type == "ODE" or (solver_type == "Reverse-time SDE" and eta == 0):
eta = 0
s_noise = 0
# Extend existing noise scalers phi(x) with eta-controlled noise scalers:
# psi(x) = x**(1-eta) * phi(x)**eta
# where eta is constant and directly scales the h^2(t) contribution.
def reverse_time_sde_noise_scaler(x):
def er_sde_noise_scaler(x: torch.Tensor) -> torch.Tensor:
return x * ((x ** 0.3).exp() + 10.0) ** eta
def reverse_time_sde_noise_scaler(x: torch.Tensor) -> torch.Tensor:
return x ** (eta + 1)
if solver_type == "ER-SDE":
# Use the default one in sample_er_sde()
noise_scaler = None
else:
noise_scaler = reverse_time_sde_noise_scaler
def ode_noise_scaler(x: torch.Tensor) -> torch.Tensor:
return x
solver_scalers = {
"ER-SDE": er_sde_noise_scaler,
"Reverse-time SDE": reverse_time_sde_noise_scaler,
"ODE": ode_noise_scaler,
}
if solver_type == "ODE" or eta == 0:
s_noise = 0.0
solver_type = "ODE"
noise_scaler = solver_scalers[solver_type]
sampler_name = "er_sde"
sampler = comfy.samplers.ksampler(sampler_name, {"s_noise": s_noise, "noise_scaler": noise_scaler, "max_stage": max_stage})
sampler = comfy.samplers.ksampler(
sampler_name,
{"s_noise": s_noise, "noise_scaler": noise_scaler, "max_stage": max_stage},
)
return io.NodeOutput(sampler)
get_sampler = execute