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 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 