You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
On non-MPS platforms (CUDA/CPU), StableDiffusionXLPipeline can call vae.decode() with fp16 latents while the VAE (or parts of it) are fp32, which causes a hard runtime error in normalization/linear layers (e.g. GroupNorm): expected scalar type Half but found Float.
This happens because the pipeline currently only aligns latents dtype when needs_upcasting is True (fp16 VAE + force_upcast), and the elif latents.dtype != self.vae.dtype: branch only handles MPS by casting the VAE to the latents dtype. On CUDA/CPU there is no dtype/device alignment, so mixed dtype can reach VAE decode.
Reproduction
Environment: diffusers==0.36.0.dev0 (observed), CUDA or CPU (non-MPS)
Steps (conceptual):
Ensure VAE is fp32 (or has fp32 submodules) while latents become fp16.
Run StableDiffusionXLPipeline.__call__ with output_type != "latent".
Pipeline reaches vae.decode(latents, ...) and errors inside VAE decoder GroupNorm/Linear due to fp16 input + fp32 weights.
A concrete regression test is included to reproduce this without GPU:
Force pipe.vae to fp32
Use callback_on_step_end to force latents to fp16
Assert that the pipeline aligns the dtype back to fp32 before calling vae.decode
Fix
When needs_upcasting is False but latents.dtype != self.vae.dtype, we now align latentsdtype/device to the VAE decode dtype/device (preferring vae.post_quant_conv parameters when available) on non-MPS platforms. This prevents mixed dtype from reaching vae.decode() and matches the intent of the upcast path.
Tests
Added test_vae_decode_aligns_latents_dtype_when_vae_is_fp32 in tests/pipelines/stable_diffusion_xl/test_stable_diffusion_xl.py.
Why this is a bug
Users can legitimately end up with fp32 VAE (stability) while latents are fp16 (performance / callbacks / schedulers). The pipeline should not crash with dtype mismatch in this scenario; it should deterministically align latents to the VAE decode dtype.
Thanks for the context and the note about deprecating VAE upcasting 1.
To clarify: the issue here is not about users intentionally running “fp32 VAE + fp16 pipeline” as a manual setup. The crash can happen whenever the VAE (or parts of it) are fp32 (which can occur for stability reasons / partial fp32 modules) while latents end up fp16 (e.g. via callback_on_step_end, scheduler/hook behavior, or external integrations). In that case, on non‑MPS platforms the current branch
elif latents.dtype != self.vae.dtype: ...
does not align anything, so mixed dtypes can reach vae.decode() and fail inside VAE decoder GroupNorm/Linear with Half/Float mismatch.
This PR adds a minimal, deterministic safety alignment at the decode boundary only when there is a mismatch, by casting latents to the VAE decode dtype/device (preferring post_quant_conv params, consistent with the existing needs_upcasting path). It doesn’t change behavior for the common case where dtypes already match.
The included regression test reproduces the scenario without GPU by forcing VAE fp32 and forcing latents fp16 via callback, then asserting that vae.decode() receives fp32 latents.
This issue has been automatically marked as stale because it has not had recent activity. If you think this still needs to be addressed please comment on this thread.
Please note that issues that do not follow the contributing guidelines are likely to be ignored.
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
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.
Summary
On non-MPS platforms (CUDA/CPU),
StableDiffusionXLPipelinecan callvae.decode()with fp16 latents while the VAE (or parts of it) are fp32, which causes a hard runtime error in normalization/linear layers (e.g.GroupNorm):expected scalar type Half but found Float.This happens because the pipeline currently only aligns
latentsdtype whenneeds_upcastingisTrue(fp16 VAE +force_upcast), and theelif latents.dtype != self.vae.dtype:branch only handles MPS by casting the VAE to the latents dtype. On CUDA/CPU there is no dtype/device alignment, so mixed dtype can reach VAE decode.Reproduction
diffusers==0.36.0.dev0(observed), CUDA or CPU (non-MPS)StableDiffusionXLPipeline.__call__withoutput_type != "latent".vae.decode(latents, ...)and errors inside VAE decoderGroupNorm/Lineardue to fp16 input + fp32 weights.A concrete regression test is included to reproduce this without GPU:
pipe.vaeto fp32callback_on_step_endto forcelatentsto fp16vae.decodeFix
When
needs_upcastingisFalsebutlatents.dtype != self.vae.dtype, we now alignlatentsdtype/device to the VAE decode dtype/device (preferringvae.post_quant_convparameters when available) on non-MPS platforms. This prevents mixed dtype from reachingvae.decode()and matches the intent of the upcast path.Tests
test_vae_decode_aligns_latents_dtype_when_vae_is_fp32intests/pipelines/stable_diffusion_xl/test_stable_diffusion_xl.py.Why this is a bug
Users can legitimately end up with fp32 VAE (stability) while latents are fp16 (performance / callbacks / schedulers). The pipeline should not crash with dtype mismatch in this scenario; it should deterministically align latents to the VAE decode dtype.