# 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 torch from ...configuration_utils import FrozenDict from ...guiders import ClassifierFreeGuidance from ...models.transformers.transformer_krea2 import Krea2Transformer2DModel from ...schedulers import FlowMatchEulerDiscreteScheduler from ...utils import logging from ..modular_pipeline import ( BlockState, LoopSequentialPipelineBlocks, ModularPipelineBlocks, PipelineState, ) from ..modular_pipeline_utils import ComponentSpec, InputParam, OutputParam from .modular_pipeline import Krea2ModularPipeline logger = logging.get_logger(__name__) # pylint: disable=invalid-name class Krea2LoopBeforeDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "Within the denoising loop: normalize the scheduler timestep into the model's flow time and broadcast it " "across the batch. Compose into the `sub_blocks` of a `Krea2DenoiseLoopWrapper`-based step." ) @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] @property def inputs(self) -> list[InputParam]: return [ InputParam(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), InputParam(name="batch_size", required=True, type_hint=int, description="Effective batch size."), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): num_train_timesteps = components.scheduler.config.num_train_timesteps block_state.timestep = (t / num_train_timesteps).expand(block_state.batch_size) return components, block_state class Krea2LoopDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "Within the denoising loop: run the `transformer` on the conditional (and, when the guider enables CFG, " "the negative) text features and combine them through the `guider`. Compose into `Krea2DenoiseStep`." ) @property def expected_components(self) -> list[ComponentSpec]: return [ ComponentSpec( "guider", ClassifierFreeGuidance, # Krea 2 uses cond-anchored CFG (`cond + scale * (cond - uncond)`), which is the # `use_original_formulation` branch of ClassifierFreeGuidance; scale 0 disables it (distilled TDM). 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(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), InputParam.template("num_inference_steps", required=True), InputParam( name="prompt_embeds", required=True, type_hint=torch.Tensor, description="Conditional stacked text features.", ), InputParam( name="prompt_embeds_mask", required=True, type_hint=torch.Tensor, description="Conditional text mask." ), InputParam( name="position_ids", required=True, type_hint=torch.Tensor, description="Shared rotary coordinates for the [text | image] sequence.", ), InputParam( name="negative_prompt_embeds", type_hint=torch.Tensor, description="Negative stacked text features." ), InputParam(name="negative_prompt_embeds_mask", type_hint=torch.Tensor, description="Negative text mask."), InputParam.template("attention_kwargs"), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): transformer = components.transformer latents = block_state.latents.to(transformer.dtype) timestep = block_state.timestep.to(transformer.dtype) guider_inputs = { "encoder_hidden_states": ( block_state.prompt_embeds.to(transformer.dtype), block_state.negative_prompt_embeds.to(transformer.dtype) if block_state.negative_prompt_embeds is not None else None, ), "encoder_attention_mask": ( block_state.prompt_embeds_mask, block_state.negative_prompt_embeds_mask, ), } 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 = {name: getattr(guider_state_batch, name) for name in guider_inputs} guider_state_batch.noise_pred = transformer( hidden_states=latents, timestep=timestep, position_ids=block_state.position_ids, attention_kwargs=block_state.attention_kwargs, return_dict=False, **cond_kwargs, )[0] components.guider.cleanup_models(components.transformer) block_state.noise_pred = components.guider(guider_state).pred return components, block_state class Krea2TurboLoopDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "Within the denoising loop: run the `transformer` on the conditional text features. The distilled Krea 2 " "turbo checkpoint runs without classifier-free guidance, so there is no negative branch or guider. Compose " "into the `sub_blocks` of `Krea2TurboDenoiseStep`." ) @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("transformer", Krea2Transformer2DModel)] @property def inputs(self) -> list[InputParam]: return [ InputParam(name="latents", required=True, type_hint=torch.Tensor, description="Packed image latents."), InputParam( name="prompt_embeds", required=True, type_hint=torch.Tensor, description="Conditional stacked text features.", ), InputParam( name="prompt_embeds_mask", required=True, type_hint=torch.Tensor, description="Conditional text mask." ), InputParam( name="position_ids", required=True, type_hint=torch.Tensor, description="Shared rotary coordinates for the [text | image] sequence.", ), InputParam.template("attention_kwargs"), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, block_state: BlockState, i: int, t: torch.Tensor): transformer = components.transformer latents = block_state.latents.to(transformer.dtype) timestep = block_state.timestep.to(transformer.dtype) block_state.noise_pred = transformer( hidden_states=latents, timestep=timestep, position_ids=block_state.position_ids, attention_kwargs=block_state.attention_kwargs, encoder_hidden_states=block_state.prompt_embeds.to(transformer.dtype), encoder_attention_mask=block_state.prompt_embeds_mask, return_dict=False, )[0] return components, block_state class Krea2LoopAfterDenoiser(ModularPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return "Within the denoising loop: scheduler step. Compose into a `Krea2DenoiseLoopWrapper`-based step." @property def expected_components(self) -> list[ComponentSpec]: return [ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler)] @property def intermediate_outputs(self) -> list[OutputParam]: return [OutputParam(name="latents", type_hint=torch.Tensor, description="The denoised 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] block_state.latents = block_state.latents.to(latents_dtype) return components, block_state class Krea2DenoiseLoopWrapper(LoopSequentialPipelineBlocks): model_name = "krea2" @property def description(self) -> str: return ( "Pipeline block that iteratively denoises the packed image latents over `timesteps`. " "The specific steps within each iteration can be customized with the `sub_blocks` attribute." ) @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="Denoising timesteps from set_timesteps.", ), InputParam.template("num_inference_steps", required=True), InputParam.template("attention_kwargs"), ] @torch.no_grad() def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState: block_state = self.get_block_state(state) 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) progress_bar.update() self.set_block_state(state, block_state) return components, state # auto_docstring class Krea2DenoiseStep(Krea2DenoiseLoopWrapper): """ Denoising loop that iteratively denoises the packed image latents over `timesteps`, running the transformer on the conditional (and, when the guider enables CFG, the negative) text features and combining them through the `guider`. Components: scheduler (`FlowMatchEulerDiscreteScheduler`) guider (`ClassifierFreeGuidance`) transformer (`Krea2Transformer2DModel`) Inputs: timesteps (`Tensor`): Denoising timesteps from set_timesteps. num_inference_steps (`int`): The number of denoising steps. attention_kwargs (`dict`, *optional*): Additional kwargs for attention processors. latents (`Tensor`): Packed image latents. batch_size (`int`): Effective batch size. prompt_embeds (`Tensor`): Conditional stacked text features. prompt_embeds_mask (`Tensor`): Conditional text mask. position_ids (`Tensor`): Shared rotary coordinates for the [text | image] sequence. negative_prompt_embeds (`Tensor`, *optional*): Negative stacked text features. negative_prompt_embeds_mask (`Tensor`, *optional*): Negative text mask. Outputs: latents (`Tensor`): The denoised latents. """ model_name = "krea2" block_classes = [Krea2LoopBeforeDenoiser, Krea2LoopDenoiser, Krea2LoopAfterDenoiser] block_names = ["before_denoiser", "denoiser", "after_denoiser"] @property def description(self) -> str: return ( "Denoising loop that iteratively denoises the packed image latents over `timesteps`, running the " "transformer on the conditional (and, when the guider enables CFG, the negative) text features and " "combining them through the `guider`." ) # auto_docstring class Krea2TurboDenoiseStep(Krea2DenoiseLoopWrapper): """ Denoising loop for the distilled Krea 2 turbo checkpoint that iteratively denoises the packed image latents over `timesteps`, running the transformer on the conditional text features. The distilled checkpoint runs without classifier-free guidance. Components: scheduler (`FlowMatchEulerDiscreteScheduler`) transformer (`Krea2Transformer2DModel`) Inputs: timesteps (`Tensor`): Denoising timesteps from set_timesteps. num_inference_steps (`int`): The number of denoising steps. attention_kwargs (`dict`, *optional*): Additional kwargs for attention processors. latents (`Tensor`): Packed image latents. batch_size (`int`): Effective batch size. prompt_embeds (`Tensor`): Conditional stacked text features. prompt_embeds_mask (`Tensor`): Conditional text mask. position_ids (`Tensor`): Shared rotary coordinates for the [text | image] sequence. Outputs: latents (`Tensor`): The denoised latents. """ model_name = "krea2" block_classes = [Krea2LoopBeforeDenoiser, Krea2TurboLoopDenoiser, Krea2LoopAfterDenoiser] block_names = ["before_denoiser", "denoiser", "after_denoiser"] @property def description(self) -> str: return ( "Denoising loop for the distilled Krea 2 turbo checkpoint that iteratively denoises the packed image " "latents over `timesteps`, running the transformer on the conditional text features. The distilled " "checkpoint runs without classifier-free guidance." )