Skip to content

UniPCMultistepScheduler returns nan with solver_type="bh1", and inf or nan with lower_order_final=False or predict_x0=False, on the final step to sigma = 0 #14887

Description

@Nicholas022400701

Describe the bug

Since #7517 final_sigmas_type defaults to "zero", so the last step goes to sigma = 0, where lambda_t = log(alpha_t) - log(sigma_t) and with it h are inf. Only the default path (solver_type="bh2", predict_x0=True, lower_order_final=True) survives that step:

  1. solver_type="bh1" returns all-nan, for every solver_order, step count and prediction_type, also with use_flow_sigmas=True. The last step runs at order 1, so in multistep_uni_p_bh_update D1s is None, pred_res = 0 and x_t = x_t_ - alpha_t * B_h * pred_res with B_h = hh = -inf, which is -inf * 0 = nan. The docstring recommends bh1 for unconditional sampling under 10 steps.
  2. lower_order_final=False keeps solver_order on the last step: the rks = (lambda_si - lambda_s0) / h are 0 and D1s = (mi - m0) / rk are inf, so the sample comes back inf at order 2, and at order 3 torch.linalg.solve raises _LinAlgError: The solver failed because the input matrix is singular. Both solver types.
  3. predict_x0=False (noise prediction) returns all-nan with both solver types: x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0 is 0 * inf there.

DPMSolverMultistepScheduler handles the same step by forcing the first-order update when final_sigmas_type == "zero" (the lower_order_final = ... or self.config.final_sigmas_type == "zero" line in step) and by rejecting algorithm_type="dpmsolver" together with final_sigmas_type="zero" in __init__. The scheduler tests do not reach this because get_scheduler_config in tests/schedulers/test_scheduler_unipc.py sets final_sigmas_type="sigma_min", and test_fp16_support, the one full loop with bh1, checks the dtype but not finiteness.

Proposed fix, patch below:

  • step: this_order = 1 on the last step when final_sigmas_type == "zero", as in DPMSolverMultistepScheduler. The first-order update there is exactly x_t = m0, the data prediction.
  • multistep_uni_p_bh_update: subtract the B_h term only when there is a correction (D1s is not None). At every other step that term was B_h * 0 = 0.
  • __init__: ValueError for predict_x0=False with final_sigmas_type="zero", worded as in DPMSolverMultistepScheduler. The first-order noise-prediction step could instead be written in its finite form, alpha_t / alpha_s0 * x - (alpha_t * sigma_s0 / alpha_s0 - sigma_t) * m0, if you would rather keep that combination working; happy to do that instead.

Checked against main over 1440 configurations (solver_order 1 to 3, bh1/bh2, the three prediction_types, Karras sigmas or not, lower_order_final, thresholding, both final_sigmas_types, 1/2/3/10/25 steps, the dummy model of the scheduler tests): the 916 outputs that are finite on main are bit-identical (32 prediction_type="sample" outputs differ only in the sign of exact zeros), and the 456 outputs that are nan, inf or a _LinAlgError on main, all with final_sigmas_type="zero", are finite. Two new tests in test_scheduler_unipc.py fail on main and pass with the patch, the file's 66 tests pass, ruff and utils/check_copies.py are clean. I can open the PR if this looks right to you.

Not touched here: with use_karras_sigmas=True and final_sigmas_type="sigma_min" the last two sigmas coincide, so h = 0, and lower_order_final=False with solver_order=3 returns nan through 0 / 0 in the higher-order terms, on main and with this patch alike. The first-order update returns the sample unchanged there.

Reproduction

import torch

from diffusers import UniPCMultistepScheduler

sample = torch.rand(1, 3, 8, 8, generator=torch.Generator().manual_seed(0))


def run(**config):
    scheduler = UniPCMultistepScheduler(**config)
    scheduler.set_timesteps(10)
    x = sample.clone()
    try:
        for t in scheduler.timesteps:
            x = scheduler.step(x * t / (t + 1), t, x).prev_sample  # the dummy model of the scheduler tests
    except Exception as e:
        return type(e).__name__
    if x.isnan().any():
        return f"nan in {x.isnan().float().mean():.0%} of the elements"
    return "inf" if x.isinf().any() else f"finite, mean |x| = {x.abs().mean():.4f}"


print("bh2, defaults                        ", run())
print("bh1                                  ", run(solver_type="bh1"))
print("bh1, order 3                         ", run(solver_type="bh1", solver_order=3))
print("bh1, flow sigmas                     ", run(solver_type="bh1", use_flow_sigmas=True, prediction_type="flow_prediction"))
print("bh2, lower_order_final=False         ", run(lower_order_final=False))
print("bh2, lower_order_final=False, order 3", run(lower_order_final=False, solver_order=3))
print("bh2, predict_x0=False                ", run(predict_x0=False))
print("bh1, final_sigmas_type='sigma_min'   ", run(solver_type="bh1", final_sigmas_type="sigma_min"))

Logs

# diffusers 0.40.0 and main
bh2, defaults                         finite, mean |x| = 0.2337
bh1                                   nan in 100% of the elements
bh1, order 3                          nan in 100% of the elements
bh1, flow sigmas                      nan in 100% of the elements
bh2, lower_order_final=False          inf
bh2, lower_order_final=False, order 3 _LinAlgError
bh2, predict_x0=False                 nan in 100% of the elements
bh1, final_sigmas_type='sigma_min'    finite, mean |x| = 0.2422

