# Copyright 2026 The HuggingFace Team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. import inspect import numpy as np import torch from ...models import AnimaTextConditioner, CosmosTransformer3DModel from ...schedulers import FlowMatchEulerDiscreteScheduler from ...utils.torch_utils import randn_tensor from ..modular_pipeline import ModularPipelineBlocks, PipelineState from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam from .modular_pipeline import AnimaModularPipeline def retrieve_timesteps( scheduler, num_inference_steps: int | None = None, device: str | torch.device | None = None, timesteps: list[int] | None = None, sigmas: list[float] | None = None, **kwargs, ): if timesteps is not None and sigmas is not None: raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values") 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" f" timestep schedules. Please check whether you are using the correct scheduler." ) 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" f" sigmas schedules. Please check whether you are using the correct scheduler." ) 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 # Copied from diffusers.modular_pipelines.z_image.before_denoise.repeat_tensor_to_batch_size def repeat_tensor_to_batch_size( input_name: str, input_tensor: torch.Tensor, batch_size: int, num_images_per_prompt: int = 1, ) -> torch.Tensor: """Repeat tensor elements to match the final batch size. This function expands a tensor's batch dimension to match the final batch size (batch_size * num_images_per_prompt) by repeating each element along dimension 0. The input tensor must have batch size 1 or batch_size. The function will: - If batch size is 1: repeat each element (batch_size * num_images_per_prompt) times - If batch size equals batch_size: repeat each element num_images_per_prompt times Args: input_name (str): Name of the input tensor (used for error messages) input_tensor (torch.Tensor): The tensor to repeat. Must have batch size 1 or batch_size. batch_size (int): The base batch size (number of prompts) num_images_per_prompt (int, optional): Number of images to generate per prompt. Defaults to 1. Returns: torch.Tensor: The repeated tensor with final batch size (batch_size * num_images_per_prompt) Raises: ValueError: If input_tensor is not a torch.Tensor or has invalid batch size Examples: tensor = torch.tensor([[1, 2, 3]]) # shape: [1, 3] repeated = repeat_tensor_to_batch_size("image", tensor, batch_size=2, num_images_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]]) - shape: [4, 3] tensor = torch.tensor([[1, 2, 3], [4, 5, 6]]) # shape: [2, 3] repeated = repeat_tensor_to_batch_size("image", tensor, batch_size=2, num_images_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [4, 5, 6], [4, 5, 6]]) - shape: [4, 3] """ # make sure input is a tensor if not isinstance(input_tensor, torch.Tensor): raise ValueError(f"`{input_name}` must be a tensor") # make sure input tensor e.g. image_latents has batch size 1 or batch_size same as prompts if input_tensor.shape[0] == 1: repeat_by = batch_size * num_images_per_prompt elif input_tensor.shape[0] == batch_size: repeat_by = num_images_per_prompt else: raise ValueError( f"`{input_name}` must have have batch size 1 or {batch_size}, but got {input_tensor.shape[0]}" ) # expand the tensor to match the batch_size * num_images_per_prompt input_tensor = input_tensor.repeat_interleave(repeat_by, dim=0) return input_tensor class AnimaTextConditioningStep(ModularPipelineBlocks): model_name = "anima" @property def description(self) -> str: return "Map Qwen text encoder states and T5 token ids to Cosmos text conditioning for Anima." @property def expected_components(self) -> list[ComponentSpec]: return [ ComponentSpec("text_conditioner", AnimaTextConditioner), ComponentSpec("transformer", CosmosTransformer3DModel), ] @property def inputs(self) -> list[InputParam]: return [ InputParam( "qwen_prompt_embeds", required=True, type_hint=torch.Tensor, description="Qwen prompt embeddings generated by the text encoder step.", ), InputParam( "qwen_attention_mask", required=True, type_hint=torch.Tensor, description="Qwen prompt attention mask generated by the text encoder step.", ), InputParam( "t5_input_ids", required=True, type_hint=torch.Tensor, description="T5 prompt token ids generated by the text encoder step.", ), InputParam( "t5_attention_mask", required=True, type_hint=torch.Tensor, description="T5 prompt attention mask generated by the text encoder step.", ), InputParam( "negative_qwen_prompt_embeds", type_hint=torch.Tensor, description="Negative Qwen prompt embeddings generated by the text encoder step.", ), InputParam( "negative_qwen_attention_mask", type_hint=torch.Tensor, description="Negative Qwen prompt attention mask generated by the text encoder step.", ), InputParam( "negative_t5_input_ids", type_hint=torch.Tensor, description="Negative T5 prompt token ids generated by the text encoder step.", ), InputParam( "negative_t5_attention_mask", type_hint=torch.Tensor, description="Negative T5 prompt attention mask generated by the text encoder step.", ), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam( "prompt_embeds", type_hint=torch.Tensor, description="Conditioned prompt embeddings generated by the Anima text conditioner.", ), OutputParam( "negative_prompt_embeds", type_hint=torch.Tensor, description="Conditioned negative prompt embeddings generated by the Anima text conditioner.", ), ] @staticmethod def _condition_prompt_embeds( components: AnimaModularPipeline, qwen_prompt_embeds: torch.Tensor, qwen_attention_mask: torch.Tensor, t5_input_ids: torch.Tensor, t5_attention_mask: torch.Tensor, device: torch.device, conditioning_dtype: torch.dtype, output_dtype: torch.dtype, ) -> torch.Tensor: prompt_embeds = components.text_conditioner( source_hidden_states=qwen_prompt_embeds.to(device=device, dtype=conditioning_dtype), target_input_ids=t5_input_ids.to(device), target_attention_mask=t5_attention_mask.to(device), source_attention_mask=qwen_attention_mask.to(device), ) return prompt_embeds.to(dtype=output_dtype, device=device) @torch.no_grad() def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) device = components._execution_device conditioning_dtype = components.text_conditioner.dtype output_dtype = components.transformer.dtype block_state.prompt_embeds = self._condition_prompt_embeds( components, qwen_prompt_embeds=block_state.qwen_prompt_embeds, qwen_attention_mask=block_state.qwen_attention_mask, t5_input_ids=block_state.t5_input_ids, t5_attention_mask=block_state.t5_attention_mask, device=device, conditioning_dtype=conditioning_dtype, output_dtype=output_dtype, ) block_state.negative_prompt_embeds = None if block_state.negative_qwen_prompt_embeds is not None: block_state.negative_prompt_embeds = self._condition_prompt_embeds( components, qwen_prompt_embeds=block_state.negative_qwen_prompt_embeds, qwen_attention_mask=block_state.negative_qwen_attention_mask, t5_input_ids=block_state.negative_t5_input_ids, t5_attention_mask=block_state.negative_t5_attention_mask, device=device, conditioning_dtype=conditioning_dtype, output_dtype=output_dtype, ) self.set_block_state(state, block_state) return components, state class AnimaTextInputStep(ModularPipelineBlocks): model_name = "anima" @property def description(self) -> str: return "Input processing step that expands Anima prompt embeddings for the requested image batch." @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("transformer", CosmosTransformer3DModel)] @property def inputs(self) -> list[InputParam]: return [ InputParam.template("num_images_per_prompt"), InputParam( "prompt_embeds", required=True, type_hint=torch.Tensor, description="Conditioned prompt embeddings generated by the Anima text conditioner.", ), InputParam( "negative_prompt_embeds", type_hint=torch.Tensor, description="Conditioned negative prompt embeddings generated by the Anima text conditioner.", ), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam( "prompt_embeds", type_hint=torch.Tensor, kwargs_type="denoiser_input_fields", description="Prompt embeddings expanded to the final denoising batch.", ), OutputParam( "negative_prompt_embeds", type_hint=torch.Tensor, kwargs_type="denoiser_input_fields", description="Negative prompt embeddings expanded to the final denoising batch.", ), OutputParam( "batch_size", type_hint=int, description="Number of input prompts before `num_images_per_prompt` expansion.", ), OutputParam("dtype", type_hint=torch.dtype, description="Dtype used by the Anima denoiser."), ] @torch.no_grad() def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) block_state.batch_size = block_state.prompt_embeds.shape[0] block_state.dtype = components.transformer.dtype _, seq_len, _ = block_state.prompt_embeds.shape block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, block_state.num_images_per_prompt, 1) block_state.prompt_embeds = block_state.prompt_embeds.view( block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 ) if block_state.negative_prompt_embeds is not None: _, seq_len, _ = block_state.negative_prompt_embeds.shape block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.repeat( 1, block_state.num_images_per_prompt, 1 ) block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.view( block_state.batch_size * block_state.num_images_per_prompt, seq_len, -1 ) self.set_block_state(state, block_state) return components, state class AnimaImageInputStep(ModularPipelineBlocks): model_name = "anima" @property def description(self) -> str: return ( "Input processing step that expands Anima image latents to the final denoising batch " "and derives height/width from the latents when not provided." ) @property def inputs(self) -> list[InputParam]: return [ InputParam.template("image_latents"), InputParam( "batch_size", required=True, type_hint=int, description="Number of input prompts before `num_images_per_prompt` expansion.", ), InputParam.template("num_images_per_prompt"), InputParam.template("height"), InputParam.template("width"), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam( "image_latents", type_hint=torch.Tensor, description="Image latents expanded to the final denoising batch.", ), OutputParam("height", type_hint=int, description="Image height used for generation."), OutputParam("width", type_hint=int, description="Image width used for generation."), ] @torch.no_grad() def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) latent_height, latent_width = block_state.image_latents.shape[-2:] block_state.height = block_state.height or latent_height * components.vae_scale_factor block_state.width = block_state.width or latent_width * components.vae_scale_factor block_state.image_latents = repeat_tensor_to_batch_size( input_name="image_latents", input_tensor=block_state.image_latents, batch_size=block_state.batch_size, num_images_per_prompt=block_state.num_images_per_prompt, ) self.set_block_state(state, block_state) return components, state class AnimaPrepareLatentsStep(ModularPipelineBlocks): model_name = "anima" @property def description(self) -> str: return "Prepare noisy image latents and padding mask for Anima denoising." @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("transformer", CosmosTransformer3DModel)] @property def inputs(self) -> list[InputParam]: return [ InputParam.template("height"), InputParam.template("width"), InputParam.template("latents"), InputParam.template("num_images_per_prompt"), InputParam.template("generator"), InputParam( "batch_size", required=True, type_hint=int, description="Number of input prompts before `num_images_per_prompt` expansion.", ), InputParam("dtype", type_hint=torch.dtype, description="Dtype used by the Anima denoiser."), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam("height", type_hint=int, description="Image height used for generation."), OutputParam("width", type_hint=int, description="Image width used for generation."), OutputParam("latents", type_hint=torch.Tensor, description="Noisy latents for the denoising process."), OutputParam("padding_mask", type_hint=torch.Tensor, description="Cosmos padding mask for image latents."), ] def check_inputs(self, components: AnimaModularPipeline, block_state): divisor = components.vae_scale_factor * 2 if block_state.height % divisor != 0 or block_state.width % divisor != 0: raise ValueError( f"`height` and `width` have to be divisible by {divisor} but are {block_state.height} and" f" {block_state.width}." ) @staticmethod def prepare_latents( batch_size: int, num_channels_latents: int, height: int, width: int, vae_scale_factor: int, dtype: torch.dtype, device: torch.device, generator: torch.Generator | list[torch.Generator] | None, latents: torch.Tensor | None = None, ) -> torch.Tensor: if latents is not None: return latents.to(device=device, dtype=dtype) latent_height = height // vae_scale_factor latent_width = width // vae_scale_factor shape = (batch_size, num_channels_latents, 1, latent_height, latent_width) 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" f" size of {batch_size}. Make sure the batch size matches the length of the generators." ) return randn_tensor(shape, generator=generator, device=device, dtype=dtype) @torch.no_grad() def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) block_state.height = block_state.height or components.default_height block_state.width = block_state.width or components.default_width self.check_inputs(components, block_state) device = components._execution_device block_state.latents = self.prepare_latents( batch_size=block_state.batch_size * block_state.num_images_per_prompt, num_channels_latents=components.num_channels_latents, height=block_state.height, width=block_state.width, vae_scale_factor=components.vae_scale_factor, dtype=torch.float32, device=device, generator=block_state.generator, latents=block_state.latents, ) block_state.padding_mask = block_state.latents.new_zeros( 1, 1, block_state.height, block_state.width, dtype=block_state.dtype ) self.set_block_state(state, block_state) return components, state # Copied from diffusers.modular_pipelines.qwenimage.before_denoise.get_timesteps def get_timesteps(scheduler, num_inference_steps, strength): # get the original timestep using init_timestep init_timestep = min(num_inference_steps * strength, num_inference_steps) t_start = int(max(num_inference_steps - init_timestep, 0)) timesteps = scheduler.timesteps[t_start * scheduler.order :] if hasattr(scheduler, "set_begin_index"): scheduler.set_begin_index(t_start * scheduler.order) return timesteps, num_inference_steps - t_start class AnimaSetTimestepsStep(ModularPipelineBlocks): model_name = "anima" @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] @property def description(self) -> str: return "Set the scheduler timesteps for Anima inference." @property def inputs(self) -> list[InputParam]: return [ InputParam.template("num_inference_steps"), InputParam.template("sigmas"), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam("timesteps", type_hint=torch.Tensor, description="Timesteps for the denoising loop."), OutputParam("num_inference_steps", type_hint=int, description="Number of denoising steps."), ] @torch.no_grad() def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) device = components._execution_device sigmas = ( np.linspace(1.0, 1 / block_state.num_inference_steps, block_state.num_inference_steps) if block_state.sigmas is None else block_state.sigmas ) block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps( components.scheduler, device=device, sigmas=sigmas, ) components.scheduler.set_begin_index(0) self.set_block_state(state, block_state) return components, state class AnimaImg2ImgSetTimestepsStep(ModularPipelineBlocks): """Set the scheduler timesteps for Anima image-to-image inference. This step computes the full timestep schedule, then slices it based on ``strength`` via ``get_timesteps()``, which also sets the scheduler's begin index. Components: scheduler (`FlowMatchEulerDiscreteScheduler`) Inputs: num_inference_steps (`int`, *optional*, defaults to 50): The number of denoising steps. sigmas (`list`, *optional*): Custom sigmas for the denoising process. strength (`float`, *optional*, defaults to 0.9): How much to transform the reference image. Outputs: timesteps (`Tensor`): Timestep schedule sliced by ``strength``. num_inference_steps (`int`): Number of denoising steps after strength-based slicing. """ model_name = "anima" @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] @property def description(self) -> str: return "Set the scheduler timesteps for Anima image-to-image inference, sliced by strength." @property def inputs(self) -> list[InputParam]: return [ InputParam.template("num_inference_steps"), InputParam.template("sigmas"), InputParam.template("strength"), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam( "timesteps", type_hint=torch.Tensor, description="Timestep schedule sliced by strength.", ), OutputParam( "num_inference_steps", type_hint=int, description="Number of denoising steps after strength-based slicing.", ), ] @torch.no_grad() def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) device = components._execution_device sigmas = ( np.linspace(1.0, 1 / block_state.num_inference_steps, block_state.num_inference_steps) if block_state.sigmas is None else block_state.sigmas ) block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps( components.scheduler, device=device, sigmas=sigmas, ) block_state.timesteps, block_state.num_inference_steps = get_timesteps( components.scheduler, block_state.num_inference_steps, block_state.strength ) self.set_block_state(state, block_state) return components, state class AnimaImg2ImgPrepareLatentsStep(ModularPipelineBlocks): """Prepares noisy latents for Anima image-to-image generation. Generates noise and mixes it with the image latents via ``scheduler.scale_noise()`` at the first sliced timestep. The image latents are expected to already be expanded to the final batch size by ``AnimaImageInputStep``. Components: scheduler (`FlowMatchEulerDiscreteScheduler`) Inputs: image_latents (`Tensor`): Encoded image latents, expanded to the final denoising batch. timesteps (`Tensor`): Timestep schedule sliced by ``strength`` from ``AnimaImg2ImgSetTimestepsStep``. generator (`Generator`, *optional*): Torch generator for deterministic generation. latents (`Tensor`, *optional*): Pre-computed noise tensor. Generated randomly if ``None``. dtype (`torch.dtype`): Dtype used by the Anima denoiser. height (`int`): Image height. width (`int`): Image width. Outputs: latents (`Tensor`): Noisy image latents for the denoising loop. padding_mask (`Tensor`): Cosmos padding mask for the image latents. """ model_name = "anima" @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] @property def description(self) -> str: return ( "Prepares noisy image-to-image latents for Anima by adding noise to the encoded " "image latents via scheduler.scale_noise()." ) @property def inputs(self) -> list[InputParam]: return [ InputParam.template("image_latents"), InputParam.template("timesteps", required=True), InputParam.template("generator"), InputParam.template("latents"), InputParam("dtype", type_hint=torch.dtype, description="Dtype used by the Anima denoiser."), InputParam.template("height"), InputParam.template("width"), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam("latents", type_hint=torch.Tensor, description="Noisy latents for the denoising loop."), OutputParam("padding_mask", type_hint=torch.Tensor, description="Cosmos padding mask for image latents."), ] @torch.no_grad() def __call__(self, components: AnimaModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) device = components._execution_device image_latents = block_state.image_latents.to(device=device, dtype=torch.float32) if block_state.latents is None: noise = randn_tensor( image_latents.shape, generator=block_state.generator, device=device, dtype=torch.float32, ) else: noise = block_state.latents.to(device=device, dtype=torch.float32) latent_timestep = block_state.timesteps[:1].repeat(image_latents.shape[0]) block_state.latents = components.scheduler.scale_noise(image_latents, latent_timestep, noise) block_state.padding_mask = block_state.latents.new_zeros( 1, 1, block_state.height, block_state.width, dtype=block_state.dtype ) self.set_block_state(state, block_state) return components, state