import inspect from typing import Callable, Dict, List, Optional, Union import numpy as np import torch from PIL import Image from transformers import ( Qwen2Tokenizer, Qwen3VLForConditionalGeneration, Qwen3VLProcessor, ) from ...callbacks import MultiPipelineCallbacks, PipelineCallback from ...image_processor import PipelineImageInput, VaeImageProcessor from ...models import AutoencoderKLWan, JoyImageEditTransformer3DModel from ...schedulers import FlowMatchEulerDiscreteScheduler from ...utils import replace_example_docstring from ...utils.torch_utils import randn_tensor from ..pipeline_utils import DiffusionPipeline from .image_processor import JoyImageEditImageProcessor from .pipeline_output import JoyImageEditPipelineOutput EXAMPLE_DOC_STRING = """ Examples: ```python >>> import torch >>> from diffusers import JoyImageEditPipeline >>> from diffusers.utils import load_image >>> model_id = "jdopensource/JoyAI-Image-Edit-Diffusers" >>> pipe = JoyImageEditPipeline.from_pretrained(model_id, torch_dtype=torch.bfloat16) >>> pipe.to("cuda") >>> image = load_image("https://huggingface.co/datasets/diffusers/docs-images/resolve/main/astronaut.jpg") >>> output = pipe( ... image=image, # pass an image for editing; omit for text-to-image generation ... prompt="Add wings to the astronaut.", ... num_inference_steps=40, ... guidance_scale=4.0, ... generator=torch.manual_seed(0), ... ) >>> output.images[0].save("joyimage_edit.png") ``` """ def retrieve_timesteps( scheduler, num_inference_steps: Optional[int] = None, device: Optional[Union[str, torch.device]] = None, timesteps: Optional[List[int]] = None, sigmas: Optional[List[float]] = None, **kwargs, ): """ Configure the scheduler and return its timestep sequence. Exactly one of ``timesteps``, ``sigmas``, or ``num_inference_steps`` should be provided to control the denoising schedule. Args: scheduler: The diffusion scheduler. num_inference_steps: Number of denoising steps (used when neither ``timesteps`` nor ``sigmas`` is given). device: Target device for the timestep tensor. timesteps: Custom discrete timesteps. sigmas: Custom sigma values (alternative to ``timesteps``). **kwargs: Additional keyword arguments forwarded to ``set_timesteps``. Returns: Tuple of (timesteps tensor, num_inference_steps int). Raises: ValueError: If both ``timesteps`` and ``sigmas`` are provided, or if the scheduler does not support the requested schedule parameterisation. """ 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: if "timesteps" not in set(inspect.signature(scheduler.set_timesteps).parameters.keys()): raise ValueError(f"{scheduler.__class__} 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: if "sigmas" not in set(inspect.signature(scheduler.set_timesteps).parameters.keys()): raise ValueError(f"{scheduler.__class__} 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 class JoyImageEditPipeline(DiffusionPipeline): """ Diffusion pipeline for image editing using the JoyImage architecture. The pipeline encodes text and image conditioning via a Qwen3-VL text encoder, denoises latents with a 3-D transformer, and decodes the result with a WAN VAE. Model offloading order: text_encoder -> transformer -> vae. """ model_cpu_offload_seq = "text_encoder->transformer->vae" _callback_tensor_inputs = ["latents", "prompt_embeds"] def __init__( self, scheduler: FlowMatchEulerDiscreteScheduler, vae: AutoencoderKLWan, text_encoder: Qwen3VLForConditionalGeneration, tokenizer: Qwen2Tokenizer, transformer: JoyImageEditTransformer3DModel, processor: Qwen3VLProcessor, text_token_max_length: int = 2048, ): """ Initialise the pipeline and register all sub-modules. Args: scheduler: Noise scheduler for the denoising process. vae: Variational autoencoder used for encoding / decoding latents. text_encoder: Qwen3-VL multimodal language model for prompt encoding. tokenizer: Tokenizer paired with the text encoder. transformer: 3-D transformer denoising network. processor: Qwen3-VL processor for multi-image prompt preparation. text_token_max_length: Maximum number of text tokens for the encoder. """ super().__init__() self.register_modules( vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, transformer=transformer, scheduler=scheduler, processor=processor, ) self.text_token_max_length = text_token_max_length self.vae_scale_factor_temporal = self.vae.config.scale_factor_temporal if getattr(self, "vae", None) else 4 self.vae_scale_factor_spatial = self.vae.config.scale_factor_spatial if getattr(self, "vae", None) else 8 self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor_spatial) self.vae_image_processor = JoyImageEditImageProcessor( vae_scale_factor=self.vae_scale_factor_spatial, ) # Prompt templates used when encoding text with / without image tokens. self.prompt_template_encode = { "image": ( "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, " "quantity, text, spatial relationships of the objects and background:<|im_end|>\n" "<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" ), "multiple_images": ( "<|im_start|>system\n \\nDescribe the image by detailing the color, shape, size, texture, " "quantity, text, spatial relationships of the objects and background:<|im_end|>\n" "{}<|im_start|>assistant\n" ), } # Number of system-prompt tokens to drop from the beginning of hidden states. self.prompt_template_encode_start_idx = { "image": 34, "multiple_images": 34, } # ------------------------------------------------------------------ # Internal helpers # ------------------------------------------------------------------ def _get_last_decoder_hidden_states(self, forward_fn, **kwargs): """ Run ``forward_fn(**kwargs)`` while capturing the **pre-norm** output of the last decoder layer via a forward hook. This model was trained on transformers 4.57, where ``Qwen3VLForConditionalGeneration``'s ``@check_model_inputs`` decorator monkey-patched each decoder layer to collect ``hidden_states``. Because ``Qwen3VLCausalLMOutputWithPast`` has no ``last_hidden_state`` field, ``tie_last_hidden_states`` had no effect and ``hidden_states[-1]`` was the **pre-norm** output of the last decoder layer. Starting from https://github.com/huggingface/transformers/pull/42609 the CausalLM forward explicitly returns ``hidden_states=outputs.hidden_states`` from the inner model. Combined with the subsequent ``@check_model_inputs`` → ``@capture_outputs`` migration (transformers 5.x), ``hidden_states`` is now captured at the ``Qwen3VLTextModel`` level where ``tie_last_hidden_states=True`` replaces ``hidden_states[-1]`` with the **post-norm** ``last_hidden_state``. The CausalLM simply passes this through, so ``hidden_states[-1]`` becomes post-norm – a ~10× scale difference (std ≈ 2 vs ≈ 21) that breaks inference. This helper bypasses both mechanisms by hooking the last decoder layer directly, returning the raw pre-norm output regardless of the transformers version. """ captured = {} def _hook(_module, _input, output): captured["hidden_states"] = output[0] if isinstance(output, tuple) else output handle = self.text_encoder.model.language_model.layers[-1].register_forward_hook(_hook) try: forward_fn(**kwargs) finally: handle.remove() return captured["hidden_states"] def _extract_masked_hidden(self, hidden_states: torch.Tensor, mask: torch.Tensor) -> tuple[torch.Tensor, ...]: """ Extract valid (non-padded) hidden states for each sequence in the batch. Args: hidden_states: Shape (B, T, D). mask: Binary attention mask of shape (B, T). Returns: Tuple of tensors, one per batch element, each of shape (valid_T, D). """ 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: Union[str, List[str]] = None, template_type: str = "image", device: Optional[torch.device] = None, dtype: Optional[torch.dtype] = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Encode text prompts using the Qwen tokenizer (text-only path). Args: prompt: A single prompt string or a list of prompt strings. template_type: Key into ``prompt_template_encode`` / ``prompt_template_encode_start_idx``. device: Target device. dtype: Target floating-point dtype. Returns: Tuple of (prompt_embeds, encoder_attention_mask) where both tensors have shape (B, max_seq_len, D) and (B, max_seq_len) respectively, zero-padded to the same length. """ device = device or self._execution_device dtype = dtype or self.text_encoder.dtype prompt = [prompt] if isinstance(prompt, str) else prompt template = self.prompt_template_encode[template_type] drop_idx = self.prompt_template_encode_start_idx[template_type] txt = [template.format(e) for e in prompt] txt_tokens = self.tokenizer( txt, max_length=self.text_token_max_length + drop_idx, padding=True, truncation=True, return_tensors="pt", ).to(device) hidden_states = self._get_last_decoder_hidden_states( self.text_encoder, input_ids=txt_tokens.input_ids, attention_mask=txt_tokens.attention_mask, ) # Drop system-prompt prefix tokens and re-pack into a padded batch. split_hidden_states = self._extract_masked_hidden(hidden_states, txt_tokens.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 = min( self.text_token_max_length, max(u.size(0) for u in split_hidden_states), max(u.size(0) for u in attn_mask_list), ) 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_multiple_images( self, prompt: Union[str, List[str]], device: Optional[torch.device] = None, num_images_per_prompt: int = 1, images: Optional[torch.Tensor] = None, prompt_embeds: Optional[torch.Tensor] = None, prompt_embeds_mask: Optional[torch.Tensor] = None, template_type: Optional[str] = "multiple_images", max_sequence_length: Optional[int] = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Encode prompts that contain inline image tokens via the Qwen processor. ``\\n`` placeholders in each prompt string are replaced by the Qwen vision special tokens before being fed to the multimodal encoder. Args: prompt: Prompt string(s), optionally containing ``\\n`` tokens. device: Target device. num_images_per_prompt: Number of outputs to generate per prompt. images: Pixel tensors corresponding to the inline image tokens. prompt_embeds: Pre-computed prompt embeddings. prompt_embeds_mask: Attention mask for pre-computed embeddings. template_type: Must be ``"multiple_images"``. max_sequence_length: If set, truncate the output to this length (keeping the last ``max_sequence_length`` tokens). Returns: Tuple of (prompt_embeds, prompt_embeds_mask). """ if template_type != "multiple_images": raise ValueError(f"Expected template_type 'multiple_images', but got '{template_type}'") 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: template = self.prompt_template_encode[template_type] drop_idx = self.prompt_template_encode_start_idx[template_type] prompt = [f"\n{p}" for p in prompt] prompt = [f"<|im_start|>user\n{p}<|im_end|>\n" for p in prompt] prompt = [p.replace("\n", "<|vision_start|><|image_pad|><|vision_end|>") for p in prompt] prompt = [template.format(p) for p in prompt] if images is not None: if not isinstance(images, list): images = [images] * len(prompt) elif len(images) < len(prompt) and len(prompt) % len(images) == 0: images = images * (len(prompt) // len(images)) inputs = self.processor( text=prompt, images=images, padding=True, return_tensors="pt", ).to(device) last_hidden_states = self._get_last_decoder_hidden_states(self.text_encoder, **inputs) prompt_embeds = last_hidden_states[:, drop_idx:] prompt_embeds_mask = inputs["attention_mask"][:, drop_idx:] if max_sequence_length is not None and prompt_embeds.shape[1] > max_sequence_length: prompt_embeds = prompt_embeds[:, -max_sequence_length:, :] prompt_embeds_mask = prompt_embeds_mask[:, -max_sequence_length:] _, 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 encode_prompt( self, prompt: Union[str, List[str]], device: Optional[torch.device] = None, num_images_per_prompt: int = 1, prompt_embeds: Optional[torch.Tensor] = None, prompt_embeds_mask: Optional[torch.Tensor] = None, max_sequence_length: int = 1024, template_type: str = "image", ) -> tuple[torch.Tensor, torch.Tensor]: """ Encode a text prompt into embeddings (text-only path). Pre-computed ``prompt_embeds`` bypass encoding entirely. Args: prompt: Prompt string or list of prompt strings. device: Target device. num_images_per_prompt: Number of outputs to generate per prompt. prompt_embeds: Pre-computed prompt embeddings. prompt_embeds_mask: Attention mask for pre-computed embeddings. max_sequence_length: Maximum output sequence length. template_type: Prompt template key (``"image"`` or ``"multiple_images"``). Returns: Tuple of (prompt_embeds, prompt_embeds_mask). """ 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, template_type, device) prompt_embeds = prompt_embeds[:, :max_sequence_length] prompt_embeds_mask = prompt_embeds_mask[:, :max_sequence_length] _, 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, ): """ Validate pipeline inputs before the forward pass. Raises: ValueError: On any invalid combination of arguments. """ 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("`callback_on_step_end_tensor_inputs` has invalid keys.") if prompt is not None and prompt_embeds is not None: raise ValueError("Cannot forward both `prompt` and `prompt_embeds`.") elif prompt is None and prompt_embeds is None: raise ValueError("Provide either `prompt` or `prompt_embeds`.") elif prompt is not None and not isinstance(prompt, (str, list)): raise ValueError("`prompt` has to be of type `str` or `list`.") if negative_prompt is not None and negative_prompt_embeds is not None: raise ValueError("Cannot forward both `negative_prompt` and `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` is required.") 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` is required.") def normalize_latents(self, latent: torch.Tensor) -> torch.Tensor: """ Normalise latents using per-channel statistics from the VAE config. Uses (latent - mean) / std when the VAE exposes ``latents_mean`` and ``latents_std``; otherwise falls back to scaling by ``scaling_factor``. Args: latent: Raw latent tensor from ``vae.encode``. Returns: Normalised latent tensor. """ if hasattr(self.vae.config, "latents_mean") and hasattr(self.vae.config, "latents_std"): latents_mean = ( torch.tensor(self.vae.config.latents_mean) .view(1, -1, 1, 1, 1) .to(device=latent.device, dtype=latent.dtype) ) latents_std = ( torch.tensor(self.vae.config.latents_std) .view(1, -1, 1, 1, 1) .to(device=latent.device, dtype=latent.dtype) ) latent = (latent - latents_mean) / latents_std else: latent = latent * self.vae.config.scaling_factor return latent def denormalize_latents(self, latent: torch.Tensor) -> torch.Tensor: """ Invert :meth:`normalize_latents` to recover the original latent scale. Args: latent: Normalised latent tensor. Returns: Latent tensor in the scale expected by ``vae.decode``. """ if hasattr(self.vae.config, "latents_mean") and hasattr(self.vae.config, "latents_std"): latents_mean = ( torch.tensor(self.vae.config.latents_mean) .view(1, -1, 1, 1, 1) .to(device=latent.device, dtype=latent.dtype) ) latents_std = ( torch.tensor(self.vae.config.latents_std) .view(1, -1, 1, 1, 1) .to(device=latent.device, dtype=latent.dtype) ) latent = latent * latents_std + latents_mean else: latent = latent / self.vae.config.scaling_factor return latent def prepare_latents( self, batch_size: int, num_channels_latents: int, height: int, width: int, video_length: int, dtype: torch.dtype, device: torch.device, generator: Optional[Union[torch.Generator, List[torch.Generator]]], latents: Optional[torch.Tensor] = None, image: Optional[List[Image.Image]] = None, enable_denormalization: bool = True, ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: """ Prepare the initial noisy latent tensor for the denoising loop. Args: batch_size: Number of samples in the batch. num_channels_latents: Latent channel dimension from the transformer config. height: Spatial height in pixels. width: Spatial width in pixels. video_length: Number of frames (1 for image inference). dtype: Floating-point dtype for the latent tensor. device: Target device. generator: RNG generator(s) for reproducible sampling. latents: Optional user-provided initial noise for the target slot. When ``None`` random noise is sampled. image: Optional list of PIL reference images to VAE-encode as conditioning slots. enable_denormalization: Whether to normalise encoded reference latents. Returns: Tuple of ``(latents, image_latents)`` where ``latents`` has shape ``(B, 1, C, T, H', W')`` and ``image_latents`` has shape ``(B, N_ref, C, T, H', W')`` or ``None`` when no reference images are given. Raises: ValueError: If ``generator`` is a list whose length differs from ``batch_size``. """ noise_shape = ( batch_size, 1, num_channels_latents, (video_length - 1) // self.vae_scale_factor_temporal + 1, int(height) // self.vae_scale_factor_spatial, int(width) // self.vae_scale_factor_spatial, ) if isinstance(generator, list) and len(generator) != batch_size: raise ValueError("Generator list length must match batch size.") if latents is None: latents = randn_tensor(noise_shape, generator=generator, device=device, dtype=dtype) else: latents = latents.to(device=device, dtype=dtype) image_latents = None if image is not None: if batch_size > len(image) and batch_size % len(image) == 0: image = image * (batch_size // len(image)) elif batch_size > len(image): raise ValueError(f"Cannot duplicate `image` of batch size {len(image)} to {batch_size} text prompts.") ref_img = [torch.from_numpy(np.array(x.convert("RGB"))) for x in image] ref_img = torch.stack(ref_img).to(device=device, dtype=dtype) ref_img = ref_img / 127.5 - 1.0 ref_img = ref_img.permute(0, 3, 1, 2).unsqueeze(2) image_latents = self.vae.encode(ref_img).latent_dist.sample() if enable_denormalization: image_latents = self.normalize_latents(image_latents) image_latents = image_latents.unsqueeze(1) # (B, 1, C, T, H', W') return latents, image_latents # ------------------------------------------------------------------ # Pipeline properties # ------------------------------------------------------------------ @property def guidance_scale(self) -> float: """Classifier-free guidance scale used in the current forward pass.""" return self._guidance_scale @property def do_classifier_free_guidance(self) -> bool: """True when guidance_scale > 1, enabling classifier-free guidance.""" return self._guidance_scale > 1 @property def num_timesteps(self) -> int: """Total number of denoising timesteps in the current forward pass.""" return self._num_timesteps @property def interrupt(self) -> bool: """When True, the denoising loop is interrupted at the next step.""" return self._interrupt # ------------------------------------------------------------------ # Forward pass # ------------------------------------------------------------------ @torch.no_grad() @replace_example_docstring(EXAMPLE_DOC_STRING) def __call__( self, image: PipelineImageInput | None = None, prompt: str | list[str] = None, height: int | None = None, width: int | None = None, num_inference_steps: int = 40, timesteps: List[int] = None, sigmas: List[float] = None, guidance_scale: float = 4.0, negative_prompt: Optional[Union[str, List[str]]] = None, num_images_per_prompt: Optional[int] = 1, generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, latents: Optional[torch.Tensor] = None, prompt_embeds: Optional[torch.Tensor] = None, prompt_embeds_mask: Optional[torch.Tensor] = None, negative_prompt_embeds: Optional[torch.Tensor] = None, negative_prompt_embeds_mask: Optional[torch.Tensor] = None, output_type: Optional[str] = "pil", return_dict: bool = True, callback_on_step_end: Optional[ Union[ Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks, ] ] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], max_sequence_length: int = 4096, enable_denormalization: bool = True, ): r""" Generate an edited image conditioned on a reference image and a text prompt. Args: prompt (`str` or `List[str]`): The prompt or prompts to guide generation. height (`int`): Height of the generated output in pixels. width (`int`): Width of the generated output in pixels. image (`PipelineImageInput`, *optional*): Reference image used for conditioning. When provided the pipeline operates in image-editing mode with ``num_items=2``. num_inference_steps (`int`, *optional*, defaults to 40): Number of denoising steps. More steps generally improve quality at the cost of slower inference. timesteps (`List[int]`, *optional*): Custom timesteps for the denoising process. When provided, ``num_inference_steps`` is inferred from the list length. sigmas (`List[float]`, *optional*): Custom sigmas for the denoising process. Mutually exclusive with ``timesteps``. guidance_scale (`float`, *optional*, defaults to 4.0): Classifier-free guidance scale. negative_prompt (`str` or `List[str]`, *optional*): Negative prompt(s) used to suppress undesired content. num_images_per_prompt (`int`, *optional*, defaults to 1): Number of generated samples per prompt. generator (`torch.Generator` or `List[torch.Generator]`, *optional*): RNG generator(s) for deterministic sampling. latents (`torch.Tensor`, *optional*): Pre-generated noisy latents for the target slot. Sampled from a Gaussian distribution when not provided. Can be used to seed generation from a specific starting noise tensor. prompt_embeds (`torch.Tensor`, *optional*): Pre-computed prompt embeddings. When provided ``prompt`` can be omitted. prompt_embeds_mask (`torch.Tensor`, *optional*): Attention mask for ``prompt_embeds``. negative_prompt_embeds (`torch.Tensor`, *optional*): Pre-computed negative prompt embeddings. negative_prompt_embeds_mask (`torch.Tensor`, *optional*): Attention mask for ``negative_prompt_embeds``. output_type (`str`, *optional*, defaults to ``"pil"``): Output format. Pass ``"latent"`` to return raw latents. return_dict (`bool`, *optional*, defaults to `True`): Whether to return a :class:`JoyImageEditPipelineOutput` or a plain tensor. callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*): Callback invoked at the end of each denoising step with signature ``(self, step: int, timestep: int, callback_kwargs: Dict)``. callback_on_step_end_tensor_inputs (`List[str]`, *optional*, defaults to ``["latents"]``): Tensor keys included in ``callback_kwargs`` for ``callback_on_step_end``. max_sequence_length (`int`, *optional*, defaults to 4096): Maximum sequence length for prompt encoding. enable_denormalization (`bool`, *optional*, defaults to `True`): Denormalise latents before VAE decoding. Examples: Returns: [`~pipelines.joyimage.JoyImageEditPipelineOutput`] or `torch.Tensor`: If ``return_dict`` is ``True``, returns a pipeline output object containing the generated image(s). Otherwise returns the image tensor directly. """ # Resize the input image to the nearest bucket resolution. # Or resize the specified height and width to the nearest bucket resolution. height, width = self.vae_image_processor.get_default_height_width(image, height, width) processed_image = None if image is not None: processed_image = self.vae_image_processor.resize_center_crop(image, (height, width)) 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, ) self._guidance_scale = guidance_scale 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 # num_items: 1 for unconditional generation, 2 for reference-image editing. num_items = 1 if image is None else 2 # Encode the conditioning prompt. if processed_image is not None: prompt_embeds, prompt_embeds_mask = self.encode_prompt_multiple_images( prompt=prompt, images=processed_image, 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, ) else: prompt_embeds, prompt_embeds_mask = self.encode_prompt( 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 self.do_classifier_free_guidance: # Build default negative prompts when none are provided. if negative_prompt is None and negative_prompt_embeds is None: negative_prompt = [""] * batch_size if processed_image is not None: negative_prompt_embeds, negative_prompt_embeds_mask = self.encode_prompt_multiple_images( prompt=negative_prompt, images=processed_image, 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, ) else: negative_prompt_embeds, negative_prompt_embeds_mask = self.encode_prompt( 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, ) timesteps, num_inference_steps = retrieve_timesteps( self.scheduler, num_inference_steps, device, timesteps, sigmas, ) num_channels_latents = self.transformer.config.in_channels noise_latents, image_latents = self.prepare_latents( batch_size * num_images_per_prompt, num_channels_latents, height, width, 1, # video_length = 1 for image inference prompt_embeds.dtype, device, generator, latents, image=( (processed_image if isinstance(processed_image, list) else [processed_image]) if processed_image is not None else None ), enable_denormalization=enable_denormalization, ) if image_latents is not None: latents = torch.cat([image_latents, noise_latents], dim=1) else: latents = noise_latents num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order self._num_timesteps = len(timesteps) with self.progress_bar(total=num_inference_steps) as progress_bar: for i, t in enumerate(timesteps): if self.interrupt: continue # Restore reference latents so they are never overwritten by the scheduler. if image_latents is not None: latents[:, : (num_items - 1)] = image_latents latent_model_input = latents t_expand = t.repeat(latent_model_input.shape[0]) noise_pred = self.transformer( hidden_states=latent_model_input, timestep=t_expand, encoder_hidden_states=prompt_embeds, return_dict=False, )[0] if self.do_classifier_free_guidance: noise_pred_uncond = self.transformer( hidden_states=latent_model_input, timestep=t_expand, encoder_hidden_states=negative_prompt_embeds, return_dict=False, )[0] comb_pred = noise_pred_uncond + self.guidance_scale * (noise_pred - noise_pred_uncond) # Rescale to match the conditional prediction norm (guidance rescaling). cond_norm = torch.norm(noise_pred, dim=2, keepdim=True) noise_norm = torch.norm(comb_pred, dim=2, keepdim=True) noise_pred = comb_pred * (cond_norm / noise_norm.clamp_min(1e-6)) latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0] 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) negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds) if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): if progress_bar is not None: progress_bar.update() if output_type != "latent": latents = latents.flatten(0, 1) if enable_denormalization: latents = self.denormalize_latents(latents) image = self.vae.decode(latents, return_dict=False)[0] image = image.unflatten(0, (batch_size * num_images_per_prompt, -1)) else: image = latents # Extract the target slot (last item) from each batch element. # (B, num_items, C, T, H, W) -> permute -> (B, num_items, T, C, H, W) -> [:, -1] -> (B, T, C, H, W) image = image.float().permute(0, 1, 3, 2, 4, 5)[:, -1].squeeze(1) image = self.image_processor.postprocess(image, output_type=output_type) self.maybe_free_model_hooks() if not return_dict: return image return JoyImageEditPipelineOutput(images=image)