Skip to content

UniPCMultistepScheduler fails in torch.linalg.solve under a float64 default dtype: the unit entry of rks takes the default dtype #14888

Description

@Nicholas022400701

Describe the bug

multistep_uni_p_bh_update and multistep_uni_c_bh_update build rks from the sigmas, which are float32, and append torch.ones((), device=device), which takes the default dtype. Under torch.set_default_dtype(torch.float64) the torch.stack(rks) promotes R to float64 while b stays float32, and the first torch.linalg.solve(R, b) (the order-2 corrector on the second step with the default config) raises

RuntimeError: linalg.solve: Expected A and B to have the same dtype, but found A of type Double and B of type Float instead

whatever the dtype of the sample (float16, float32 and float64 all fail). DPMSolverMultistepScheduler, DEISMultistepScheduler and SASolverScheduler run under a float64 default dtype.

Fix, patch below: torch.ones_like(h) in both places, so the unit entry has the dtype and device of the other rks entries. Under the float32 default this is the same tensor as before: 540 configurations (solver_order 1 to 3, bh1/bh2, the three prediction_types, Karras sigmas or not, predict_x0, thresholding, 1/2/3/10/25 steps) are bit-identical to main. test_default_dtype_float64 fails on main with the error above and passes with the patch, the file's 64 tests pass, ruff is clean. I can open the PR if this looks right to you.

Reproduction

import torch

from diffusers import UniPCMultistepScheduler

torch.set_default_dtype(torch.float64)

scheduler = UniPCMultistepScheduler()
scheduler.set_timesteps(10)
sample = torch.rand(1, 3, 8, 8, generator=torch.Generator().manual_seed(0), dtype=torch.float32)
for t in scheduler.timesteps:
    sample = scheduler.step(sample * t / (t + 1), t, sample).prev_sample
print(sample.dtype, sample.abs().mean())

Logs

Traceback (most recent call last):
  File "repro.py", line 11, in <module>
    sample = scheduler.step(sample * t / (t + 1), t, sample).prev_sample
  File ".../diffusers/schedulers/scheduling_unipc_multistep.py", line 1191, in step
    sample = self.multistep_uni_c_bh_update(
  File ".../diffusers/schedulers/scheduling_unipc_multistep.py", line 1078, in multistep_uni_c_bh_update
    rhos_c = torch.linalg.solve(R, b).to(device).to(x.dtype)
RuntimeError: linalg.solve: Expected A and B to have the same dtype, but found A of type Double and B of type Float instead

With the patch: torch.float32 tensor(0.2337, dtype=torch.float32), the same value as under the float32 default dtype.

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 test and this report. I have read and checked the report, the patch and the test myself and I will answer questions personally.

Patch

git diff against main (2 files, +15 -2)
diff --git a/src/diffusers/schedulers/scheduling_unipc_multistep.py b/src/diffusers/schedulers/scheduling_unipc_multistep.py
index 5c2cbcc..e8f5af1 100644
--- a/src/diffusers/schedulers/scheduling_unipc_multistep.py
+++ b/src/diffusers/schedulers/scheduling_unipc_multistep.py
@@ -903,7 +903,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin):
             rks.append(rk)
             D1s.append((mi - m0) / rk)
 
-        rks.append(torch.ones((), device=device))
+        rks.append(torch.ones_like(h))
         rks = torch.stack(rks)
 
         R = []
@@ -1038,7 +1038,7 @@ class UniPCMultistepScheduler(SchedulerMixin, ConfigMixin):
             rks.append(rk)
             D1s.append((mi - m0) / rk)
 
-        rks.append(torch.ones((), device=device))
+        rks.append(torch.ones_like(h))
         rks = torch.stack(rks)
 
         R = []
diff --git a/tests/schedulers/test_scheduler_unipc.py b/tests/schedulers/test_scheduler_unipc.py
index ac7e1d3..92ff4e7 100644
--- a/tests/schedulers/test_scheduler_unipc.py
+++ b/tests/schedulers/test_scheduler_unipc.py
@@ -257,6 +257,19 @@ class UniPCMultistepSchedulerTest(SchedulerCommonTest):
 
                     assert sample.dtype == torch.float16
 
+    def test_default_dtype_float64(self):
+        # the unit entry of `rks` took the default dtype, the other entries the dtype of the sigmas
+        default_dtype = torch.get_default_dtype()
+        torch.set_default_dtype(torch.float64)
+        try:
+            sample = self.full_loop(solver_order=3)
+        finally:
+            torch.set_default_dtype(default_dtype)
+        result_mean = torch.mean(torch.abs(sample))
+
+        assert sample.dtype == torch.float64
+        assert abs(result_mean.item() - torch.mean(torch.abs(self.full_loop(solver_order=3))).item()) < 1e-3
+
     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