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:
parent
295094b4b5
commit
99a4a5887b
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Reference in New Issue