# with the patch
bh2, defaults                         finite, mean |x| = 0.2337
bh1                                   finite, mean |x| = 0.2392
bh1, order 3                          finite, mean |x| = 0.2393
bh1, flow sigmas                      finite, mean |x| = 0.1884
bh2, lower_order_final=False          finite, mean |x| = 0.2337
bh2, lower_order_final=False, order 3 finite, mean |x| = 0.2349
bh2, predict_x0=False                 ValueError: `final_sigmas_type` zero is not supported for `predict_x0=False`. Please choose `sigma_min` instead.
bh1, final_sigmas_type='sigma_min'    finite, mean |x| = 0.2422

System Info

  • diffusers 0.40.0 (PyPI) and main at e0abab8
  • torch 2.14.0+cpu, numpy 2.2.6
  • Python 3.13, Linux

AI disclosure: I used an AI coding agent to help find this, to write the reproduction, the patch, the tests and this report. I have read and checked the report, the patch and the tests myself and I will answer questions personally.

Patch

git diff against main (2 files, +52 -4)
diff --git a/src/diffusers/schedulers/scheduling_unipc_multistep.py b/src/diffusers/schedulers/scheduling_unipc_multistep.py
index 5c2cbcc..c73d7e1 100644
--- a/src/diffusers/schedulers/scheduling_unipc_multistep.py
+++ b/src/diffusers/schedulers/scheduling_unipc_multistep.py
@@ -275,6 +275,11 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin):
             else:
                 raise NotImplementedError(f"{solver_type} is not implemented for {self.__class__}")
 
+        if not predict_x0 and final_sigmas_type == "zero":
+            raise ValueError(
+                f"`final_sigmas_type` {final_sigmas_type} is not supported for `predict_x0=False`. Please choose `sigma_min` instead."
+            )
+
         self.predict_x0 = predict_x0
         # setable values
         self.num_inference_steps = None
@@ -945,16 +950,16 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin):
             x_t_ = sigma_t / sigma_s0 * x - alpha_t * h_phi_1 * m0
             if D1s is not None:
                 pred_res = torch.einsum("k,bkc...->bc...", rhos_p, D1s)
+                x_t = x_t_ - alpha_t * B_h * pred_res
             else:
-                pred_res = 0
-            x_t = x_t_ - alpha_t * B_h * pred_res
+                x_t = x_t_
         else:
             x_t_ = alpha_t / alpha_s0 * x - sigma_t * h_phi_1 * m0
             if D1s is not None:
                 pred_res = torch.einsum("k,bkc...->bc...", rhos_p, D1s)
+                x_t = x_t_ - sigma_t * B_h * pred_res
             else:
-                pred_res = 0
-            x_t = x_t_ - sigma_t * B_h * pred_res
+                x_t = x_t_
 
         x_t = x_t.to(x.dtype)
         return x_t
@@ -1210,6 +1215,10 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin):
         else:
             this_order = self.config.solver_order
 
+        if self.config.final_sigmas_type == "zero" and self.step_index == len(self.timesteps) - 1:
+            # the step to sigma = 0 has h = inf, only the first-order update is finite there
+            this_order = 1
+
         self.this_order = min(this_order, self.lower_order_nums + 1)  # warmup for multistep
         assert self.this_order > 0
 
diff --git a/tests/schedulers/test_scheduler_unipc.py b/tests/schedulers/test_scheduler_unipc.py
index ac7e1d3..fbcaf49 100644
--- a/tests/schedulers/test_scheduler_unipc.py
+++ b/tests/schedulers/test_scheduler_unipc.py
@@ -257,6 +257,45 @@ class UniPCMultistepSchedulerTest(SchedulerCommonTest):
 
                     assert sample.dtype == torch.float16
 
+    def test_last_step_to_sigma_zero_is_the_data_prediction(self):
+        # h is inf on the step to sigma = 0: bh1 evaluated B_h * 0 with B_h = -inf and returned nan, and
+        # lower_order_final=False kept the higher-order terms there and returned inf or hit a singular system
+        for solver_type, solver_order, lower_order_final in [("bh1", 2, True), ("bh1", 3, False), ("bh2", 3, False)]:
+            scheduler_class = self.scheduler_classes[0]
+            scheduler_config = self.get_scheduler_config(
+                solver_type=solver_type,
+                solver_order=solver_order,
+                lower_order_final=lower_order_final,
+                final_sigmas_type="zero",
+            )
+            scheduler = scheduler_class(**scheduler_config)
+
+            num_inference_steps = 10
+            model = self.dummy_model()
+            sample = self.dummy_sample_deter
+            scheduler.set_timesteps(num_inference_steps)
+
+            for t in scheduler.timesteps[:-1]:
+                residual = model(sample, t)
+                sample = scheduler.step(residual, t, sample).prev_sample
+
+            t = scheduler.timesteps[-1]
+            residual = model(sample, t)
+            data_prediction = scheduler.convert_model_output(residual, sample=sample)
+            sample = scheduler.step(residual, t, sample).prev_sample
+
+            assert torch.isfinite(sample).all(), (solver_type, solver_order, lower_order_final)
+            assert torch.equal(sample, data_prediction), (solver_type, solver_order, lower_order_final)
+
+    def test_predict_x0_false_requires_final_sigmas_type_sigma_min(self):
+        # the noise-prediction update evaluates sigma_t * expm1(h) = 0 * inf on the step to sigma = 0
+        scheduler_class = self.scheduler_classes[0]
+        with self.assertRaises(ValueError):
+            scheduler_class(**self.get_scheduler_config(predict_x0=False, final_sigmas_type="zero"))
+
+        sample = self.full_loop(predict_x0=False, final_sigmas_type="sigma_min")
+        assert torch.isfinite(sample).all()
+
     def test_full_loop_with_noise(self):
         scheduler_class = self.scheduler_classes[0]
         scheduler_config = self.get_scheduler_config()

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions