Team Ai
Modelpublic

diffusers-modular/krea2-edit

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
inputs.py407 linesDownload Raw Back to root
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