| import inspect |
| import math |
| from typing import Any, Callable, Dict, List, Optional, Union |
|
|
| import numpy as np |
| import torch |
| from transformers import Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer, Qwen2VLProcessor |
|
|
| from diffusers.image_processor import PipelineImageInput, VaeImageProcessor |
| from diffusers.loaders import QwenImageLoraLoaderMixin |
| from diffusers.models import AutoencoderKLQwenImage, QwenImageTransformer2DModel |
| from diffusers.schedulers import FlowMatchEulerDiscreteScheduler |
| from diffusers.utils import is_torch_xla_available, logging, replace_example_docstring |
| from diffusers.utils.torch_utils import randn_tensor |
| from diffusers.pipelines.pipeline_utils import DiffusionPipeline |
| from diffusers.pipelines.qwenimage.pipeline_output import QwenImagePipelineOutput |
|
|
| if is_torch_xla_available(): |
| import torch_xla.core.xla_model as xm |
| XLA_AVAILABLE = True |
| else: |
| XLA_AVAILABLE = False |
|
|
| logger = logging.get_logger(__name__) |
|
|
| CONDITION_IMAGE_SIZE = 384 * 384 |
| VAE_IMAGE_SIZE = 1024 * 1024 |
|
|
|
|
| def calculate_shift(image_seq_len, base_seq_len=256, max_seq_len=4096, base_shift=0.5, max_shift=1.15): |
| m = (max_shift - base_shift) / (max_seq_len - base_seq_len) |
| b = base_shift - m * base_seq_len |
| return image_seq_len * m + b |
|
|
|
|
| def retrieve_timesteps(scheduler, num_inference_steps=None, device=None, timesteps=None, sigmas=None, **kwargs): |
| if timesteps is not None and sigmas is not None: |
| raise ValueError("Only one of `timesteps` or `sigmas` can be passed.") |
| if timesteps is not None: |
| accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) |
| if not accepts_timesteps: |
| raise ValueError(f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom timesteps.") |
| scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) |
| timesteps = scheduler.timesteps |
| num_inference_steps = len(timesteps) |
| elif sigmas is not None: |
| accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) |
| if not accept_sigmas: |
| raise ValueError(f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom sigmas.") |
| scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs) |
| timesteps = scheduler.timesteps |
| num_inference_steps = len(timesteps) |
| else: |
| scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) |
| timesteps = scheduler.timesteps |
| return timesteps, num_inference_steps |
|
|
|
|
| def retrieve_latents(encoder_output, generator=None, sample_mode="sample"): |
| if hasattr(encoder_output, "latent_dist") and sample_mode == "sample": |
| return encoder_output.latent_dist.sample(generator) |
| elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax": |
| return encoder_output.latent_dist.mode() |
| elif hasattr(encoder_output, "latents"): |
| return encoder_output.latents |
| else: |
| raise AttributeError("Could not access latents of provided encoder_output") |
|
|
|
|
| def calculate_dimensions(target_area, ratio): |
| width = math.sqrt(target_area * ratio) |
| height = width / ratio |
| width = round(width / 32) * 32 |
| height = round(height / 32) * 32 |
| return width, height |
|
|
|
|
| class QwenImageEditPlusPipeline(DiffusionPipeline, QwenImageLoraLoaderMixin): |
| model_cpu_offload_seq = "text_encoder->transformer->vae" |
| _callback_tensor_inputs = ["latents", "prompt_embeds"] |
|
|
| def __init__(self, scheduler, vae, text_encoder, tokenizer, processor, transformer): |
| super().__init__() |
| self.register_modules( |
| vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, |
| processor=processor, transformer=transformer, scheduler=scheduler, |
| ) |
| self.vae_scale_factor = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8 |
| self.latent_channels = self.vae.config.z_dim if getattr(self, "vae", None) else 16 |
| self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor * 2) |
| self.tokenizer_max_length = 1024 |
| self.prompt_template_encode = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" |
| self.prompt_template_encode_start_idx = 64 |
| self.default_sample_size = 128 |
|
|
| def _extract_masked_hidden(self, hidden_states, mask): |
| bool_mask = mask.bool() |
| valid_lengths = bool_mask.sum(dim=1) |
| selected = hidden_states[bool_mask] |
| return torch.split(selected, valid_lengths.tolist(), dim=0) |
|
|
| def _get_qwen_prompt_embeds(self, prompt=None, image=None, device=None, dtype=None): |
| device = device or self._execution_device |
| dtype = dtype or self.text_encoder.dtype |
| prompt = [prompt] if isinstance(prompt, str) else prompt |
| img_prompt_template = "Picture {}: <|vision_start|><|image_pad|><|vision_end|>" |
| if isinstance(image, list): |
| base_img_prompt = "" |
| for i, img in enumerate(image): |
| base_img_prompt += img_prompt_template.format(i + 1) |
| elif image is not None: |
| base_img_prompt = img_prompt_template.format(1) |
| else: |
| base_img_prompt = "" |
| template = self.prompt_template_encode |
| drop_idx = self.prompt_template_encode_start_idx |
| txt = [template.format(base_img_prompt + e) for e in prompt] |
| model_inputs = self.processor(text=txt, images=image, padding=True, return_tensors="pt").to(device) |
| outputs = self.text_encoder( |
| input_ids=model_inputs.input_ids, attention_mask=model_inputs.attention_mask, |
| pixel_values=model_inputs.pixel_values, image_grid_thw=model_inputs.image_grid_thw, |
| output_hidden_states=True, |
| ) |
| hidden_states = outputs.hidden_states[-1] |
| split_hidden_states = self._extract_masked_hidden(hidden_states, model_inputs.attention_mask) |
| split_hidden_states = [e[drop_idx:] for e in split_hidden_states] |
| attn_mask_list = [torch.ones(e.size(0), dtype=torch.long, device=e.device) for e in split_hidden_states] |
| max_seq_len = max([e.size(0) for e in split_hidden_states]) |
| prompt_embeds = torch.stack( |
| [torch.cat([u, u.new_zeros(max_seq_len - u.size(0), u.size(1))]) for u in split_hidden_states] |
| ) |
| encoder_attention_mask = torch.stack( |
| [torch.cat([u, u.new_zeros(max_seq_len - u.size(0))]) for u in attn_mask_list] |
| ) |
| prompt_embeds = prompt_embeds.to(dtype=dtype, device=device) |
| return prompt_embeds, encoder_attention_mask |
|
|
| def encode_prompt(self, prompt, image=None, device=None, num_images_per_prompt=1, |
| prompt_embeds=None, prompt_embeds_mask=None, max_sequence_length=1024): |
| device = device or self._execution_device |
| prompt = [prompt] if isinstance(prompt, str) else prompt |
| batch_size = len(prompt) if prompt_embeds is None else prompt_embeds.shape[0] |
| if prompt_embeds is None: |
| prompt_embeds, prompt_embeds_mask = self._get_qwen_prompt_embeds(prompt, image, device) |
| _, seq_len, _ = prompt_embeds.shape |
| prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) |
| prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1) |
| prompt_embeds_mask = prompt_embeds_mask.repeat(1, num_images_per_prompt, 1) |
| prompt_embeds_mask = prompt_embeds_mask.view(batch_size * num_images_per_prompt, seq_len) |
| return prompt_embeds, prompt_embeds_mask |
|
|
| def check_inputs(self, prompt, height, width, negative_prompt=None, prompt_embeds=None, |
| negative_prompt_embeds=None, prompt_embeds_mask=None, negative_prompt_embeds_mask=None, |
| callback_on_step_end_tensor_inputs=None, max_sequence_length=None): |
| if height % (self.vae_scale_factor * 2) != 0 or width % (self.vae_scale_factor * 2) != 0: |
| logger.warning(f"`height` and `width` have to be divisible by {self.vae_scale_factor * 2} but are {height} and {width}. Dimensions will be resized accordingly") |
| if callback_on_step_end_tensor_inputs is not None and not all(k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs): |
| raise ValueError(f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}") |
| if prompt is not None and prompt_embeds is not None: |
| raise ValueError(f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}.") |
| elif prompt is None and prompt_embeds is None: |
| raise ValueError("Provide either `prompt` or `prompt_embeds`. Cannot leave both undefined.") |
| elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): |
| raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") |
| if negative_prompt is not None and negative_prompt_embeds is not None: |
| raise ValueError(f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}.") |
| if prompt_embeds is not None and prompt_embeds_mask is None: |
| raise ValueError("If `prompt_embeds` are provided, `prompt_embeds_mask` also have to be passed.") |
| if negative_prompt_embeds is not None and negative_prompt_embeds_mask is None: |
| raise ValueError("If `negative_prompt_embeds` are provided, `negative_prompt_embeds_mask` also have to be passed.") |
| if max_sequence_length is not None and max_sequence_length > 1024: |
| raise ValueError(f"`max_sequence_length` cannot be greater than 1024 but is {max_sequence_length}") |
|
|
| @staticmethod |
| def _pack_latents(latents, batch_size, num_channels_latents, height, width): |
| latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2) |
| latents = latents.permute(0, 2, 4, 1, 3, 5) |
| latents = latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4) |
| return latents |
|
|
| @staticmethod |
| def _unpack_latents(latents, height, width, vae_scale_factor): |
| batch_size, num_patches, channels = latents.shape |
| height = 2 * (int(height) // (vae_scale_factor * 2)) |
| width = 2 * (int(width) // (vae_scale_factor * 2)) |
| latents = latents.view(batch_size, height // 2, width // 2, channels // 4, 2, 2) |
| latents = latents.permute(0, 3, 1, 4, 2, 5) |
| latents = latents.reshape(batch_size, channels // (2 * 2), 1, height, width) |
| return latents |
|
|
| def _encode_vae_image(self, image, generator): |
| if isinstance(generator, list): |
| image_latents = [ |
| retrieve_latents(self.vae.encode(image[i:i+1]), generator=generator[i], sample_mode="argmax") |
| for i in range(image.shape[0]) |
| ] |
| image_latents = torch.cat(image_latents, dim=0) |
| else: |
| image_latents = retrieve_latents(self.vae.encode(image), generator=generator, sample_mode="argmax") |
| latents_mean = torch.tensor(self.vae.config.latents_mean).view(1, self.latent_channels, 1, 1, 1).to(image_latents.device, image_latents.dtype) |
| latents_std = torch.tensor(self.vae.config.latents_std).view(1, self.latent_channels, 1, 1, 1).to(image_latents.device, image_latents.dtype) |
| return (image_latents - latents_mean) / latents_std |
|
|
| def prepare_latents(self, images, batch_size, num_channels_latents, height, width, dtype, device, generator, latents=None): |
| height = 2 * (int(height) // (self.vae_scale_factor * 2)) |
| width = 2 * (int(width) // (self.vae_scale_factor * 2)) |
| shape = (batch_size, 1, num_channels_latents, height, width) |
| image_latents = None |
| if images is not None: |
| if not isinstance(images, list): |
| images = [images] |
| all_image_latents = [] |
| for image in images: |
| image = image.to(device=device, dtype=dtype) |
| if image.shape[1] != self.latent_channels: |
| image_latents = self._encode_vae_image(image=image, generator=generator) |
| else: |
| image_latents = image |
| if batch_size > image_latents.shape[0] and batch_size % image_latents.shape[0] == 0: |
| additional_image_per_prompt = batch_size // image_latents.shape[0] |
| image_latents = torch.cat([image_latents] * additional_image_per_prompt, dim=0) |
| elif batch_size > image_latents.shape[0] and batch_size % image_latents.shape[0] != 0: |
| raise ValueError(f"Cannot duplicate `image` of batch size {image_latents.shape[0]} to {batch_size} text prompts.") |
| else: |
| image_latents = torch.cat([image_latents], dim=0) |
| image_latent_height, image_latent_width = image_latents.shape[3:] |
| image_latents = self._pack_latents(image_latents, batch_size, num_channels_latents, image_latent_height, image_latent_width) |
| all_image_latents.append(image_latents) |
| image_latents = torch.cat(all_image_latents, dim=1) |
| if isinstance(generator, list) and len(generator) != batch_size: |
| raise ValueError(f"You have passed a list of generators of length {len(generator)}, but requested an effective batch size of {batch_size}.") |
| if latents is None: |
| latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) |
| latents = self._pack_latents(latents, batch_size, num_channels_latents, height, width) |
| else: |
| latents = latents.to(device=device, dtype=dtype) |
| return latents, image_latents |
|
|
| @property |
| def guidance_scale(self): |
| return self._guidance_scale |
|
|
| @property |
| def attention_kwargs(self): |
| return self._attention_kwargs |
|
|
| @property |
| def num_timesteps(self): |
| return self._num_timesteps |
|
|
| @property |
| def current_timestep(self): |
| return self._current_timestep |
|
|
| @property |
| def interrupt(self): |
| return self._interrupt |
|
|
| @torch.no_grad() |
| def __call__(self, image=None, prompt=None, negative_prompt=None, true_cfg_scale=4.0, |
| height=None, width=None, num_inference_steps=50, sigmas=None, |
| guidance_scale=None, num_images_per_prompt=1, generator=None, latents=None, |
| prompt_embeds=None, prompt_embeds_mask=None, negative_prompt_embeds=None, |
| negative_prompt_embeds_mask=None, output_type="pil", return_dict=True, |
| attention_kwargs=None, callback_on_step_end=None, |
| callback_on_step_end_tensor_inputs=["latents"], max_sequence_length=512): |
| image_size = image[-1].size if isinstance(image, list) else image.size |
| calculated_width, calculated_height = calculate_dimensions(1024 * 1024, image_size[0] / image_size[1]) |
| height = height or calculated_height |
| width = width or calculated_width |
| multiple_of = self.vae_scale_factor * 2 |
| width = width // multiple_of * multiple_of |
| height = height // multiple_of * multiple_of |
| self.check_inputs(prompt, height, width, negative_prompt=negative_prompt, |
| prompt_embeds=prompt_embeds, negative_prompt_embeds=negative_prompt_embeds, |
| prompt_embeds_mask=prompt_embeds_mask, negative_prompt_embeds_mask=negative_prompt_embeds_mask, |
| callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs, |
| max_sequence_length=max_sequence_length) |
| self._guidance_scale = guidance_scale |
| self._attention_kwargs = attention_kwargs |
| self._current_timestep = None |
| self._interrupt = False |
| if prompt is not None and isinstance(prompt, str): |
| batch_size = 1 |
| elif prompt is not None and isinstance(prompt, list): |
| batch_size = len(prompt) |
| else: |
| batch_size = prompt_embeds.shape[0] |
| device = self._execution_device |
|
|
| if image is not None and not (isinstance(image, torch.Tensor) and image.size(1) == self.latent_channels): |
| if not isinstance(image, list): |
| image = [image] |
| condition_image_sizes = [] |
| condition_images = [] |
| vae_image_sizes = [] |
| vae_images = [] |
| for img in image: |
| image_width, image_height = img.size |
| condition_width, condition_height = calculate_dimensions(CONDITION_IMAGE_SIZE, image_width / image_height) |
| vae_width, vae_height = calculate_dimensions(VAE_IMAGE_SIZE, image_width / image_height) |
| condition_image_sizes.append((condition_width, condition_height)) |
| vae_image_sizes.append((vae_width, vae_height)) |
| condition_images.append(self.image_processor.resize(img, condition_height, condition_width)) |
| vae_images.append(self.image_processor.preprocess(img, vae_height, vae_width).unsqueeze(2)) |
|
|
| has_neg_prompt = negative_prompt is not None or ( |
| negative_prompt_embeds is not None and negative_prompt_embeds_mask is not None |
| ) |
| do_true_cfg = true_cfg_scale > 1 and has_neg_prompt |
| if true_cfg_scale > 1 and not has_neg_prompt: |
| logger.warning(f"true_cfg_scale is passed as {true_cfg_scale}, but classifier-free guidance is not enabled.") |
| elif true_cfg_scale <= 1 and has_neg_prompt: |
| logger.warning("negative_prompt is passed but classifier-free guidance is not enabled since true_cfg_scale <= 1") |
|
|
| prompt_embeds, prompt_embeds_mask = self.encode_prompt( |
| image=condition_images, prompt=prompt, prompt_embeds=prompt_embeds, |
| prompt_embeds_mask=prompt_embeds_mask, device=device, |
| num_images_per_prompt=num_images_per_prompt, max_sequence_length=max_sequence_length, |
| ) |
| if do_true_cfg: |
| negative_prompt_embeds, negative_prompt_embeds_mask = self.encode_prompt( |
| image=condition_images, prompt=negative_prompt, prompt_embeds=negative_prompt_embeds, |
| prompt_embeds_mask=negative_prompt_embeds_mask, device=device, |
| num_images_per_prompt=num_images_per_prompt, max_sequence_length=max_sequence_length, |
| ) |
|
|
| num_channels_latents = self.transformer.config.in_channels // 4 |
| latents, image_latents = self.prepare_latents( |
| vae_images, batch_size * num_images_per_prompt, num_channels_latents, |
| height, width, prompt_embeds.dtype, device, generator, latents, |
| ) |
| img_shapes = [ |
| [(1, height // self.vae_scale_factor // 2, width // self.vae_scale_factor // 2), |
| *[(1, vae_height // self.vae_scale_factor // 2, vae_width // self.vae_scale_factor // 2) |
| for vae_width, vae_height in vae_image_sizes]] |
| ] * batch_size |
|
|
| sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas |
| image_seq_len = latents.shape[1] |
| mu = calculate_shift(image_seq_len, self.scheduler.config.get("base_image_seq_len", 256), |
| self.scheduler.config.get("max_image_seq_len", 4096), |
| self.scheduler.config.get("base_shift", 0.5), |
| self.scheduler.config.get("max_shift", 1.15)) |
| timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, sigmas=sigmas, mu=mu) |
| num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) |
| self._num_timesteps = len(timesteps) |
|
|
| if self.transformer.config.guidance_embeds and guidance_scale is None: |
| raise ValueError("guidance_scale is required for guidance-distilled model.") |
| elif self.transformer.config.guidance_embeds: |
| guidance = torch.full([1], guidance_scale, device=device, dtype=torch.float32).expand(latents.shape[0]) |
| elif not self.transformer.config.guidance_embeds and guidance_scale is not None: |
| guidance = None |
| else: |
| guidance = None |
|
|
| if self.attention_kwargs is None: |
| self._attention_kwargs = {} |
|
|
| txt_seq_lens = prompt_embeds_mask.sum(dim=1).tolist() if prompt_embeds_mask is not None else None |
| image_rotary_emb = self.transformer.pos_embed(img_shapes, txt_seq_lens, device=latents.device) |
| if do_true_cfg: |
| negative_txt_seq_lens = negative_prompt_embeds_mask.sum(dim=1).tolist() if negative_prompt_embeds_mask is not None else None |
| uncond_image_rotary_emb = self.transformer.pos_embed(img_shapes, negative_txt_seq_lens, device=latents.device) |
| else: |
| uncond_image_rotary_emb = None |
|
|
| self.scheduler.set_begin_index(0) |
| with self.progress_bar(total=num_inference_steps) as progress_bar: |
| for i, t in enumerate(timesteps): |
| if self.interrupt: |
| continue |
| self._current_timestep = t |
| latent_model_input = latents |
| if image_latents is not None: |
| latent_model_input = torch.cat([latents, image_latents], dim=1) |
| timestep = t.expand(latents.shape[0]).to(latents.dtype) |
| with self.transformer.cache_context("cond"): |
| noise_pred = self.transformer( |
| hidden_states=latent_model_input, timestep=timestep / 1000, |
| guidance=guidance, encoder_hidden_states_mask=prompt_embeds_mask, |
| encoder_hidden_states=prompt_embeds, image_rotary_emb=image_rotary_emb, |
| attention_kwargs=self.attention_kwargs, return_dict=False, |
| )[0] |
| noise_pred = noise_pred[:, :latents.size(1)] |
| if do_true_cfg: |
| with self.transformer.cache_context("uncond"): |
| neg_noise_pred = self.transformer( |
| hidden_states=latent_model_input, timestep=timestep / 1000, |
| guidance=guidance, encoder_hidden_states_mask=negative_prompt_embeds_mask, |
| encoder_hidden_states=negative_prompt_embeds, image_rotary_emb=uncond_image_rotary_emb, |
| attention_kwargs=self.attention_kwargs, return_dict=False, |
| )[0] |
| neg_noise_pred = neg_noise_pred[:, :latents.size(1)] |
| comb_pred = neg_noise_pred + true_cfg_scale * (noise_pred - neg_noise_pred) |
| noise_pred = comb_pred * (torch.norm(noise_pred, dim=-1, keepdim=True) / torch.norm(comb_pred, dim=-1, keepdim=True)) |
| latents_dtype = latents.dtype |
| latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] |
| if latents.dtype != latents_dtype: |
| if torch.backends.mps.is_available(): |
| latents = latents.to(latents_dtype) |
| if callback_on_step_end is not None: |
| callback_kwargs = {} |
| for k in callback_on_step_end_tensor_inputs: |
| callback_kwargs[k] = locals()[k] |
| callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) |
| latents = callback_outputs.pop("latents", latents) |
| prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) |
| if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): |
| progress_bar.update() |
| if XLA_AVAILABLE: |
| xm.mark_step() |
|
|
| self._current_timestep = None |
| if output_type == "latent": |
| image = latents |
| else: |
| latents = self._unpack_latents(latents, height, width, self.vae_scale_factor) |
| latents = latents.to(self.vae.dtype) |
| latents_mean = torch.tensor(self.vae.config.latents_mean).view(1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype) |
| latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(latents.device, latents.dtype) |
| latents = latents / latents_std + latents_mean |
| image = self.vae.decode(latents, return_dict=False)[0][:, :, 0] |
| image = self.image_processor.postprocess(image, output_type=output_type) |
| self.maybe_free_model_hooks() |
| if not return_dict: |
| return (image,) |
| return QwenImagePipelineOutput(images=image) |
|
|