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 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 