diffusers-modular/krea2-edit
0
1# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7# http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14 15 16import torch17 18from diffusers.modular_pipelines.modular_pipeline import ModularPipelineBlocks, PipelineState19from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam20from .modular_pipeline import Krea2ModularPipeline, Krea2Pachifier21 22 23# Copied from diffusers.modular_pipelines.qwenimage.inputs.repeat_tensor_to_batch_size24def repeat_tensor_to_batch_size(25 input_name: str,26 input_tensor: torch.Tensor,27 batch_size: int,28 num_images_per_prompt: int = 1,29) -> torch.Tensor:30 """Repeat tensor elements to match the final batch size.31 32 This function expands a tensor's batch dimension to match the final batch size (batch_size * num_images_per_prompt)33 by repeating each element along dimension 0.34 35 The input tensor must have batch size 1 or batch_size. The function will:36 - If batch size is 1: repeat each element (batch_size * num_images_per_prompt) times37 - If batch size equals batch_size: repeat each element num_images_per_prompt times38 39 Args:40 input_name (str): Name of the input tensor (used for error messages)41 input_tensor (torch.Tensor): The tensor to repeat. Must have batch size 1 or batch_size.42 batch_size (int): The base batch size (number of prompts)43 num_images_per_prompt (int, optional): Number of images to generate per prompt. Defaults to 1.44 45 Returns:46 torch.Tensor: The repeated tensor with final batch size (batch_size * num_images_per_prompt)47 48 Raises:49 ValueError: If input_tensor is not a torch.Tensor or has invalid batch size50 51 Examples:52 tensor = torch.tensor([[1, 2, 3]]) # shape: [1, 3] repeated = repeat_tensor_to_batch_size("image", tensor,53 batch_size=2, num_images_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]]) - shape:54 [4, 3]55 56 tensor = torch.tensor([[1, 2, 3], [4, 5, 6]]) # shape: [2, 3] repeated = repeat_tensor_to_batch_size("image",57 tensor, batch_size=2, num_images_per_prompt=2) repeated # tensor([[1, 2, 3], [1, 2, 3], [4, 5, 6], [4, 5, 6]])58 - shape: [4, 3]59 """60 # make sure input is a tensor61 if not isinstance(input_tensor, torch.Tensor):62 raise ValueError(f"`{input_name}` must be a tensor")63 64 # make sure input tensor e.g. image_latents has batch size 1 or batch_size same as prompts65 if input_tensor.shape[0] == 1:66 repeat_by = batch_size * num_images_per_prompt67 elif input_tensor.shape[0] == batch_size:68 repeat_by = num_images_per_prompt69 else:70 raise ValueError(71 f"`{input_name}` must have have batch size 1 or {batch_size}, but got {input_tensor.shape[0]}"72 )73 74 # expand the tensor to match the batch_size * num_images_per_prompt75 input_tensor = input_tensor.repeat_interleave(repeat_by, dim=0)76 77 return input_tensor78 79 80# Copied from diffusers.modular_pipelines.qwenimage.inputs.calculate_dimension_from_latents81def calculate_dimension_from_latents(latents: torch.Tensor, vae_scale_factor: int) -> tuple[int, int]:82 """Calculate image dimensions from latent tensor dimensions.83 84 This function converts latent space dimensions to image space dimensions by multiplying the latent height and width85 by the VAE scale factor.86 87 Args:88 latents (torch.Tensor): The latent tensor. Must have 4 or 5 dimensions.89 Expected shapes: [batch, channels, height, width] or [batch, channels, frames, height, width]90 vae_scale_factor (int): The scale factor used by the VAE to compress images.91 Typically 8 for most VAEs (image is 8x larger than latents in each dimension)92 93 Returns:94 tuple[int, int]: The calculated image dimensions as (height, width)95 96 Raises:97 ValueError: If latents tensor doesn't have 4 or 5 dimensions98 99 """100 # make sure the latents are not packed101 if latents.ndim != 4 and latents.ndim != 5:102 raise ValueError(f"unpacked latents must have 4 or 5 dimensions, but got {latents.ndim}")103 104 latent_height, latent_width = latents.shape[-2:]105 106 height = latent_height * vae_scale_factor107 width = latent_width * vae_scale_factor108 109 return height, width110 111 112class Krea2TextInputsStep(ModularPipelineBlocks):113 model_name = "krea2"114 115 @property116 def description(self) -> str:117 summary_section = (118 "Text input processing step that standardizes text embeddings for the pipeline.\n"119 "This step:\n"120 " 1. Determines `batch_size` and `dtype` based on `prompt_embeds`\n"121 " 2. Ensures all text embeddings have consistent batch sizes (batch_size * num_images_per_prompt)"122 )123 124 # Placement guidance125 placement_section = "\n\nThis block should be placed after all encoder steps to process the text embeddings before they are used in subsequent pipeline steps."126 127 return summary_section + placement_section128 129 @property130 def inputs(self) -> list[InputParam]:131 return [132 InputParam.template("num_images_per_prompt"),133 InputParam.template("prompt_embeds"),134 InputParam.template("prompt_embeds_mask"),135 InputParam.template("negative_prompt_embeds"),136 InputParam.template("negative_prompt_embeds_mask"),137 ]138 139 @property140 def intermediate_outputs(self) -> list[OutputParam]:141 return [142 OutputParam(name="batch_size", type_hint=int, description="The batch size of the prompt embeddings"),143 OutputParam(name="dtype", type_hint=torch.dtype, description="The data type of the prompt embeddings"),144 OutputParam.template("prompt_embeds", note="batch-expanded"),145 OutputParam.template("prompt_embeds_mask", note="batch-expanded"),146 OutputParam.template("negative_prompt_embeds", note="batch-expanded"),147 OutputParam.template("negative_prompt_embeds_mask", note="batch-expanded"),148 ]149 150 @staticmethod151 def check_inputs(152 prompt_embeds,153 prompt_embeds_mask,154 negative_prompt_embeds,155 negative_prompt_embeds_mask,156 ):157 if prompt_embeds.ndim != 4:158 raise ValueError(159 f"`prompt_embeds` must have 4 dimensions (batch_size, text_seq_len, num_text_layers, text_hidden_dim), but got {prompt_embeds.ndim}"160 )161 162 if negative_prompt_embeds is not None and negative_prompt_embeds_mask is None:163 raise ValueError("`negative_prompt_embeds_mask` is required when `negative_prompt_embeds` is not None")164 165 if negative_prompt_embeds is None and negative_prompt_embeds_mask is not None:166 raise ValueError("cannot pass `negative_prompt_embeds_mask` without `negative_prompt_embeds`")167 168 if prompt_embeds_mask.shape[0] != prompt_embeds.shape[0]:169 raise ValueError("`prompt_embeds_mask` must have the same batch size as `prompt_embeds`")170 171 elif negative_prompt_embeds is not None and negative_prompt_embeds.shape[0] != prompt_embeds.shape[0]:172 raise ValueError("`negative_prompt_embeds` must have the same batch size as `prompt_embeds`")173 174 elif (175 negative_prompt_embeds_mask is not None and negative_prompt_embeds_mask.shape[0] != prompt_embeds.shape[0]176 ):177 raise ValueError("`negative_prompt_embeds_mask` must have the same batch size as `prompt_embeds`")178 179 def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:180 block_state = self.get_block_state(state)181 182 self.check_inputs(183 prompt_embeds=block_state.prompt_embeds,184 prompt_embeds_mask=block_state.prompt_embeds_mask,185 negative_prompt_embeds=block_state.negative_prompt_embeds,186 negative_prompt_embeds_mask=block_state.negative_prompt_embeds_mask,187 )188 189 block_state.batch_size = block_state.prompt_embeds.shape[0]190 block_state.dtype = block_state.prompt_embeds.dtype191 192 # Krea 2 prompt embeddings are 4D: (batch_size, text_seq_len, num_text_layers, text_hidden_dim)193 _, seq_len, num_text_layers, dim = block_state.prompt_embeds.shape194 195 block_state.prompt_embeds = block_state.prompt_embeds.repeat(1, block_state.num_images_per_prompt, 1, 1)196 block_state.prompt_embeds = block_state.prompt_embeds.view(197 block_state.batch_size * block_state.num_images_per_prompt, seq_len, num_text_layers, dim198 )199 200 block_state.prompt_embeds_mask = block_state.prompt_embeds_mask.repeat(1, block_state.num_images_per_prompt)201 block_state.prompt_embeds_mask = block_state.prompt_embeds_mask.view(202 block_state.batch_size * block_state.num_images_per_prompt, seq_len203 )204 205 if block_state.negative_prompt_embeds is not None:206 _, seq_len, num_text_layers, dim = block_state.negative_prompt_embeds.shape207 block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.repeat(208 1, block_state.num_images_per_prompt, 1, 1209 )210 block_state.negative_prompt_embeds = block_state.negative_prompt_embeds.view(211 block_state.batch_size * block_state.num_images_per_prompt, seq_len, num_text_layers, dim212 )213 214 block_state.negative_prompt_embeds_mask = block_state.negative_prompt_embeds_mask.repeat(215 1, block_state.num_images_per_prompt216 )217 block_state.negative_prompt_embeds_mask = block_state.negative_prompt_embeds_mask.view(218 block_state.batch_size * block_state.num_images_per_prompt, seq_len219 )220 221 self.set_block_state(state, block_state)222 223 return components, state224 225 226class Krea2AdditionalInputsStep(ModularPipelineBlocks):227 model_name = "krea2"228 229 def __init__(230 self,231 image_latent_inputs: list[InputParam] | None = None,232 additional_batch_inputs: list[InputParam] | None = None,233 ):234 """Input processing step for additional (non-text) inputs.235 236 For image latent inputs: updates height/width if None, patchifies, and expands batch size. For additional batch237 inputs: expands batch dimensions to match the final batch size.238 239 Args:240 image_latent_inputs (list[InputParam], optional): Image latent inputs to process. Defaults to241 `[InputParam.template("image_latents")]`.242 additional_batch_inputs (list[InputParam], optional): Additional tensor inputs that only need batch243 expansion. Defaults to `[]`.244 """245 # by default, process `image_latents`246 if image_latent_inputs is None:247 image_latent_inputs = [InputParam.template("image_latents")]248 if additional_batch_inputs is None:249 additional_batch_inputs = []250 251 if not isinstance(image_latent_inputs, list):252 raise ValueError(f"image_latent_inputs must be a list, but got {type(image_latent_inputs)}")253 else:254 for input_param in image_latent_inputs:255 if not isinstance(input_param, InputParam):256 raise ValueError(f"image_latent_inputs must be a list of InputParam, but got {type(input_param)}")257 258 if not isinstance(additional_batch_inputs, list):259 raise ValueError(f"additional_batch_inputs must be a list, but got {type(additional_batch_inputs)}")260 else:261 for input_param in additional_batch_inputs:262 if not isinstance(input_param, InputParam):263 raise ValueError(264 f"additional_batch_inputs must be a list of InputParam, but got {type(input_param)}"265 )266 267 self._image_latent_inputs = image_latent_inputs268 self._additional_batch_inputs = additional_batch_inputs269 super().__init__()270 271 @property272 def description(self) -> str:273 summary_section = (274 "Input processing step that:\n"275 " 1. For image latent inputs: Updates height/width if None, patchifies, and expands batch size\n"276 " 2. For additional batch inputs: Expands batch dimensions to match final batch size"277 )278 279 inputs_info = ""280 if self._image_latent_inputs or self._additional_batch_inputs:281 inputs_info = "\n\nConfigured inputs:"282 if self._image_latent_inputs:283 inputs_info += f"\n - Image latent inputs: {[p.name for p in self._image_latent_inputs]}"284 if self._additional_batch_inputs:285 inputs_info += f"\n - Additional batch inputs: {[p.name for p in self._additional_batch_inputs]}"286 287 placement_section = "\n\nThis block should be placed after the encoder steps and the text input step."288 289 return summary_section + inputs_info + placement_section290 291 @property292 def expected_components(self) -> list[ComponentSpec]:293 return [294 ComponentSpec("pachifier", Krea2Pachifier, default_creation_method="from_config"),295 ]296 297 @property298 def inputs(self) -> list[InputParam]:299 inputs = [300 InputParam.template("num_images_per_prompt"),301 InputParam.template("batch_size"),302 InputParam.template("height"),303 InputParam.template("width"),304 ]305 # default is `image_latents`306 inputs += self._image_latent_inputs + self._additional_batch_inputs307 308 return inputs309 310 @property311 def intermediate_outputs(self) -> list[OutputParam]:312 outputs = [313 OutputParam(314 name="image_height",315 type_hint=int,316 description="The image height calculated from the image latents dimension",317 ),318 OutputParam(319 name="image_width",320 type_hint=int,321 description="The image width calculated from the image latents dimension",322 ),323 ]324 325 # `height`/`width` are not new outputs, but they will be updated if any image latent inputs are provided326 if len(self._image_latent_inputs) > 0:327 outputs.append(328 OutputParam(name="height", type_hint=int, description="if not provided, updated to image height")329 )330 outputs.append(331 OutputParam(name="width", type_hint=int, description="if not provided, updated to image width")332 )333 334 # image latent inputs are modified in place (patchified and batch-expanded)335 for input_param in self._image_latent_inputs:336 outputs.append(337 OutputParam(338 name=input_param.name,339 type_hint=input_param.type_hint,340 description=input_param.description + " (patchified and batch-expanded)",341 )342 )343 344 # additional batch inputs (batch-expanded only)345 for input_param in self._additional_batch_inputs:346 outputs.append(347 OutputParam(348 name=input_param.name,349 type_hint=input_param.type_hint,350 description=input_param.description + " (batch-expanded)",351 )352 )353 354 return outputs355 356 def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:357 block_state = self.get_block_state(state)358 359 # Process image latent inputs360 for input_param in self._image_latent_inputs:361 image_latent_input_name = input_param.name362 image_latent_tensor = getattr(block_state, image_latent_input_name)363 if image_latent_tensor is None:364 continue365 366 # 1. Calculate height/width from latents and update if not provided367 height, width = calculate_dimension_from_latents(image_latent_tensor, components.vae_scale_factor)368 block_state.height = block_state.height or height369 block_state.width = block_state.width or width370 371 if not hasattr(block_state, "image_height"):372 block_state.image_height = height373 if not hasattr(block_state, "image_width"):374 block_state.image_width = width375 376 # 2. Patchify377 image_latent_tensor = components.pachifier.pack_latents(image_latent_tensor)378 379 # 3. Expand batch size380 image_latent_tensor = repeat_tensor_to_batch_size(381 input_name=image_latent_input_name,382 input_tensor=image_latent_tensor,383 num_images_per_prompt=block_state.num_images_per_prompt,384 batch_size=block_state.batch_size,385 )386 387 setattr(block_state, image_latent_input_name, image_latent_tensor)388 389 # Process additional batch inputs (only batch expansion)390 for input_param in self._additional_batch_inputs:391 input_name = input_param.name392 input_tensor = getattr(block_state, input_name)393 if input_tensor is None:394 continue395 396 input_tensor = repeat_tensor_to_batch_size(397 input_name=input_name,398 input_tensor=input_tensor,399 num_images_per_prompt=block_state.num_images_per_prompt,400 batch_size=block_state.batch_size,401 )402 403 setattr(block_state, input_name, input_tensor)404 405 self.set_block_state(state, block_state)406 return components, state407 