[diffusion] Fix Z-Image accuracy - #29742
Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces support for batched inference, attention masking, and custom RMS normalization (ZImageRMSNorm) to match the official Z-Image implementation. It also fixes a batch-dimension issue in the Qwen3 encoder's default position IDs and adds corresponding unit tests. The reviewer suggests caching the batched freqs_cis in the Z-Image model's forward pass to avoid redundant computations across denoising steps, which would improve serving throughput.
Important
The consumer version of Gemini Code Assist on GitHub is being sunset. Starting June 18, 2026, new organization installations will be blocked, and all code review activity will officially cease on July 17, 2026.
For more details on the timeline and next steps, please review the Help Documentation.
| if len(input_images) > 1 and get_sp_world_size() == 1: | ||
| freqs_cis = self._build_batched_freqs_cis( | ||
| input_images, | ||
| input_cap_feats, | ||
| patch_size, | ||
| f_patch_size, | ||
| image_target_len=x.shape[1], | ||
| cap_target_len=cap_feats.shape[1], | ||
| ) |
There was a problem hiding this comment.
In the current implementation, _build_batched_freqs_cis is called on every single denoising step when batching is enabled (len(input_images) > 1). Since the shapes and devices of the input images and caption features do not change across denoising steps within a request, rebuilding the batched freqs_cis on every step introduces redundant overhead (such as coordinate grid creation, rotary embedding computation, padding, and stacking).
We can cache the batched freqs_cis based on the input shapes, devices, and patch parameters to completely avoid this redundant computation and improve serving throughput.
if len(input_images) > 1 and get_sp_world_size() == 1:
cache_key = (
len(input_images),
tuple(img.shape for img in input_images),
tuple(cap.shape for cap in input_cap_feats),
patch_size,
f_patch_size,
x.shape[1],
cap_feats.shape[1],
device,
)
if (
getattr(self, "_cached_batched_freqs_cis_key", None) == cache_key
and getattr(self, "_cached_batched_freqs_cis", None) is not None
):
freqs_cis = self._cached_batched_freqs_cis
else:
freqs_cis = self._build_batched_freqs_cis(
input_images,
input_cap_feats,
patch_size,
f_patch_size,
image_target_len=x.shape[1],
cap_target_len=cap_feats.shape[1],
)
self._cached_batched_freqs_cis_key = cache_key
self._cached_batched_freqs_cis = freqs_cis|
/tag-and-rerun-ci |
There was a problem hiding this comment.
excellent. could you help confirm:
- this kind of problem only happens with Z-Image, and,
- does this ground truth needs to be updated, or we just need to tighten the consistency threshold? repro script is here
|
|
||
| image_pos = torch.arange(image_target_len, device=device).unsqueeze(0) | ||
| cap_pos = torch.arange(cap_target_len, device=device).unsqueeze(0) | ||
| image_len = torch.tensor(image_lengths, device=device).unsqueeze(1) |
|
/rerun-failed-ci |
i tested across some other models, and didn't seem to see the same problem - gt also seems to be okay and i don't think it needs to be updated, let's see if ci passes |
92f2b43 to
34cd67e
Compare
|
/rerun-failed-ci |
d7b91b6 to
15c3f12
Compare
|
could you resolve the conflict? cheers |
11ecbc5 to
e4d663c
Compare
|
/rerun-failed-ci |
|
/tag-and-rerun-ci |
Co-authored-by: Mick <mickjagger19@icloud.com>
Motivation
Fix #28502 , and general accuracy issues with Z-image-Turbo
Modifications
With
--batching-mode dynamic,Tongyi-MAI/Z-Image-Turbocould produce severely degraded images when requests were batched.Initially, i found part of this issue was that qwen3 text encoder default
position_idswere shaped[1, seq]even whenhidden_stateswere[batch, seq, dim], causing RoPE to be applied with the wrong flattened token layout for batched prompts,But following this, I realize that our current z-image turbo implementation is severely degraded, even for singleton generations, compared to the native pytorch implementation:
Firstly, z-image sampling and normalization differed from the native implementation above; we use the scheduler default sigma path, outer autocast, and shared fp32-accumulating RMSNorm, changing the denoising trajectory and bf16 activations for this model
Secondly, we padded batched images/captions to shared batch len but did not preserve RoPE offsets for each request or mask extra batch-only padding in attention, which causes mixed-prompt dynamic batches to attend to invalid tokens and use wrong image positions
Accuracy Tests
Tested on 1xh100
Prompts tested:
Draw Donald Duck and Mickey Mouse fishing by the river.Crayon Shin chan walks with a puppy on the street.Naruto and Luffy sit side by side on a large boulder. On the left, Naruto has a smile on his face and raises one hand to make a victory gesture. On the right, Luffy wears a wide-brimmed straw hat, an unbuttoned short-sleeved jacket, shorts and flip-flops, with a long sword strapped to his back. He is also grinning happily. Behind them lie a wide expanse of water and the sky dotted with a few clouds, and patches of grass grow beside the boulder.SpongeBob SquarePants and Patrick sat at the table having dinner together, and there was a TV in the middle of the house with an old phone next to it.Mario is driving a go kart on the highway, surrounded by many low trees.Speed Tests and Profiling
1xh100, 1024x1024, 9 steps single generation
863.27 ms851.57 msbatch generation size 5:
4907.5 ms4893.5 ms3749.0 ms416.3 ms4781.3 ms4765.0 ms3691.5 ms409.9 msabout 2% faster across single and batch generation
Checklist
CI States
Latest PR Test (Base): ❌ Run #28860860153
Latest PR Test (Extra): ❌ Run #28860859883