multimodalart's picture
multimodalart HF Staff
Sync the split MiniMax-H3 Spaces
186aa49 verified
Raw
History Blame Contribute Delete
28.5 kB
# 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