Team Ai
Modelpublic

diffusers-modular/krea2-edit

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
denoise.py408 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 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