Fix UniPCMultistepScheduler under a float64 default dtype - #14920
Open
DawnofGenX wants to merge 2 commits into
Open
DawnofGenX wants to merge 2 commits into
DawnofGenX wants to merge 2 commits into
Conversation
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
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Fixes #14888
multistep_uni_p_bh_updateandmultistep_uni_c_bh_updatebuildrksfrom the sigmas (float32) and appendtorch.ones(()), which takes the default dtype. Undertorch.set_default_dtype(torch.float64),torch.stack(rks)promotesRto float64 whilebstays float32, and the firsttorch.linalg.solve(R, b)raises:This happens whatever the dtype of the sample (float16, float32 and float64 all fail).
DPMSolverMultistepScheduler,DEISMultistepSchedulerandSASolverSchedulerrun 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 otherrksentries.Under the default float32 this is the same tensor as before: I compared 144 configurations (
solver_order1–3 ×bh1/bh2× the threeprediction_types × Karras sigmas or not ×predict_x0×thresholding, 10 steps) against pristinemainand all 144 outputs are bit-identical.Test
test_default_dtype_float64intests/schedulers/test_scheduler_unipc.py, modeled on the existingtest_fp16_support. It runs the 10-step loop undertorch.set_default_dtype(torch.float64)acrosssolver_order1–3, both solver types and all three prediction types, asserting the output dtype and finiteness.mainwith thelinalg.solveerror above; passes with the fix.test_compatiblesandtest_beta_sigmasfail onmaintoo, with an unrelatedImportErrorfrom optional dependencies.)Before submitting
self-reviewskill on the diff — notes below.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 fromblock (the nearest markers inscheduling_unipc_multistep.pyare at lines 711 and 1099), so nomake fix-copiespropagation 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,DEISMultistepSchedulerandSASolverSchedulerwere 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.