From 3f5059f310e981affa613a94b6fdc3ff385b9ec2 Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Fri, 10 Apr 2026 00:07:43 +0100 Subject: [PATCH] fix sign-destroying clamp in ER-SDE higher-order corrections The 2nd/3rd order derivative denominators (lambda_curr - prev_lambda) are always negative since lambda decreases during denoising. The .clamp(min=1e-10) forced these to a tiny positive value, causing d_x0 to explode by ~1e18 and producing NaN output. - Add _safe_div helper that returns zero for near-zero denominators while preserving the sign for normal values - Replace .clamp(min=1e-10) with _safe_div in both 2nd and 3rd order --- modules/schedulers/scheduler_ersde.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/modules/schedulers/scheduler_ersde.py b/modules/schedulers/scheduler_ersde.py index a488edf29..df58a78ee 100644 --- a/modules/schedulers/scheduler_ersde.py +++ b/modules/schedulers/scheduler_ersde.py @@ -306,6 +306,13 @@ class ERSDEScheduler(SchedulerMixin, ConfigMixin): return torch.where(torch.abs(x) < eps, torch.zeros_like(x), x) return 0.0 if abs(x) < eps else x + @staticmethod + def _safe_div(numerator, denominator, eps=1e-10): + """Division with sign-preserving zero protection for lambda differences.""" + if denominator.abs() < eps: + return torch.zeros_like(numerator) + return numerator / denominator + def _compute_fn(self, x): """Evaluate the noise scaling function on a scalar or tensor.""" if isinstance(x, (int, float)): @@ -485,7 +492,7 @@ class ERSDEScheduler(SchedulerMixin, ConfigMixin): elif effective_order == 2: # 2nd order: add first derivative correction - d_x0 = (x0 - self.old_x0) / (lambda_curr - self.prev_lambda).clamp(min=1e-10) + d_x0 = self._safe_div(x0 - self.old_x0, lambda_curr - self.prev_lambda) s_int = self._compute_integral(lambda_next, lambda_curr) s_int = s_int.to(device) delta_lambda = lambda_next - lambda_curr @@ -494,8 +501,8 @@ class ERSDEScheduler(SchedulerMixin, ConfigMixin): elif effective_order == 3: # 3rd order: add second derivative correction - d_x0 = (x0 - self.old_x0) / (lambda_curr - self.prev_lambda).clamp(min=1e-10) - dd_x0 = 2.0 * (d_x0 - self.old_d_x0) / (lambda_curr - self.prev_prev_lambda).clamp(min=1e-10) + d_x0 = self._safe_div(x0 - self.old_x0, lambda_curr - self.prev_lambda) + dd_x0 = self._safe_div(2.0 * (d_x0 - self.old_d_x0), lambda_curr - self.prev_prev_lambda) s_int = self._compute_integral(lambda_next, lambda_curr).to(device) s_d_int = self._compute_derivative_integral(lambda_next, lambda_curr).to(device) delta_lambda = lambda_next - lambda_curr