# Copyright 2026 Krea AI and The HuggingFace Team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. import inspect import torch from diffusers.configuration_utils import FrozenDict from diffusers.guiders import ClassifierFreeGuidance from .transformer_krea2 import Krea2Transformer2DModel from diffusers.schedulers import FlowMatchEulerDiscreteScheduler from diffusers.utils import logging from diffusers.modular_pipelines.modular_pipeline import BlockState, LoopSequentialPipelineBlocks, ModularPipelineBlocks, PipelineState from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam from .modular_pipeline import Krea2ModularPipeline logger = logging.get_logger(__name__) # ==================== # 1. LOOP STEPS (run at each denoising step) # ==================== # loop step:before denoiser class Krea2LoopBeforeDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "step within the denoising loop that prepares the latent input for the denoiser. " "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " "object (e.g. `Krea2DenoiseLoopWrapper`)" ) @property def inputs(self) -> list[InputParam]: return [ InputParam( name="latents", required=True, type_hint=torch.Tensor, description="The initial latents to use for the denoising process. Can be generated in prepare_latent step.", ), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): # one timestep block_state.timestep = t.expand(block_state.latents.shape[0]).to(block_state.latents.dtype) block_state.latent_model_input = block_state.latents return components, block_state # loop step:before denoiser (edit) -- appends the clean reference tokens to the denoiser input each step class Krea2EditLoopBeforeDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "step within the denoising loop that prepares the latent input for the edit denoiser: it appends the " "packed clean reference tokens after the noisy image tokens. This block should be used to compose the " "`sub_blocks` attribute of a `LoopSequentialPipelineBlocks` object (e.g. `Krea2EditDenoiseStep`)." ) @property def inputs(self) -> list[InputParam]: return [ InputParam( name="latents", required=True, type_hint=torch.Tensor, description="The initial latents to use for the denoising process. Can be generated in prepare_latent step.", ), InputParam( name="reference_latents", required=True, type_hint=torch.Tensor, description="Packed clean reference tokens to append to the denoiser sequence. Can be generated in the reference latents step.", ), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): block_state.timestep = t.expand(block_state.latents.shape[0]).to(block_state.latents.dtype) # Reference tokens are shared across the batch; expand and append them after the noisy image tokens. reference_latents = block_state.reference_latents.expand(block_state.latents.shape[0], -1, -1) block_state.latent_model_input = torch.cat( [block_state.latents, reference_latents.to(block_state.latents.dtype)], dim=1 ) return components, block_state # loop step:denoiser class Krea2LoopDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "step within the denoising loop that denoise the latent input for the denoiser. " "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " "object (e.g. `Krea2DenoiseLoopWrapper`)" ) @property def expected_components(self) -> list[ComponentSpec]: return [ ComponentSpec( "guider", ClassifierFreeGuidance, config=FrozenDict({"guidance_scale": 4.5, "use_original_formulation": True}), default_creation_method="from_config", ), ComponentSpec("transformer", Krea2Transformer2DModel), ] @property def inputs(self) -> list[InputParam]: return [ InputParam.template("denoiser_input_fields"), InputParam( "position_ids", required=True, type_hint=torch.Tensor, description="The rotary coordinates for the combined text-image sequence. Can be generated in prepare_rope_inputs step.", ), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): guider_inputs = { "encoder_hidden_states": ( getattr(block_state, "prompt_embeds", None), getattr(block_state, "negative_prompt_embeds", None), ), "encoder_attention_mask": ( getattr(block_state, "prompt_embeds_mask", None), getattr(block_state, "negative_prompt_embeds_mask", None), ), } transformer_args = set(inspect.signature(components.transformer.forward).parameters.keys()) additional_cond_kwargs = {} for field_name, field_value in block_state.denoiser_input_fields.items(): if field_name in transformer_args and field_name not in guider_inputs: additional_cond_kwargs[field_name] = field_value block_state.additional_cond_kwargs.update(additional_cond_kwargs) components.guider.set_state(step=i, num_inference_steps=block_state.num_inference_steps, timestep=t) guider_state = components.guider.prepare_inputs(guider_inputs) for guider_state_batch in guider_state: components.guider.prepare_models(components.transformer) cond_kwargs = {input_name: getattr(guider_state_batch, input_name) for input_name in guider_inputs.keys()} guider_state_batch.noise_pred = components.transformer( hidden_states=block_state.latent_model_input, timestep=block_state.timestep / 1000, return_dict=False, **cond_kwargs, **block_state.additional_cond_kwargs, )[0] components.guider.cleanup_models(components.transformer) guider_output = components.guider(guider_state) block_state.noise_pred = guider_output.pred return components, block_state # loop step:after denoiser class Krea2LoopAfterDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "step within the denoising loop that updates the latents. " "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " "object (e.g. `Krea2DenoiseLoopWrapper`)" ) @property def expected_components(self) -> list[ComponentSpec]: return [ ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam.template("latents"), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): latents_dtype = block_state.latents.dtype block_state.latents = components.scheduler.step( block_state.noise_pred, t, block_state.latents, return_dict=False, )[0] if block_state.latents.dtype != latents_dtype: if torch.backends.mps.is_available(): # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272 block_state.latents = block_state.latents.to(latents_dtype) return components, block_state class Krea2LoopAfterDenoiserInpaint(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "step within the denoising loop that updates the latents using mask and image_latents for inpainting. " "This block should be used to compose the `sub_blocks` attribute of a `LoopSequentialPipelineBlocks` " "object (e.g. `Krea2DenoiseLoopWrapper`)" ) @property def inputs(self) -> list[InputParam]: return [ InputParam( "mask", required=True, type_hint=torch.Tensor, description="The mask to use for the inpainting process. Can be generated in inpaint prepare latents step.", ), InputParam.template("image_latents"), InputParam( "initial_noise", required=True, type_hint=torch.Tensor, description="The initial noise to use for the inpainting process. Can be generated in inpaint prepare latents step.", ), ] @property def intermediate_outputs(self) -> list[OutputParam]: return [ OutputParam.template("latents"), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): block_state.init_latents_proper = block_state.image_latents if i < len(block_state.timesteps) - 1: block_state.noise_timestep = block_state.timesteps[i + 1] block_state.init_latents_proper = components.scheduler.scale_noise( block_state.init_latents_proper, torch.tensor([block_state.noise_timestep]), block_state.initial_noise ) block_state.latents = ( 1 - block_state.mask ) * block_state.init_latents_proper + block_state.mask * block_state.latents return components, block_state # ==================== # 2. DENOISE LOOP WRAPPER: define the denoising loop logic # ==================== class Krea2DenoiseLoopWrapper(LoopSequentialPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "Pipeline block that iteratively denoise the latents over `timesteps`. " "The specific steps with each iteration can be customized with `sub_blocks` attributes" ) @property def loop_expected_components(self) -> list[ComponentSpec]: return [ ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler), ] @property def loop_inputs(self) -> list[InputParam]: return [ InputParam( name="timesteps", required=True, type_hint=torch.Tensor, description="The timesteps to use for the denoising process. Can be generated in set_timesteps step.", ), InputParam.template("num_inference_steps", required=True), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) block_state.num_warmup_steps = max( len(block_state.timesteps) - block_state.num_inference_steps * components.scheduler.order, 0 ) block_state.additional_cond_kwargs = {} with self.progress_bar(total=block_state.num_inference_steps) as progress_bar: for i, t in enumerate(block_state.timesteps): components, block_state = self.loop_step(components, block_state, i=i, t=t) if i == len(block_state.timesteps) - 1 or ( (i + 1) > block_state.num_warmup_steps and (i + 1) % components.scheduler.order == 0 ): progress_bar.update() self.set_block_state(state, block_state) return components, state # ==================== # 3. DENOISE STEPS: compose the denoising loop with loop wrapper + loop steps # ==================== # Krea 2 (text2image, image2image) class Krea2DenoiseStep(Krea2DenoiseLoopWrapper): model_name = "krea2" block_classes = [ Krea2LoopBeforeDenoiser, Krea2LoopDenoiser, Krea2LoopAfterDenoiser, ] block_names = ["before_denoiser", "denoiser", "after_denoiser"] @property def description(self) -> str: return ( "Denoise step that iteratively denoise the latents.\n" "Its loop logic is defined in `Krea2DenoiseLoopWrapper.__call__` method\n" "At each iteration, it runs blocks defined in `sub_blocks` sequencially:\n" " - `Krea2LoopBeforeDenoiser`\n" " - `Krea2LoopDenoiser`\n" " - `Krea2LoopAfterDenoiser`\n" "This block supports text2image and image2image tasks for Krea 2." ) # Krea 2 (inpainting) class Krea2InpaintDenoiseStep(Krea2DenoiseLoopWrapper): model_name = "krea2" block_classes = [ Krea2LoopBeforeDenoiser, Krea2LoopDenoiser, Krea2LoopAfterDenoiser, Krea2LoopAfterDenoiserInpaint, ] block_names = ["before_denoiser", "denoiser", "after_denoiser", "after_denoiser_inpaint"] @property def description(self) -> str: return ( "Denoise step that iteratively denoise the latents. \n" "Its loop logic is defined in `Krea2DenoiseLoopWrapper.__call__` method \n" "At each iteration, it runs blocks defined in `sub_blocks` sequencially:\n" " - `Krea2LoopBeforeDenoiser`\n" " - `Krea2LoopDenoiser`\n" " - `Krea2LoopAfterDenoiser`\n" " - `Krea2LoopAfterDenoiserInpaint`\n" "This block supports inpainting tasks for Krea 2." ) # Krea 2 (reference-image edit) class Krea2EditDenoiseStep(Krea2DenoiseLoopWrapper): model_name = "krea2" block_classes = [ Krea2EditLoopBeforeDenoiser, Krea2LoopDenoiser, Krea2LoopAfterDenoiser, ] block_names = ["before_denoiser", "denoiser", "after_denoiser"] @property def description(self) -> str: return ( "Denoise step that iteratively denoise the latents for the reference-image edit task.\n" "Its loop logic is defined in `Krea2DenoiseLoopWrapper.__call__` method\n" "At each iteration, it runs blocks defined in `sub_blocks` sequencially:\n" " - `Krea2EditLoopBeforeDenoiser` (appends the clean reference tokens)\n" " - `Krea2LoopDenoiser`\n" " - `Krea2LoopAfterDenoiser`\n" "This block supports reference-image (edit) generation for Krea 2." )