Fix QwenImage21 MPS temporal padding corruption - #14902
vinitsonawane45 wants to merge 2 commits into
Conversation
|
Thanks for turning this around so quickly. The code change matches what I validated (same logic, applied as a monkeypatch on 0121a91), so that part looks good to me. A few notes: 1. The regression test passes with the old code too. It uses a 1×1×3×1×1 tensor on CPU, far below the size where MPS breaks, so it only checks that the new implementation equals the old one. A test that actually fails without the fix needs MPS and a large enough tensor. This one fails on the old code and passes with your change on my M5 Max (torch 2.14.0): @pytest.mark.skipif(not torch.backends.mps.is_available(), reason="the F.pad bug is specific to MPS")
def test_temporal_downsample_matches_cpu_on_mps(self):
torch.manual_seed(0)
downsample = QwenImage21AvgDown3D(in_channels=192, out_channels=384, factor_t=2, factor_s=2)
# 256 x 256 = 65,536 trailing elements: the size where temporal F.pad breaks on MPS
sample = torch.randn(1, 192, 1, 256, 256)
torch.testing.assert_close(downsample(sample.to("mps")).cpu(), downsample(sample))It won't run in CI without MPS, but it will on any Apple Silicon dev machine. Keeping your CPU equivalence test next to it makes sense. 2. Correction to my earlier comment. I wrote that it fails "for any padded axis of a 5-D tensor"; I had only checked that on a single-channel tensor 3. Validation numbers (same logic as the PR, MPS bf16 vs. CPU fp32):
To be precise about the PR description: I validated the same change as a monkeypatch, not your branch itself. Happy to run the branch directly if that helps the review. |
|
Thanks @ck71 for the detailed feedback. I’ve added the requested MPS-specific regression test at the 256×256 / 65,536 boundary while keeping the existing CPU equivalence test. The update is pushed as commit Please take another look when you have a chance, and I’m happy to make any further changes you recommend. |
|
Thanks! I ran the branch itself (7ef0f38) on the M5 Max this time: both |
Summary
Avoid
F.padfor temporal padding inQwenImage21AvgDown3D.On MPS,
torch.nn.functional.padcan produce incorrect results for sufficiently large 5-D tensors. Qwen-Image 2.1 can hit this condition during VAE encoding at certain image aspect ratios and resolutions, causing severe output degradation.This change replaces the temporal
F.padoperation with equivalent zero-padding usingtorch.cat.Root cause
The issue was reproduced on MPS when the affected 5-D tensor dimensions reached the relevant size threshold. The corruption occurs inside the temporal padding operation in
QwenImage21AvgDown3D.The problem is aspect-ratio dependent. At
output_resolution=1024, several aspect ratios can cross the affected tensor-size boundary, while others remain below it.Changes
F.padoperation withtorch.catand a zero-initialized tensor.256 × 256/65,536boundary where the original implementation produces incorrect results.Validation
git diff --check: passed7ef0f38) on an M5 Max.The full pytest collection was blocked locally by an existing environment incompatibility involving
huggingface_hub(resolve_revisionimport error); no dependency changes were made as part of this fix.Related issues
Fixes #14858
Related to #14859