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
This commit is contained in:
Rayane 2026-03-23 21:43:08 +01:00 committed by GitHub
parent 295094b4b5
commit 99a4a5887b
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 21 additions and 11 deletions

View File

@ -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

View File

@ -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,

View File

@ -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(