Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions src/diffusers/schedulers/scheduling_unipc_multistep.py
Original file line number Diff line number Diff line change
Expand Up @@ -903,7 +903,7 @@ def multistep_uni_p_bh_update(
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 = []
Expand Down Expand Up @@ -1038,7 +1038,7 @@ def multistep_uni_c_bh_update(
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 = []
Expand Down
35 changes: 35 additions & 0 deletions tests/schedulers/test_scheduler_unipc.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,6 +257,41 @@ def test_fp16_support(self):

assert sample.dtype == torch.float16

def test_default_dtype_float64(self):
"""The unit entry of ``rks`` must take the sample's dtype, not the default dtype.

``torch.ones(())`` takes the default dtype, so under
``torch.set_default_dtype(torch.float64)`` the stacked ``R`` was promoted
to float64 while ``b`` stayed float32 and ``torch.linalg.solve`` raised.
"""
default_dtype = torch.get_default_dtype()
torch.set_default_dtype(torch.float64)
try:
for order in [1, 2, 3]:
for solver_type in ["bh1", "bh2"]:
for prediction_type in ["epsilon", "sample", "v_prediction"]:
scheduler_class = self.scheduler_classes[0]
scheduler_config = self.get_scheduler_config(
prediction_type=prediction_type,
solver_order=order,
solver_type=solver_type,
)
scheduler = scheduler_class(**scheduler_config)

num_inference_steps = 10
model = self.dummy_model()
sample = self.dummy_sample_deter.to(torch.float64)
scheduler.set_timesteps(num_inference_steps)

for i, t in enumerate(scheduler.timesteps):
residual = model(sample, t)
sample = scheduler.step(residual, t, sample).prev_sample

assert sample.dtype == torch.float64
assert torch.isfinite(sample).all()
finally:
torch.set_default_dtype(default_dtype)

def test_full_loop_with_noise(self):
scheduler_class = self.scheduler_classes[0]
scheduler_config = self.get_scheduler_config()
Expand Down
Loading