From 6cc5dde5c208c4218ec58cd693468d815658160e Mon Sep 17 00:00:00 2001 From: Rayane <40967731+Rasaboun@users.noreply.github.com> Date: Mon, 23 Mar 2026 17:18:34 +0100 Subject: [PATCH 1/2] Fix Qwen Image mask handling --- .../diffusion_models/qwen_image/qwen_image.py | 10 +++++----- .../diffusion_models/qwen_image/qwen_image_edit.py | 11 +++++------ .../qwen_image/qwen_image_edit_plus.py | 13 +++++-------- toolkit/prompt_utils.py | 7 +++++++ 4 files changed, 22 insertions(+), 19 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 99ecd6ad2f..0a40912deb 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py @@ -280,11 +280,11 @@ def callback_on_step_end(pipe, i, t, callback_kwargs): gen_config.height = int(gen_config.height // sc * sc) img = pipeline( prompt_embeds=conditional_embeds.text_embeds, - prompt_embeds_mask=conditional_embeds.attention_mask.to( + prompt_embeds_mask=conditional_embeds.get_attention_mask( self.device_torch, dtype=torch.int64 ), negative_prompt_embeds=unconditional_embeds.text_embeds, - negative_prompt_embeds_mask=unconditional_embeds.attention_mask.to( + negative_prompt_embeds_mask=unconditional_embeds.get_attention_mask( self.device_torch, dtype=torch.int64 ), height=gen_config.height, @@ -324,10 +324,10 @@ def get_noise_prediction( img_shapes = [[(1, img_h2, img_w2)]] * batch_size enc_hs = text_embeddings.text_embeds.to(self.device_torch, self.torch_dtype) - prompt_embeds_mask = text_embeddings.attention_mask.to( + prompt_embeds_mask = text_embeddings.get_attention_mask( self.device_torch, dtype=torch.int64 ) - txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() + txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None noise_pred = self.transformer( hidden_states=latent_model_input.to( @@ -336,7 +336,7 @@ def get_noise_prediction( timestep=(timestep / 1000).detach(), guidance=None, encoder_hidden_states=enc_hs.detach(), - encoder_hidden_states_mask=prompt_embeds_mask.detach(), + encoder_hidden_states_mask=prompt_embeds_mask.detach() if prompt_embeds_mask is not None else None, 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.py b/extensions_built_in/diffusion_models/qwen_image/qwen_image_edit.py index bcc8d735c0..e34369a93e 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 @@ -118,11 +118,11 @@ def callback_on_step_end(pipe, i, t, callback_kwargs): img = pipeline( image=control_img, prompt_embeds=conditional_embeds.text_embeds, - prompt_embeds_mask=conditional_embeds.attention_mask.to( + prompt_embeds_mask=conditional_embeds.get_attention_mask( self.device_torch, dtype=torch.int64 ), negative_prompt_embeds=unconditional_embeds.text_embeds, - negative_prompt_embeds_mask=unconditional_embeds.attention_mask.to( + negative_prompt_embeds_mask=unconditional_embeds.get_attention_mask( self.device_torch, dtype=torch.int64 ), height=gen_config.height, @@ -246,19 +246,18 @@ def get_noise_prediction( latent_model_input = torch.cat([latent_model_input, control], dim=1) batch_size = latent_model_input.shape[0] - prompt_embeds_mask = text_embeddings.attention_mask.to( + prompt_embeds_mask = text_embeddings.get_attention_mask( self.device_torch, dtype=torch.int64 ) - txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() + txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None 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() if prompt_embeds_mask is not None else None, 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 8272ee464f..f0d0cec7be 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 @@ -136,11 +136,11 @@ def callback_on_step_end(pipe, i, t, callback_kwargs): img = pipeline( image=control_img_list, prompt_embeds=conditional_embeds.text_embeds, - prompt_embeds_mask=conditional_embeds.attention_mask.to( + prompt_embeds_mask=conditional_embeds.get_attention_mask( self.device_torch, dtype=torch.int64 ), negative_prompt_embeds=unconditional_embeds.text_embeds, - negative_prompt_embeds_mask=unconditional_embeds.attention_mask.to( + negative_prompt_embeds_mask=unconditional_embeds.get_attention_mask( self.device_torch, dtype=torch.int64 ), height=gen_config.height, @@ -318,14 +318,11 @@ def get_noise_prediction( latent_model_input = torch.cat(packed_latents_with_controls_list, dim=0) - prompt_embeds_mask = text_embeddings.attention_mask.to( + prompt_embeds_mask = text_embeddings.get_attention_mask( self.device_torch, dtype=torch.int64 ) - txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() + txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None 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( @@ -334,7 +331,7 @@ def get_noise_prediction( timestep=(timestep / 1000).detach(), guidance=None, encoder_hidden_states=enc_hs.detach(), - encoder_hidden_states_mask=prompt_embeds_mask.detach(), + encoder_hidden_states_mask=prompt_embeds_mask.detach() if prompt_embeds_mask is not None else None, img_shapes=img_shapes, txt_seq_lens=txt_seq_lens, return_dict=False, diff --git a/toolkit/prompt_utils.py b/toolkit/prompt_utils.py index ef3a7da811..05cd358470 100644 --- a/toolkit/prompt_utils.py +++ b/toolkit/prompt_utils.py @@ -35,6 +35,13 @@ def __init__(self, args: Union[Tuple[torch.Tensor], List[torch.Tensor], torch.Te self.attention_mask = attention_mask + def get_attention_mask(self, *args, **kwargs): + if self.attention_mask is None: + return None + if isinstance(self.attention_mask, list) or isinstance(self.attention_mask, tuple): + return [t.to(*args, **kwargs) for t in self.attention_mask] + return self.attention_mask.to(*args, **kwargs) + def to(self, *args, **kwargs): if isinstance(self.text_embeds, list) or isinstance(self.text_embeds, tuple): self.text_embeds = [t.to(*args, **kwargs) for t in self.text_embeds] From c9a99551dcd4ff1a3a7d150deee42d956a891709 Mon Sep 17 00:00:00 2001 From: Rayane <40967731+Rasaboun@users.noreply.github.com> Date: Mon, 23 Mar 2026 21:21:07 +0100 Subject: [PATCH 2/2] 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 | 26 +++++++++++-------- .../qwen_image/qwen_image_edit.py | 15 +++++++---- .../qwen_image/qwen_image_edit_plus.py | 15 +++++++---- toolkit/prompt_utils.py | 7 ----- 4 files changed, 35 insertions(+), 28 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 0a40912deb..0e1e0ee8f5 100644 --- a/extensions_built_in/diffusion_models/qwen_image/qwen_image.py +++ b/extensions_built_in/diffusion_models/qwen_image/qwen_image.py @@ -280,11 +280,11 @@ def callback_on_step_end(pipe, i, t, callback_kwargs): gen_config.height = int(gen_config.height // sc * sc) img = pipeline( prompt_embeds=conditional_embeds.text_embeds, - prompt_embeds_mask=conditional_embeds.get_attention_mask( + prompt_embeds_mask=conditional_embeds.attention_mask.to( self.device_torch, dtype=torch.int64 ), negative_prompt_embeds=unconditional_embeds.text_embeds, - negative_prompt_embeds_mask=unconditional_embeds.get_attention_mask( + negative_prompt_embeds_mask=unconditional_embeds.attention_mask.to( self.device_torch, dtype=torch.int64 ), height=gen_config.height, @@ -324,10 +324,10 @@ def get_noise_prediction( img_shapes = [[(1, img_h2, img_w2)]] * batch_size enc_hs = text_embeddings.text_embeds.to(self.device_torch, self.torch_dtype) - prompt_embeds_mask = text_embeddings.get_attention_mask( + prompt_embeds_mask = text_embeddings.attention_mask.to( self.device_torch, dtype=torch.int64 ) - txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None + txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() noise_pred = self.transformer( hidden_states=latent_model_input.to( @@ -336,7 +336,7 @@ def get_noise_prediction( timestep=(timestep / 1000).detach(), guidance=None, encoder_hidden_states=enc_hs.detach(), - encoder_hidden_states_mask=prompt_embeds_mask.detach() if prompt_embeds_mask is not None else None, + encoder_hidden_states_mask=prompt_embeds_mask.detach(), img_shapes=img_shapes, txt_seq_lens=txt_seq_lens, return_dict=False, @@ -355,12 +355,16 @@ def get_prompt_embeds(self, prompt: str) -> PromptEmbeds: 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 e34369a93e..724f9e7bb8 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 @@ -118,11 +118,11 @@ def callback_on_step_end(pipe, i, t, callback_kwargs): img = pipeline( image=control_img, prompt_embeds=conditional_embeds.text_embeds, - prompt_embeds_mask=conditional_embeds.get_attention_mask( + prompt_embeds_mask=conditional_embeds.attention_mask.to( self.device_torch, dtype=torch.int64 ), negative_prompt_embeds=unconditional_embeds.text_embeds, - negative_prompt_embeds_mask=unconditional_embeds.get_attention_mask( + negative_prompt_embeds_mask=unconditional_embeds.attention_mask.to( self.device_torch, dtype=torch.int64 ), height=gen_config.height, @@ -197,6 +197,11 @@ def get_prompt_embeds(self, prompt: str, control_images=None) -> PromptEmbeds: 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 @@ -246,10 +251,10 @@ def get_noise_prediction( latent_model_input = torch.cat([latent_model_input, control], dim=1) batch_size = latent_model_input.shape[0] - prompt_embeds_mask = text_embeddings.get_attention_mask( + prompt_embeds_mask = text_embeddings.attention_mask.to( self.device_torch, dtype=torch.int64 ) - txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None + txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() enc_hs = text_embeddings.text_embeds.to(self.device_torch, self.torch_dtype) noise_pred = self.transformer( @@ -257,7 +262,7 @@ def get_noise_prediction( timestep=timestep / 1000, guidance=None, encoder_hidden_states=enc_hs, - encoder_hidden_states_mask=prompt_embeds_mask.detach() if prompt_embeds_mask is not None else None, + 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 f0d0cec7be..74a6e293a7 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 @@ -136,11 +136,11 @@ def callback_on_step_end(pipe, i, t, callback_kwargs): img = pipeline( image=control_img_list, prompt_embeds=conditional_embeds.text_embeds, - prompt_embeds_mask=conditional_embeds.get_attention_mask( + prompt_embeds_mask=conditional_embeds.attention_mask.to( self.device_torch, dtype=torch.int64 ), negative_prompt_embeds=unconditional_embeds.text_embeds, - negative_prompt_embeds_mask=unconditional_embeds.get_attention_mask( + negative_prompt_embeds_mask=unconditional_embeds.attention_mask.to( self.device_torch, dtype=torch.int64 ), height=gen_config.height, @@ -192,6 +192,11 @@ def get_prompt_embeds(self, prompt: str, control_images=None) -> PromptEmbeds: 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 @@ -318,10 +323,10 @@ def get_noise_prediction( latent_model_input = torch.cat(packed_latents_with_controls_list, dim=0) - prompt_embeds_mask = text_embeddings.get_attention_mask( + prompt_embeds_mask = text_embeddings.attention_mask.to( self.device_torch, dtype=torch.int64 ) - txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None + txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() enc_hs = text_embeddings.text_embeds.to(self.device_torch, self.torch_dtype) noise_pred = self.transformer( @@ -331,7 +336,7 @@ def get_noise_prediction( timestep=(timestep / 1000).detach(), guidance=None, encoder_hidden_states=enc_hs.detach(), - encoder_hidden_states_mask=prompt_embeds_mask.detach() if prompt_embeds_mask is not None else None, + encoder_hidden_states_mask=prompt_embeds_mask.detach(), img_shapes=img_shapes, txt_seq_lens=txt_seq_lens, return_dict=False, diff --git a/toolkit/prompt_utils.py b/toolkit/prompt_utils.py index 05cd358470..ef3a7da811 100644 --- a/toolkit/prompt_utils.py +++ b/toolkit/prompt_utils.py @@ -35,13 +35,6 @@ def __init__(self, args: Union[Tuple[torch.Tensor], List[torch.Tensor], torch.Te self.attention_mask = attention_mask - def get_attention_mask(self, *args, **kwargs): - if self.attention_mask is None: - return None - if isinstance(self.attention_mask, list) or isinstance(self.attention_mask, tuple): - return [t.to(*args, **kwargs) for t in self.attention_mask] - return self.attention_mask.to(*args, **kwargs) - def to(self, *args, **kwargs): if isinstance(self.text_embeds, list) or isinstance(self.text_embeds, tuple): self.text_embeds = [t.to(*args, **kwargs) for t in self.text_embeds]