Skip to content

Fix UniPCMultistepScheduler under a float64 default dtype - #14920

Open
DawnofGenX wants to merge 2 commits into
huggingface:mainfrom
DawnofGenX:fix-unipc-default-dtype-14888
Open

DawnofGenX wants to merge 2 commits into
huggingface:mainfrom
DawnofGenX:fix-unipc-default-dtype-14888

Conversation

@DawnofGenX

Copy link
Copy Markdown

Fixes #14888

multistep_uni_p_bh_update and multistep_uni_c_bh_update build rks from the sigmas (float32) and append torch.ones(()), which takes the default dtype. Under torch.set_default_dtype(torch.float64), torch.stack(rks) promotes R to float64 while b stays float32, and the first torch.linalg.solve(R, b) 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

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

Change

torch.ones((), device=device) → torch.ones_like(h) at both sites, so the unit entry takes the dtype and device of the other rks entries.

Under the default float32 this is the same tensor as before: I compared 144 configurations (solver_order 1–3 × bh1/bh2 × the three prediction_types × Karras sigmas or not × predict_x0 × thresholding, 10 steps) against pristine main and all 144 outputs are bit-identical.

Test

test_default_dtype_float64 in tests/schedulers/test_scheduler_unipc.py, modeled on the existing test_fp16_support. It runs the 10-step loop under torch.set_default_dtype(torch.float64) across solver_order 1–3, both solver types and all three prediction types, asserting the output dtype and finiteness.

  • Fails on main with the linalg.solve error above; passes with the fix.
  • The file's tests: 61 passed. (test_compatibles and test_beta_sigmas fail on main too, with an unrelated ImportError from optional dependencies.)

Before submitting

Self-review notes

Ran against .ai/references/review-rules.md.

Blocking issues: none.

Non-blocking issues: none.

Copied code: neither call site is inside a # Copied from block (the nearest markers in scheduling_unipc_multistep.py are at lines 711 and 1099), so no make fix-copies propagation is required.

Scope: grep -rn "torch.ones((), device=device)" src/diffusers/schedulers/ returns only these two sites, so no other scheduler carries this pattern. DPMSolverMultistepScheduler, DEISMultistepScheduler and SASolverScheduler were checked and already run under a float64 default dtype.

Dtype impact: the fix is a strict narrowing of the unit entry's dtype to match its neighbours. Under the default float32 the tensor is unchanged (144/144 configurations bit-identical), so there is no precision regression on the existing path.

Covers the float64 default-dtype path that fails with a dtype mismatch in torch.linalg.solve. Modeled on test_fp16_support.
multistep_uni_p_bh_update and multistep_uni_c_bh_update built rks from the
float32 sigmas and appended torch.ones(()), which takes the default dtype.
Under torch.set_default_dtype(torch.float64) torch.stack promoted R to
float64 while b stayed float32, so torch.linalg.solve raised:

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

Use torch.ones_like(h) so the unit entry has the dtype and device of the
other rks entries. Under the default float32 this is the same tensor as
before, so existing behaviour is unchanged.

Fixes huggingface#14888

This branch has not been deployed

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

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

1 participant