From 99a4a5887ba1629e1876f5a1f1f67c79d1a79b7e Mon Sep 17 00:00:00 2001 From: Rayane <40967731+Rasaboun@users.noreply.github.com> Date: Mon, 23 Mar 2026 21:43:08 +0100 Subject: [PATCH] Fix Qwen attention mask crash with diffusers >=0.37 (#748) * Fix Qwen Image mask handling * Fix Qwen attention mask crash with diffusers >=0.37 diffusers v0.37 (PR #12987) optimizes all-ones attention masks to None in encode_prompt() when there is no padding. This breaks ai-toolkit's Qwen extensions which call .to() on the mask unconditionally. Fix: reconstruct the all-ones mask at the boundary (get_prompt_embeds) right after encode_prompt() returns. This keeps the rest of the code unchanged and works with both old and new diffusers versions. Also removes redundant duplicate mask assignments in qwen_image_edit and qwen_image_edit_plus. Fixes #740 --- .../diffusion_models/qwen_image/qwen_image.py | 16 ++++++++++------ .../qwen_image/qwen_image_edit.py | 8 ++++++-- .../qwen_image/qwen_image_edit_plus.py | 8 +++++--- 3 files changed, 21 insertions(+), 11 deletions(-) diff --git a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py index 99ecd6ad..0e1e0ee8 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py @@ -355,12 +355,16 @@ class QwenImageModel(BaseModel): if self.pipeline.text_encoder.device != self.device_torch: self.pipeline.text_encoder.to(self.device_torch) - max_sequence_length = 1024 - - prompt_embeds, prompt_embeds_mask = self.pipeline._get_qwen_prompt_embeds(prompt, self.device_torch) - prompt_embeds = prompt_embeds[:, :max_sequence_length] - prompt_embeds_mask = prompt_embeds_mask[:, :max_sequence_length] - + prompt_embeds, prompt_embeds_mask = self.pipeline.encode_prompt( + prompt, + device=self.device_torch, + num_images_per_prompt=1, + ) + # diffusers >=0.37 returns None when all tokens are valid (no padding) + if prompt_embeds_mask is None: + prompt_embeds_mask = torch.ones( + prompt_embeds.shape[:2], device=prompt_embeds.device, dtype=torch.int64 + ) pe = PromptEmbeds(prompt_embeds) pe.attention_mask = prompt_embeds_mask return pe diff --git a/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit.py b/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit.py index bcc8d735..724f9e7b 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit.py @@ -197,6 +197,11 @@ class QwenImageEditModel(QwenImageModel): device=self.device_torch, num_images_per_prompt=1, ) + # diffusers >=0.37 returns None when all tokens are valid (no padding) + if prompt_embeds_mask is None: + prompt_embeds_mask = torch.ones( + prompt_embeds.shape[:2], device=prompt_embeds.device, dtype=torch.int64 + ) pe = PromptEmbeds(prompt_embeds) pe.attention_mask = prompt_embeds_mask return pe @@ -251,14 +256,13 @@ class QwenImageEditModel(QwenImageModel): ) txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() enc_hs = text_embeddings.text_embeds.to(self.device_torch, self.torch_dtype) - prompt_embeds_mask = text_embeddings.attention_mask.to(self.device_torch, dtype=torch.int64) noise_pred = self.transformer( hidden_states=latent_model_input.to(self.device_torch, self.torch_dtype), timestep=timestep / 1000, guidance=None, encoder_hidden_states=enc_hs, - encoder_hidden_states_mask=prompt_embeds_mask, + encoder_hidden_states_mask=prompt_embeds_mask.detach(), img_shapes=img_shapes, txt_seq_lens=txt_seq_lens, return_dict=False, diff --git a/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit_plus.py b/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit_plus.py index 8272ee46..74a6e293 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit_plus.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit_plus.py @@ -192,6 +192,11 @@ class QwenImageEditPlusModel(QwenImageModel): device=self.device_torch, num_images_per_prompt=1, ) + # diffusers >=0.37 returns None when all tokens are valid (no padding) + if prompt_embeds_mask is None: + prompt_embeds_mask = torch.ones( + prompt_embeds.shape[:2], device=prompt_embeds.device, dtype=torch.int64 + ) pe = PromptEmbeds(prompt_embeds) pe.attention_mask = prompt_embeds_mask return pe @@ -323,9 +328,6 @@ class QwenImageEditPlusModel(QwenImageModel): ) txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() enc_hs = text_embeddings.text_embeds.to(self.device_torch, self.torch_dtype) - prompt_embeds_mask = text_embeddings.attention_mask.to( - self.device_torch, dtype=torch.int64 - ) noise_pred = self.transformer( hidden_states=latent_model_input.to(