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(