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 15import inspect16 17import torch18 19from diffusers.configuration_utils import FrozenDict20from diffusers.guiders import ClassifierFreeGuidance21from .transformer_krea2 import Krea2Transformer2DModel22from diffusers.schedulers import FlowMatchEulerDiscreteScheduler23from diffusers.utils import logging24from diffusers.modular_pipelines.modular_pipeline import BlockState, LoopSequentialPipelineBlocks, ModularPipelineBlocks, PipelineState25from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam26from .modular_pipeline import Krea2ModularPipeline27 28 29logger = logging.get_logger(__name__)30 31# ====================32# 1. LOOP STEPS (run at each denoising step)33# ====================34 35 36# loop step:before denoiser37class Krea2LoopBeforeDenoiser(ModularPipelineBlocks):38 model_name = "krea2"39 40 @property41 def description(self) -> str:42 return (43 "step within the denoising loop that prepares the latent input for the denoiser. "44 "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` "45 "object (e.g. `Krea2DenoiseLoopWrapper`)"46 )47 48 @property49 def inputs(self) -> list[InputParam]:50 return [51 InputParam(52 name="latents",53 required=True,54 type_hint=torch.Tensor,55 description="The initial latents to use for the denoising process. Can be generated in prepare_latent step.",56 ),57 ]58 59 @torch.no_grad()60 def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor):61 # one timestep62 block_state.timestep = t.expand(block_state.latents.shape[0]).to(block_state.latents.dtype)63 block_state.latent_model_input = block_state.latents64 return components, block_state65 66 67# loop step:before denoiser (edit) -- appends the clean reference tokens to the denoiser input each step68class Krea2EditLoopBeforeDenoiser(ModularPipelineBlocks):69 model_name = "krea2"70 71 @property72 def description(self) -> str:73 return (74 "step within the denoising loop that prepares the latent input for the edit denoiser: it appends the "75 "packed clean reference tokens after the noisy image tokens. This block should be used to compose the "76 "`sub_blocks` attribute of a `LoopSequentialPipelineBlocks` object (e.g. `Krea2EditDenoiseStep`)."77 )78 79 @property80 def inputs(self) -> list[InputParam]:81 return [82 InputParam(83 name="latents",84 required=True,85 type_hint=torch.Tensor,86 description="The initial latents to use for the denoising process. Can be generated in prepare_latent step.",87 ),88 InputParam(89 name="reference_latents",90 required=True,91 type_hint=torch.Tensor,92 description="Packed clean reference tokens to append to the denoiser sequence. Can be generated in the reference latents step.",93 ),94 ]95 96 @torch.no_grad()97 def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor):98 block_state.timestep = t.expand(block_state.latents.shape[0]).to(block_state.latents.dtype)99 # Reference tokens are shared across the batch; expand and append them after the noisy image tokens.100 reference_latents = block_state.reference_latents.expand(block_state.latents.shape[0], -1, -1)101 block_state.latent_model_input = torch.cat(102 [block_state.latents, reference_latents.to(block_state.latents.dtype)], dim=1103 )104 return components, block_state105 106 107# loop step:denoiser108class Krea2LoopDenoiser(ModularPipelineBlocks):109 model_name = "krea2"110 111 @property112 def description(self) -> str:113 return (114 "step within the denoising loop that denoise the latent input for the denoiser. "115 "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` "116 "object (e.g. `Krea2DenoiseLoopWrapper`)"117 )118 119 @property120 def expected_components(self) -> list[ComponentSpec]:121 return [122 ComponentSpec(123 "guider",124 ClassifierFreeGuidance,125 config=FrozenDict({"guidance_scale": 4.5, "use_original_formulation": True}),126 default_creation_method="from_config",127 ),128 ComponentSpec("transformer", Krea2Transformer2DModel),129 ]130 131 @property132 def inputs(self) -> list[InputParam]:133 return [134 InputParam.template("denoiser_input_fields"),135 InputParam(136 "position_ids",137 required=True,138 type_hint=torch.Tensor,139 description="The rotary coordinates for the combined text-image sequence. Can be generated in prepare_rope_inputs step.",140 ),141 ]142 143 @torch.no_grad()144 def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor):145 guider_inputs = {146 "encoder_hidden_states": (147 getattr(block_state, "prompt_embeds", None),148 getattr(block_state, "negative_prompt_embeds", None),149 ),150 "encoder_attention_mask": (151 getattr(block_state, "prompt_embeds_mask", None),152 getattr(block_state, "negative_prompt_embeds_mask", None),153 ),154 }155 156 transformer_args = set(inspect.signature(components.transformer.forward).parameters.keys())157 additional_cond_kwargs = {}158 for field_name, field_value in block_state.denoiser_input_fields.items():159 if field_name in transformer_args and field_name not in guider_inputs:160 additional_cond_kwargs[field_name] = field_value161 block_state.additional_cond_kwargs.update(additional_cond_kwargs)162 163 components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t)164 guider_state = components.guider.prepare_inputs(guider_inputs)165 166 for guider_state_batch in guider_state:167 components.guider.prepare_models(components.transformer)168 cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()}169 170 guider_state_batch.noise_pred = components.transformer(171 hidden_states=block_state.latent_model_input,172 timestep=block_state.timestep / 1000,173 return_dict=False,174 **cond_kwargs,175 **block_state.additional_cond_kwargs,176 )[0]177 178 components.guider.cleanup_models(components.transformer)179 180 guider_output = components.guider(guider_state)181 block_state.noise_pred = guider_output.pred182 183 return components, block_state184 185 186# loop step:after denoiser187class Krea2LoopAfterDenoiser(ModularPipelineBlocks):188 model_name = "krea2"189 190 @property191 def description(self) -> str:192 return (193 "step within the denoising loop that updates the latents. "194 "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` "195 "object (e.g. `Krea2DenoiseLoopWrapper`)"196 )197 198 @property199 def expected_components(self) -> list[ComponentSpec]:200 return [201 ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler),202 ]203 204 @property205 def intermediate_outputs(self) -> list[OutputParam]:206 return [207 OutputParam.template("latents"),208 ]209 210 @torch.no_grad()211 def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor):212 latents_dtype = block_state.latents.dtype213 block_state.latents = components.scheduler.step(214 block_state.noise_pred,215 t,216 block_state.latents,217 return_dict=False,218 )[0]219 220 if block_state.latents.dtype != latents_dtype:221 if torch.backends.mps.is_available():222 # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272223 block_state.latents = block_state.latents.to(latents_dtype)224 225 return components, block_state226 227 228class Krea2LoopAfterDenoiserInpaint(ModularPipelineBlocks):229 model_name = "krea2"230 231 @property232 def description(self) -> str:233 return (234 "step within the denoising loop that updates the latents using mask and image_latents for inpainting. "235 "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` "236 "object (e.g. `Krea2DenoiseLoopWrapper`)"237 )238 239 @property240 def inputs(self) -> list[InputParam]:241 return [242 InputParam(243 "mask",244 required=True,245 type_hint=torch.Tensor,246 description="The mask to use for the inpainting process. Can be generated in inpaint prepare latents step.",247 ),248 InputParam.template("image_latents"),249 InputParam(250 "initial_noise",251 required=True,252 type_hint=torch.Tensor,253 description="The initial noise to use for the inpainting process. Can be generated in inpaint prepare latents step.",254 ),255 ]256 257 @property258 def intermediate_outputs(self) -> list[OutputParam]:259 return [260 OutputParam.template("latents"),261 ]262 263 @torch.no_grad()264 def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor):265 block_state.init_latents_proper = block_state.image_latents266 if i < len(block_state.timesteps) - 1:267 block_state.noise_timestep = block_state.timesteps[i + 1]268 block_state.init_latents_proper = components.scheduler.scale_noise(269 block_state.init_latents_proper, torch.tensor([block_state.noise_timestep]), block_state.initial_noise270 )271 272 block_state.latents = (273 1 - block_state.mask274 ) * block_state.init_latents_proper + block_state.mask * block_state.latents275 276 return components, block_state277 278 279# ====================280# 2. DENOISE LOOP WRAPPER: define the denoising loop logic281# ====================282class Krea2DenoiseLoopWrapper(LoopSequentialPipelineBlocks):283 model_name = "krea2"284 285 @property286 def description(self) -> str:287 return (288 "Pipeline block that iteratively denoise the latents over `timesteps`. "289 "The specific steps with each iteration can be customized with `sub_blocks` attributes"290 )291 292 @property293 def loop_expected_components(self) -> list[ComponentSpec]:294 return [295 ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler),296 ]297 298 @property299 def loop_inputs(self) -> list[InputParam]:300 return [301 InputParam(302 name="timesteps",303 required=True,304 type_hint=torch.Tensor,305 description="The timesteps to use for the denoising process. Can be generated in set_timesteps step.",306 ),307 InputParam.template("num_inference_steps", required=True),308 ]309 310 @torch.no_grad()311 def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:312 block_state = self.get_block_state(state)313 314 block_state.num_warmup_steps = max(315 len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, 0316 )317 318 block_state.additional_cond_kwargs = {}319 320 with self.progress_bar(total=block_state.num_inference_steps) as progress_bar:321 for i, t in enumerate(block_state.timesteps):322 components, block_state = self.loop_step(components, block_state, i=i, t=t)323 if i == len(block_state.timesteps) - 1 or (324 (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0325 ):326 progress_bar.update()327 328 self.set_block_state(state, block_state)329 330 return components, state331 332 333# ====================334# 3. DENOISE STEPS: compose the denoising loop with loop wrapper + loop steps335# ====================336 337 338# Krea 2 (text2image, image2image)339class Krea2DenoiseStep(Krea2DenoiseLoopWrapper):340 model_name = "krea2"341 342 block_classes = [343 Krea2LoopBeforeDenoiser,344 Krea2LoopDenoiser,345 Krea2LoopAfterDenoiser,346 ]347 block_names = ["before_denoiser", "denoiser", "after_denoiser"]348 349 @property350 def description(self) -> str:351 return (352 "Denoise step that iteratively denoise the latents.\n"353 "Its loop logic is defined in `Krea2DenoiseLoopWrapper.__call__` method\n"354 "At each iteration, it runs blocks defined in `sub_blocks` sequencially:\n"355 " - `Krea2LoopBeforeDenoiser`\n"356 " - `Krea2LoopDenoiser`\n"357 " - `Krea2LoopAfterDenoiser`\n"358 "This block supports text2image and image2image tasks for Krea 2."359 )360 361 362# Krea 2 (inpainting)363class Krea2InpaintDenoiseStep(Krea2DenoiseLoopWrapper):364 model_name = "krea2"365 block_classes = [366 Krea2LoopBeforeDenoiser,367 Krea2LoopDenoiser,368 Krea2LoopAfterDenoiser,369 Krea2LoopAfterDenoiserInpaint,370 ]371 block_names = ["before_denoiser", "denoiser", "after_denoiser", "after_denoiser_inpaint"]372 373 @property374 def description(self) -> str:375 return (376 "Denoise step that iteratively denoise the latents. \n"377 "Its loop logic is defined in `Krea2DenoiseLoopWrapper.__call__` method \n"378 "At each iteration, it runs blocks defined in `sub_blocks` sequencially:\n"379 " - `Krea2LoopBeforeDenoiser`\n"380 " - `Krea2LoopDenoiser`\n"381 " - `Krea2LoopAfterDenoiser`\n"382 " - `Krea2LoopAfterDenoiserInpaint`\n"383 "This block supports inpainting tasks for Krea 2."384 )385 386 387# Krea 2 (reference-image edit)388class Krea2EditDenoiseStep(Krea2DenoiseLoopWrapper):389 model_name = "krea2"390 block_classes = [391 Krea2EditLoopBeforeDenoiser,392 Krea2LoopDenoiser,393 Krea2LoopAfterDenoiser,394 ]395 block_names = ["before_denoiser", "denoiser", "after_denoiser"]396 397 @property398 def description(self) -> str:399 return (400 "Denoise step that iteratively denoise the latents for the reference-image edit task.\n"401 "Its loop logic is defined in `Krea2DenoiseLoopWrapper.__call__` method\n"402 "At each iteration, it runs blocks defined in `sub_blocks` sequencially:\n"403 " - `Krea2EditLoopBeforeDenoiser` (appends the clean reference tokens)\n"404 " - `Krea2LoopDenoiser`\n"405 " - `Krea2LoopAfterDenoiser`\n"406 "This block supports reference-image (edit) generation for Krea 2."407 )408 