Hacktoberfest 2026: the issues maintainers tagged for October, open and beginner-friendly. Browse Hacktoberfest issues

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

Open Beginner friendly
#14,888 0 comments 0 reactions 0 assignees View on GitHub

Maintainers usually reply within 1 day

Nobody has claimed this yet.

Assessment

Difficulty
2/5
Estimated time
1-3 hours
Newbie friendliness
88/100
Issue type
Bug
Clarity
Clearly specified
Activity status
Active
Tech stack
python, pytorch

Research direction

Start in src/diffusers/schedulers/scheduling_unipc_multistep.py at the rks construction in multistep_uni_p_bh_update and multistep_uni_c_bh_update, then read tests/schedulers/test_scheduler_unipc.py and run the default-dtype coverage. Done means the scheduler runs with torch.float64 as the default dtype without a linalg.solve dtype error, while existing behavior and tests remain unchanged.

Written by the indexing model from the issue text.

Description

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()
Dominant language
Python
Stars
34.6k
Forks
7.4k
Avg merge
4d 8h
Merged PRs (30d)
50

Getting set up

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

More from huggingface/diffusers

All issues in huggingface/diffusers

Similar issues

More Python issues

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.