Repository navigation
Qwen Image 2.1 transformer is incompatible with torch.compile #14821
Description
Activity
Hi @abel1502, thanks for the report and the repro.
I agree with the diagnosis.
pipe.transformer.compile(..., fullgraph=True)tracesQwenImage21Transformer2DModel.forward, which callsQwenImage21Rope.forward, and that doesis_image_token = image_pad_mask.tolist()(transformer_qwenimage21.pyaround line 686). The traceback matches the source onmain. The inductor flags in the repro are unrelated.Root cause:
image_pad_maskis bool. It istorch.repeat_interleave(img_mask[0], repeats), and the pipeline buildsimg_maskas a bool tensor. Dynamo only tracesTensor.tolist()forint8,int16,int32, andint64(gb0109), so a bool mask fails before the rest of the forward is traced.fullgraph=Trueturns that into a hard error. This isQwenImage21Rope, which runs for both attention processors.Casting the mask with
.to(torch.int32).tolist()only satisfies the dtype check. Integertolist()is lowered to one.item()per token, and the Pythonlist.index/rangewalk then branches on those values. That is still a host sync and still not one graph. This file already treatstolist()as a device sync in_qwenimage21_prefix_segments.The same forward has three more host-dependent spots, so fixing RoPE alone still leaves this repro unable to compile:
build_token_metadatausesnonzero, so the index vector length depends on mask values.use_kv_cachedefaults toTrue, so step 0 is"extract"and later steps are"cached".prefix_len = int((~target_token_mask).sum())syncs a scalar on every step and then slices with it.- On the default
QwenImage21AttnProcessor,_qwenimage21_prefix_segmentscallstolist()on the whole prefix.
compile_repeated_blocks()does not cover this. RoPE runs in the parentforward, and_repeated_blocksis onlyQwenImage21TransformerBlock.Proposed approach:
img_shapesalready gives each block's(frame, height, width)as Python ints. The mask only says where those blocks sit between text tokens. The current walk is:- text tokens get
frame = height = width, advancing by 1 - an image block shares one frame id and a zero-centered height/width grid
- after the block, the position advances by
max(height, width)
Build that with static-shape tensor ops. Let
maskbe 0/1 andseq_len = mask.shape[0]:img_before = cumsum(mask) - masktext_before = cumsum(1 - mask) - (1 - mask)- for each block, once
seenimage tokens have been consumed, addmax(height, width)whereverimg_before >= seen frame_index = text_before + that advance
Checked against the current Python walk for text/image/text, adjacent images, and a trailing text tail. For example, mask
F F T T T T Fwith one2x2block produces frame ids[0, 1, 2, 2, 2, 2, 4].The centered grid does not need the mask. Per block:
start = -(length - length // 2) torch.arange(start, start + length)
Repeat height across width and width down height, then concatenate in block order.
masked_scatterwrites that static vector onto theTruepositions.height_indexandwidth_indexstart as clones offrame_index, so text tokens keep the frame id on all three axes.torch._check(mask.sum() == image_token_count)keeps the scatter length static for Dynamo. Eager keeps today'sValueErrorwhenimg_shapesand the mask disagree.The other three spots, in the same change:
- The image-token count is already
sum(math.prod(shape) for shape in img_shapes). Buildimage_idsandtarget_token_maskwithmasked_scatterof a static id vector, and use the sametorch._checkinstead ofnonzero. - The pipeline appends the target as a suffix, so
prefix_len = seq_len - math.prod(img_shapes[0][-1])matchesint((~target_token_mask).sum())with no.item(). - Prefix segment count is static: one text gap and one image block per condition image, then the text before the target. Block starts come from
argmaxon the image cumsum. The SDPA loop stays a Python loop over that static count, with 0-dim tensor slice bounds. A zero-length text gap (two image blocks with no text between them, whichbuild_token_metadataalready calls out) is an empty slice, andtorch.catdrops it.
QwenImage21Ropealso rebindsself.freqswith.to(device)on every forward. Those tables are fixed at init and are not buffers, so they stay unregistered.QwenEmbedRope._get_device_freqsalready caches this transfer behindlru_cache_unless_export. I would use that here so the first call moves them and later calls do not.Tests:
On the tiny config already in
tests/models/transformers/test_models_transformer_qwenimage21.py:- Eager RoPE matches the current Python implementation for text/image/text, adjacent images, odd and even spatial sizes, and a target-only sequence.
torch.compile(model, fullgraph=True)matches eager, then a second call witherror_on_recompile=True.- Same for
QwenImage21AttnProcessorandQwenImage21FlexAttnProcessor. - KV cache
"extract"then"cached"underfullgraph=True. - Add
TorchCompileTesterMixinto this model. The Qwen-Image 1.x transformer tests already have it. This 2.1 file does not, which is consistent with compile being broken.
I would be glad to open a PR for this once the approach looks right.
Reacted by abel1502
Describe the bug
QwenImage21Ropefeatures a call toimage_pad_mask.tolist(). When the transformer in the Qwen Image 2.1 pipeline istorch.compile-ed and then run, it causes the following error:Reproduction
System Info
Who can help?
@naykun @yiyixuxu , based on https://github.com/huggingface/diffusers/blame/main/src/diffusers/models/transformers/transformer_qwenimage21.py