Team Ai
Modelpublic

diffusers-modular/krea2-edit

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
modular_pipeline.py136 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 16from diffusers.configuration_utils import ConfigMixin, register_to_config17from diffusers.modular_pipelines.modular_pipeline import ModularPipeline18 19 20# Copied from diffusers.modular_pipelines.qwenimage.modular_pipeline.QwenImagePachifier with QwenImage->Krea221class Krea2Pachifier(ConfigMixin):22    """23    A class to pack and unpack latents for Krea2.24    """25 26    config_name = "config.json"27 28    @register_to_config29    def __init__(self, patch_size: int = 2):30        super().__init__()31 32    def pack_latents(self, latents):33        if latents.ndim != 4 and latents.ndim != 5:34            raise ValueError(f"Latents must have 4 or 5 dimensions, but got {latents.ndim}")35 36        if latents.ndim == 4:37            latents = latents.unsqueeze(2)38 39        batch_size, num_channels_latents, num_latent_frames, latent_height, latent_width = latents.shape40        patch_size = self.config.patch_size41 42        if latent_height % patch_size != 0 or latent_width % patch_size != 0:43            raise ValueError(44                f"Latent height and width must be divisible by {patch_size}, but got {latent_height} and {latent_width}"45            )46 47        latents = latents.view(48            batch_size,49            num_channels_latents,50            latent_height // patch_size,51            patch_size,52            latent_width // patch_size,53            patch_size,54        )55        latents = latents.permute(56            0, 2, 4, 1, 3, 557        )  # Batch_size, num_patches_height, num_patches_width, num_channels_latents, patch_size, patch_size58        latents = latents.reshape(59            batch_size,60            (latent_height // patch_size) * (latent_width // patch_size),61            num_channels_latents * patch_size * patch_size,62        )63 64        return latents65 66    def unpack_latents(self, latents, height, width, vae_scale_factor=8):67        if latents.ndim != 3:68            raise ValueError(f"Latents must have 3 dimensions, but got {latents.ndim}")69 70        batch_size, num_patches, channels = latents.shape71        patch_size = self.config.patch_size72 73        # VAE applies 8x compression on images but we must also account for packing which requires74        # latent height and width to be divisible by 2.75        height = patch_size * (int(height) // (vae_scale_factor * patch_size))76        width = patch_size * (int(width) // (vae_scale_factor * patch_size))77 78        latents = latents.view(79            batch_size,80            height // patch_size,81            width // patch_size,82            channels // (patch_size * patch_size),83            patch_size,84            patch_size,85        )86        latents = latents.permute(0, 3, 1, 4, 2, 5)87 88        latents = latents.reshape(batch_size, channels // (patch_size * patch_size), 1, height, width)89 90        return latents91 92 93class Krea2ModularPipeline(ModularPipeline):94    """95    A ModularPipeline for Krea 2.96 97    > [!WARNING] > This is an experimental feature and is likely to change in the future.98    """99 100    default_blocks_name = "Krea2AutoBlocks"101 102    @property103    def default_height(self):104        return self.default_sample_size * self.vae_scale_factor105 106    @property107    def default_width(self):108        return self.default_sample_size * self.vae_scale_factor109 110    @property111    def default_sample_size(self):112        return 128113 114    @property115    def vae_scale_factor(self):116        vae_scale_factor = 8117        if hasattr(self, "vae") and self.vae is not None:118            vae_scale_factor = 2 ** len(self.vae.temperal_downsample)119        return vae_scale_factor120 121    @property122    def num_channels_latents(self):123        num_channels_latents = 16124        if hasattr(self, "transformer") and self.transformer is not None:125            num_channels_latents = self.transformer.config.in_channels // 4126        return num_channels_latents127 128    @property129    def requires_unconditional_embeds(self):130        requires_unconditional_embeds = False131 132        if hasattr(self, "guider") and self.guider is not None:133            requires_unconditional_embeds = self.guider._enabled and self.guider.num_conditions > 1134 135        return requires_unconditional_embeds136