Team Ai
Modelpublic

diffusers-modular/krea2-edit

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
before_denoise.py636 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 numpy as np18import torch19 20from diffusers.schedulers import FlowMatchEulerDiscreteScheduler21from diffusers.utils.torch_utils import randn_tensor22from diffusers.modular_pipelines.modular_pipeline import ModularPipelineBlocks, PipelineState23from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam24from .modular_pipeline import Krea2ModularPipeline, Krea2Pachifier25 26 27# Copied from diffusers.pipelines.flux.pipeline_flux.calculate_shift28def calculate_shift(29    image_seq_len,30    base_seq_len: int = 256,31    max_seq_len: int = 4096,32    base_shift: float = 0.5,33    max_shift: float = 1.15,34):35    m = (max_shift - base_shift) / (max_seq_len - base_seq_len)36    b = base_shift - m * base_seq_len37    mu = image_seq_len * m + b38    return mu39 40 41# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps42def retrieve_timesteps(43    scheduler,44    num_inference_steps: int | None = None,45    device: str | torch.device | None = None,46    timesteps: list[int] | None = None,47    sigmas: list[float] | None = None,48    **kwargs,49):50    r"""51    Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles52    custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.53 54    Args:55        scheduler (`SchedulerMixin`):56            The scheduler to get timesteps from.57        num_inference_steps (`int`):58            The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`59            must be `None`.60        device (`str` or `torch.device`, *optional*):61            The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.62        timesteps (`list[int]`, *optional*):63            Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,64            `num_inference_steps` and `sigmas` must be `None`.65        sigmas (`list[float]`, *optional*):66            Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,67            `num_inference_steps` and `timesteps` must be `None`.68 69    Returns:70        `tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the71        second element is the number of inference steps.72    """73    if timesteps is not None and sigmas is not None:74        raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")75    if timesteps is not None:76        accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())77        if not accepts_timesteps:78            raise ValueError(79                f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"80                f" timestep schedules. Please check whether you are using the correct scheduler."81            )82        scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)83        timesteps = scheduler.timesteps84        num_inference_steps = len(timesteps)85    elif sigmas is not None:86        accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())87        if not accept_sigmas:88            raise ValueError(89                f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"90                f" sigmas schedules. Please check whether you are using the correct scheduler."91            )92        scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)93        timesteps = scheduler.timesteps94        num_inference_steps = len(timesteps)95    else:96        scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)97        timesteps = scheduler.timesteps98    return timesteps, num_inference_steps99 100 101# Copied from diffusers.modular_pipelines.qwenimage.before_denoise.get_timesteps102def get_timesteps(scheduler, num_inference_steps, strength):103    # get the original timestep using init_timestep104    init_timestep = min(num_inference_steps * strength, num_inference_steps)105 106    t_start = int(max(num_inference_steps - init_timestep, 0))107    timesteps = scheduler.timesteps[t_start * scheduler.order :]108    if hasattr(scheduler, "set_begin_index"):109        scheduler.set_begin_index(t_start * scheduler.order)110 111    return timesteps, num_inference_steps - t_start112 113 114# ====================115# 1. PREPARE LATENTS116# ====================117 118 119class Krea2PrepareLatentsStep(ModularPipelineBlocks):120    model_name = "krea2"121 122    @property123    def description(self) -> str:124        return "Prepare initial random noise for the generation process"125 126    @property127    def expected_components(self) -> list[ComponentSpec]:128        return [129            ComponentSpec("pachifier", Krea2Pachifier, default_creation_method="from_config"),130        ]131 132    @property133    def inputs(self) -> list[InputParam]:134        return [135            InputParam.template("latents"),136            InputParam.template("height"),137            InputParam.template("width"),138            InputParam.template("num_images_per_prompt"),139            InputParam.template("generator"),140            InputParam.template("batch_size"),141            InputParam.template("dtype"),142        ]143 144    @property145    def intermediate_outputs(self) -> list[OutputParam]:146        return [147            OutputParam(name="height", type_hint=int, description="if not set, updated to default value"),148            OutputParam(name="width", type_hint=int, description="if not set, updated to default value"),149            OutputParam(150                name="latents",151                type_hint=torch.Tensor,152                description="The initial latents to use for the denoising process",153            ),154        ]155 156    @staticmethod157    def check_inputs(height, width, vae_scale_factor):158        if height is not None and height % (vae_scale_factor * 2) != 0:159            raise ValueError(f"Height must be divisible by {vae_scale_factor * 2} but is {height}")160 161        if width is not None and width % (vae_scale_factor * 2) != 0:162            raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}")163 164    @torch.no_grad()165    def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:166        block_state = self.get_block_state(state)167 168        self.check_inputs(169            height=block_state.height,170            width=block_state.width,171            vae_scale_factor=components.vae_scale_factor,172        )173 174        device = components._execution_device175        batch_size = block_state.batch_size * block_state.num_images_per_prompt176 177        # we can update the height and width here since it's used to generate the initial178        block_state.height = block_state.height or components.default_height179        block_state.width = block_state.width or components.default_width180 181        # VAE applies 8x compression on images but we must also account for packing which requires182        # latent height and width to be divisible by 2.183        latent_height = 2 * (int(block_state.height) // (components.vae_scale_factor * 2))184        latent_width = 2 * (int(block_state.width) // (components.vae_scale_factor * 2))185 186        shape = (batch_size, components.num_channels_latents, 1, latent_height, latent_width)187        if isinstance(block_state.generator, list) and len(block_state.generator) != batch_size:188            raise ValueError(189                f"You have passed a list of generators of length {len(block_state.generator)}, but requested an effective batch"190                f" size of {batch_size}. Make sure the batch size matches the length of the generators."191            )192        if block_state.latents is None:193            block_state.latents = randn_tensor(194                shape, generator=block_state.generator, device=device, dtype=block_state.dtype195            )196            block_state.latents = components.pachifier.pack_latents(block_state.latents)197 198        self.set_block_state(state, block_state)199        return components, state200 201 202class Krea2PrepareLatentsWithStrengthStep(ModularPipelineBlocks):203    model_name = "krea2"204 205    @property206    def description(self) -> str:207        return "Step that adds noise to image latents for image-to-image/inpainting. Should be run after set_timesteps, prepare_latents. Both noise and image latents should already be patchified."208 209    @property210    def expected_components(self) -> list[ComponentSpec]:211        return [212            ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler),213        ]214 215    @property216    def inputs(self) -> list[InputParam]:217        return [218            InputParam(219                name="latents",220                required=True,221                type_hint=torch.Tensor,222                description="The initial random noised, can be generated in prepare latent step.",223            ),224            InputParam.template("image_latents", note="Can be generated from vae encoder and updated in input step."),225            InputParam(226                name="timesteps",227                required=True,228                type_hint=torch.Tensor,229                description="The timesteps to use for the denoising process. Can be generated in set_timesteps step.",230            ),231        ]232 233    @property234    def intermediate_outputs(self) -> list[OutputParam]:235        return [236            OutputParam(237                name="initial_noise",238                type_hint=torch.Tensor,239                description="The initial random noised used for inpainting denoising.",240            ),241            OutputParam(242                name="latents",243                type_hint=torch.Tensor,244                description="The scaled noisy latents to use for inpainting/image-to-image denoising.",245            ),246        ]247 248    @staticmethod249    def check_inputs(image_latents, latents):250        if image_latents.shape[0] != latents.shape[0]:251            raise ValueError(252                f"`image_latents` must have have same batch size as `latents`, but got {image_latents.shape[0]} and {latents.shape[0]}"253            )254 255        if image_latents.ndim != 3:256            raise ValueError(f"`image_latents` must have 3 dimensions (patchified), but got {image_latents.ndim}")257 258    @torch.no_grad()259    def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:260        block_state = self.get_block_state(state)261 262        self.check_inputs(263            image_latents=block_state.image_latents,264            latents=block_state.latents,265        )266 267        # prepare latent timestep268        latent_timestep = block_state.timesteps[:1].repeat(block_state.latents.shape[0])269 270        # make copy of initial_noise271        block_state.initial_noise = block_state.latents272 273        # scale noise274        block_state.latents = components.scheduler.scale_noise(275            block_state.image_latents, latent_timestep, block_state.latents276        )277 278        self.set_block_state(state, block_state)279 280        return components, state281 282 283class Krea2CreateMaskLatentsStep(ModularPipelineBlocks):284    model_name = "krea2"285 286    @property287    def description(self) -> str:288        return "Step that creates mask latents from preprocessed mask_image by interpolating to latent space."289 290    @property291    def expected_components(self) -> list[ComponentSpec]:292        return [293            ComponentSpec("pachifier", Krea2Pachifier, default_creation_method="from_config"),294        ]295 296    @property297    def inputs(self) -> list[InputParam]:298        return [299            InputParam(300                name="processed_mask_image",301                required=True,302                type_hint=torch.Tensor,303                description="The processed mask to use for the inpainting process.",304            ),305            InputParam.template("height", required=True),306            InputParam.template("width", required=True),307            InputParam.template("dtype"),308        ]309 310    @property311    def intermediate_outputs(self) -> list[OutputParam]:312        return [313            OutputParam(314                name="mask", type_hint=torch.Tensor, description="The mask to use for the inpainting process."315            ),316        ]317 318    @torch.no_grad()319    def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:320        block_state = self.get_block_state(state)321 322        device = components._execution_device323 324        # VAE applies 8x compression on images but we must also account for packing which requires325        # latent height and width to be divisible by 2.326 327        height_latents = 2 * (int(block_state.height) // (components.vae_scale_factor * 2))328        width_latents = 2 * (int(block_state.width) // (components.vae_scale_factor * 2))329 330        block_state.mask = torch.nn.functional.interpolate(331            block_state.processed_mask_image,332            size=(height_latents, width_latents),333        )334 335        block_state.mask = block_state.mask.unsqueeze(2)336        block_state.mask = block_state.mask.repeat(1, components.num_channels_latents, 1, 1, 1)337        block_state.mask = block_state.mask.to(device=device, dtype=block_state.dtype)338 339        block_state.mask = components.pachifier.pack_latents(block_state.mask)340 341        self.set_block_state(state, block_state)342 343        return components, state344 345 346# ====================347# 2. SET TIMESTEPS348# ====================349 350 351class Krea2SetTimestepsStep(ModularPipelineBlocks):352    model_name = "krea2"353 354    @property355    def description(self) -> str:356        return "Step that sets the scheduler's timesteps for text-to-image generation. Should be run after prepare latents step."357 358    @property359    def expected_components(self) -> list[ComponentSpec]:360        return [361            ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler),362        ]363 364    @property365    def inputs(self) -> list[InputParam]:366        return [367            InputParam.template("num_inference_steps", default=28),368            InputParam.template("sigmas"),369            InputParam(370                name="mu",371                type_hint=float,372                description="Fixed timestep shift for the scheduler. Pass `1.15` for the few-step distilled (TDM/turbo) checkpoint; if not provided, computed from the image sequence length (base checkpoint behavior).",373            ),374            InputParam(375                name="latents",376                required=True,377                type_hint=torch.Tensor,378                description="The initial random noised latents for the denoising process. Can be generated in prepare latents step.",379            ),380        ]381 382    @property383    def intermediate_outputs(self) -> list[OutputParam]:384        return [385            OutputParam(386                name="timesteps", type_hint=torch.Tensor, description="The timesteps to use for the denoising process"387            ),388        ]389 390    def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:391        block_state = self.get_block_state(state)392 393        device = components._execution_device394        sigmas = (395            np.linspace(1.0, 1 / block_state.num_inference_steps, block_state.num_inference_steps)396            if block_state.sigmas is None397            else block_state.sigmas398        )399 400        mu = block_state.mu401        if mu is None:402            mu = calculate_shift(403                image_seq_len=block_state.latents.shape[1],404                base_seq_len=components.scheduler.config.get("base_image_seq_len", 256),405                max_seq_len=components.scheduler.config.get("max_image_seq_len", 6400),406                base_shift=components.scheduler.config.get("base_shift", 0.5),407                max_shift=components.scheduler.config.get("max_shift", 1.15),408            )409        block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps(410            scheduler=components.scheduler,411            num_inference_steps=block_state.num_inference_steps,412            device=device,413            sigmas=sigmas,414            mu=mu,415        )416 417        components.scheduler.set_begin_index(0)418 419        self.set_block_state(state, block_state)420 421        return components, state422 423 424class Krea2SetTimestepsWithStrengthStep(ModularPipelineBlocks):425    model_name = "krea2"426 427    @property428    def description(self) -> str:429        return "Step that sets the scheduler's timesteps for image-to-image generation, and inpainting. Should be run after prepare latents step."430 431    @property432    def expected_components(self) -> list[ComponentSpec]:433        return [434            ComponentSpec("scheduler", FlowMatchEulerDiscreteScheduler),435        ]436 437    @property438    def inputs(self) -> list[InputParam]:439        return [440            InputParam.template("num_inference_steps", default=28),441            InputParam.template("sigmas"),442            InputParam(443                name="mu",444                type_hint=float,445                description="Fixed timestep shift for the scheduler. Pass `1.15` for the few-step distilled (TDM/turbo) checkpoint; if not provided, computed from the image sequence length (base checkpoint behavior).",446            ),447            InputParam(448                "latents",449                required=True,450                type_hint=torch.Tensor,451                description="The latents to use for the denoising process. Can be generated in prepare latents step.",452            ),453            InputParam.template("strength"),454        ]455 456    @property457    def intermediate_outputs(self) -> list[OutputParam]:458        return [459            OutputParam(460                name="timesteps",461                type_hint=torch.Tensor,462                description="The timesteps to use for the denoising process.",463            ),464            OutputParam(465                name="num_inference_steps",466                type_hint=int,467                description="The number of denoising steps to perform at inference time. Updated based on strength.",468            ),469        ]470 471    def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:472        block_state = self.get_block_state(state)473 474        device = components._execution_device475        sigmas = (476            np.linspace(1.0, 1 / block_state.num_inference_steps, block_state.num_inference_steps)477            if block_state.sigmas is None478            else block_state.sigmas479        )480 481        mu = block_state.mu482        if mu is None:483            mu = calculate_shift(484                image_seq_len=block_state.latents.shape[1],485                base_seq_len=components.scheduler.config.get("base_image_seq_len", 256),486                max_seq_len=components.scheduler.config.get("max_image_seq_len", 6400),487                base_shift=components.scheduler.config.get("base_shift", 0.5),488                max_shift=components.scheduler.config.get("max_shift", 1.15),489            )490        block_state.timesteps, block_state.num_inference_steps = retrieve_timesteps(491            scheduler=components.scheduler,492            num_inference_steps=block_state.num_inference_steps,493            device=device,494            sigmas=sigmas,495            mu=mu,496        )497 498        block_state.timesteps, block_state.num_inference_steps = get_timesteps(499            scheduler=components.scheduler,500            num_inference_steps=block_state.num_inference_steps,501            strength=block_state.strength,502        )503 504        self.set_block_state(state, block_state)505 506        return components, state507 508 509# ====================510# 3. OTHER INPUTS FOR DENOISER511# ====================512 513## RoPE inputs for denoiser514 515 516class Krea2RoPEInputsStep(ModularPipelineBlocks):517    model_name = "krea2"518 519    @property520    def description(self) -> str:521        return (522            "Step that prepares the rotary position ids for the denoising process. Text tokens sit at the origin, "523            "image tokens carry their `(0, h, w)` latent-grid coordinates. Should be placed after prepare_latents step."524        )525 526    @property527    def expected_components(self) -> list[ComponentSpec]:528        return [529            ComponentSpec("pachifier", Krea2Pachifier, default_creation_method="from_config"),530        ]531 532    @property533    def inputs(self) -> list[InputParam]:534        return [535            InputParam.template("height", required=True),536            InputParam.template("width", required=True),537            InputParam.template("prompt_embeds_mask"),538        ]539 540    @property541    def intermediate_outputs(self) -> list[OutputParam]:542        return [543            OutputParam(544                name="position_ids",545                kwargs_type="denoiser_input_fields",546                type_hint=torch.Tensor,547                description="The rotary coordinates of shape (text_seq_len + grid_height * grid_width, 3) for the combined text-image sequence.",548            ),549        ]550 551    def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:552        block_state = self.get_block_state(state)553 554        device = components._execution_device555        patch_size = components.pachifier.config.patch_size556 557        text_seq_len = block_state.prompt_embeds_mask.shape[1]558        grid_height = block_state.height // (components.vae_scale_factor * patch_size)559        grid_width = block_state.width // (components.vae_scale_factor * patch_size)560 561        text_ids = torch.zeros(text_seq_len, 3, device=device)562        image_ids = torch.zeros(grid_height, grid_width, 3, device=device)563        image_ids[..., 1] = torch.arange(grid_height, device=device)[:, None]564        image_ids[..., 2] = torch.arange(grid_width, device=device)[None, :]565        image_ids = image_ids.reshape(grid_height * grid_width, 3)566        block_state.position_ids = torch.cat([text_ids, image_ids], dim=0)567 568        self.set_block_state(state, block_state)569 570        return components, state571 572 573class Krea2EditRoPEInputsStep(ModularPipelineBlocks):574    model_name = "krea2"575 576    @property577    def description(self) -> str:578        return (579            "Step that prepares the rotary position ids for the edit task: the `[text, image]` coordinates followed "580            "by the reference tokens' coordinates (each reference on its own frame axis). Should be placed after "581            "prepare_latents and the reference latents step."582        )583 584    @property585    def expected_components(self) -> list[ComponentSpec]:586        return [587            ComponentSpec("pachifier", Krea2Pachifier, default_creation_method="from_config"),588        ]589 590    @property591    def inputs(self) -> list[InputParam]:592        return [593            InputParam.template("height", required=True),594            InputParam.template("width", required=True),595            InputParam.template("prompt_embeds_mask"),596            InputParam(597                "reference_position_ids",598                required=True,599                type_hint=torch.Tensor,600                description="Rotary coordinates for the reference tokens. Can be generated in the reference latents step.",601            ),602        ]603 604    @property605    def intermediate_outputs(self) -> list[OutputParam]:606        return [607            OutputParam(608                name="position_ids",609                kwargs_type="denoiser_input_fields",610                type_hint=torch.Tensor,611                description="The rotary coordinates of shape (text_seq_len + grid_height * grid_width + ref_seq_len, 3) "612                "for the combined text-image-reference sequence.",613            ),614        ]615 616    def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:617        block_state = self.get_block_state(state)618 619        device = components._execution_device620        patch_size = components.pachifier.config.patch_size621 622        text_seq_len = block_state.prompt_embeds_mask.shape[1]623        grid_height = block_state.height // (components.vae_scale_factor * patch_size)624        grid_width = block_state.width // (components.vae_scale_factor * patch_size)625 626        text_ids = torch.zeros(text_seq_len, 3, device=device)627        image_ids = torch.zeros(grid_height, grid_width, 3, device=device)628        image_ids[..., 1] = torch.arange(grid_height, device=device)[:, None]629        image_ids[..., 2] = torch.arange(grid_width, device=device)[None, :]630        image_ids = image_ids.reshape(grid_height * grid_width, 3)631        block_state.position_ids = torch.cat([text_ids, image_ids, block_state.reference_position_ids], dim=0)632 633        self.set_block_state(state, block_state)634 635        return components, state636