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 15"""16Text and VAE encoder blocks for Krea 2 pipelines.17"""18 19import math20 21import numpy as np22import PIL.Image23import torch24import torch.nn.functional as F25from transformers import Qwen2Tokenizer, Qwen3VLModel, Qwen3VLProcessor26 27from diffusers.configuration_utils import FrozenDict28from diffusers.guiders import ClassifierFreeGuidance29from diffusers.image_processor import InpaintProcessor, VaeImageProcessor30from diffusers.models import AutoencoderKLQwenImage31from diffusers.utils import logging32from diffusers.modular_pipelines.modular_pipeline import ModularPipelineBlocks, PipelineState33from diffusers.modular_pipelines.modular_pipeline_utils import ComponentSpec, InputParam, OutputParam34from .modular_pipeline import Krea2ModularPipeline, Krea2Pachifier35 36 37logger = logging.get_logger(__name__)38 39# Text conditioning uses the Qwen-Image chat template, tokenized as a fixed-length block: the prompt is padded to a40# fixed length first and the assistant suffix is appended after the padding (matching how the model was sampled at41# training time). The first `KREA2_PROMPT_TEMPLATE_START_IDX` (system prefix) tokens are dropped from the encoder42# outputs.43KREA2_PROMPT_TEMPLATE_PREFIX = (44 "<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, "45 "spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n"46)47KREA2_PROMPT_TEMPLATE_SUFFIX = "<|im_end|>\n<|im_start|>assistant\n"48KREA2_PROMPT_TEMPLATE_START_IDX = 3449KREA2_PROMPT_TEMPLATE_NUM_SUFFIX_TOKENS = 550 51# Indices into the text encoder's `hidden_states` tuple (0 is the embedding output) whose states are stacked per token52# and fed to the transformer's text fusion stage. These are the Krea 2 (Qwen3-VL-4B) taps; must have53# `transformer.config.num_text_layers` entries.54KREA2_TEXT_ENCODER_SELECT_LAYERS = (2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35)55 56 57def get_krea2_prompt_embeds(58 text_encoder,59 tokenizer,60 prompt: str | list[str],61 text_encoder_select_layers: tuple[int, ...] = KREA2_TEXT_ENCODER_SELECT_LAYERS,62 prompt_template_prefix: str = KREA2_PROMPT_TEMPLATE_PREFIX,63 prompt_template_suffix: str = KREA2_PROMPT_TEMPLATE_SUFFIX,64 prompt_template_start_idx: int = KREA2_PROMPT_TEMPLATE_START_IDX,65 prompt_template_num_suffix_tokens: int = KREA2_PROMPT_TEMPLATE_NUM_SUFFIX_TOKENS,66 max_sequence_length: int = 512,67 device: torch.device | None = None,68):69 """Tokenize `prompt` into the fixed-length Krea 2 layout and tap the selected encoder hidden states.70 71 Returns a `(prompt_embeds, prompt_embeds_mask)` tuple of shapes `(batch_size, text_seq_len, num_text_layers,72 text_hidden_dim)` and `(batch_size, text_seq_len)` (bool).73 """74 prompt = [prompt] if isinstance(prompt, str) else prompt75 prefix_idx = prompt_template_start_idx76 text = [prompt_template_prefix + e for e in prompt]77 text_tokens = tokenizer(78 text,79 truncation=True,80 padding="max_length",81 max_length=max_sequence_length + prefix_idx - prompt_template_num_suffix_tokens,82 return_tensors="pt",83 ).to(device)84 suffix_tokens = tokenizer([prompt_template_suffix] * len(text), return_tensors="pt").to(device)85 86 input_ids = torch.cat([text_tokens.input_ids, suffix_tokens.input_ids], dim=1)87 attention_mask = torch.cat([text_tokens.attention_mask, suffix_tokens.attention_mask], dim=1).bool()88 89 # Krea 2 pads in the middle of the template (`[prefix | prompt | PAD | suffix]`), so the suffix tokens sit90 # downstream of the padding. The text features must use positions that count only real tokens (padding does91 # not consume a position) to match how the model was trained; otherwise the suffix gets a shifted mRoPE phase.92 # `Qwen3VLModel`'s default raw-index positions would place the suffix at ~max_length instead. Build the93 # cumulative-valid-token positions explicitly and broadcast across the 3 mRoPE axes (T/H/W are equal for text).94 position_ids = (attention_mask.long().cumsum(dim=-1) - 1).clamp(min=0)95 position_ids = position_ids.unsqueeze(0).expand(3, -1, -1)96 97 outputs = text_encoder(98 input_ids=input_ids,99 attention_mask=attention_mask,100 position_ids=position_ids,101 output_hidden_states=True,102 )103 hidden_states = torch.stack([outputs.hidden_states[i] for i in text_encoder_select_layers], dim=2)104 105 prompt_embeds = hidden_states[:, prefix_idx:]106 prompt_embeds_mask = attention_mask[:, prefix_idx:]107 return prompt_embeds, prompt_embeds_mask108 109 110# Reference images ride in the user message ahead of the prompt through named vision placeholders; the Qwen3-VL111# processor expands each `<|image_pad|>` into the image's token grid so the text conditioning "sees" the references.112KREA2_EDIT_IMAGE_PLACEHOLDER = "Picture {}: <|vision_start|><|image_pad|><|vision_end|>"113 114 115def to_chw_tensor(image) -> torch.Tensor:116 """Convert a PIL image / numpy array / CHW tensor to a float CHW tensor in [0, 1]."""117 if isinstance(image, torch.Tensor):118 t = image.squeeze(0) if image.ndim == 4 else image119 t = t.float()120 if t.min() < 0: # assume [-1, 1]121 t = (t + 1.0) / 2.0122 return t.clamp(0, 1)123 if isinstance(image, np.ndarray):124 image = PIL.Image.fromarray(image)125 image = image.convert("RGB")126 arr = np.asarray(image).astype(np.float32) / 255.0127 return torch.from_numpy(arr).permute(2, 0, 1)128 129 130def prep_vl_images(images: list[torch.Tensor], max_pixels: int) -> list[torch.Tensor]:131 """Resize reference images for the Qwen3-VL pass: aspect-preserving downscale (never upscaled) to fit132 `max_pixels` total area. The MLLM only needs a coarse view of the references; high-res detail flows through the133 VAE reference latents."""134 prepped = []135 for img in images:136 h, w = img.shape[1], img.shape[2]137 scale = min(1.0, math.sqrt(max_pixels / (h * w)))138 nh, nw = max(round(h * scale), 28), max(round(w * scale), 28)139 if (nh, nw) != (h, w):140 img = F.interpolate(img.unsqueeze(0).float(), size=(nh, nw), mode="bicubic", antialias=True).squeeze(0)141 img = img.clamp(0, 1)142 prepped.append(img.float())143 return prepped144 145 146def get_krea2_edit_prompt_embeds(147 text_encoder,148 tokenizer,149 processor,150 prompt: str | list[str],151 images: list[torch.Tensor] | None = None,152 text_encoder_select_layers: tuple[int, ...] = KREA2_TEXT_ENCODER_SELECT_LAYERS,153 prompt_template_prefix: str = KREA2_PROMPT_TEMPLATE_PREFIX,154 prompt_template_suffix: str = KREA2_PROMPT_TEMPLATE_SUFFIX,155 prompt_template_start_idx: int = KREA2_PROMPT_TEMPLATE_START_IDX,156 max_sequence_length: int = 512,157 device: torch.device | None = None,158):159 """Encode prompts for the edit task, embedding reference images (a coarse VL view of each) into the text160 conditioning through the Qwen3-VL vision tower.161 162 Unlike `get_krea2_prompt_embeds`, prompts are tokenized at their natural (unpadded) length: the processor163 expands each `<|image_pad|>` placeholder into a run of vision tokens that must stay intact, so no truncation or164 fixed padding is applied before encoding. All prompts share the same `images`. Returns a165 `(prompt_embeds, prompt_embeds_mask)` tuple of shapes `(batch_size, text_seq_len, num_text_layers,166 text_hidden_dim)` and `(batch_size, text_seq_len)` (bool), right-padded across the batch.167 """168 prompt = [prompt] if isinstance(prompt, str) else prompt169 prefix_idx = prompt_template_start_idx170 171 # The suffix is tokenized separately so it lands after the (image +) prompt tokens.172 suffix_inputs = tokenizer([prompt_template_suffix], return_tensors="pt").to(device)173 suffix_ids = suffix_inputs["input_ids"]174 suffix_mask = suffix_inputs["attention_mask"].bool()175 176 image_prompt = ""177 if images:178 image_prompt = "".join(KREA2_EDIT_IMAGE_PLACEHOLDER.format(i + 1) for i in range(len(images)))179 180 features = []181 for p in prompt:182 text = prompt_template_prefix + image_prompt + p183 extra_inputs = {}184 if images:185 # No truncation: the expanded image-pad runs must stay intact.186 inputs = processor(text=[text], images=list(images), return_tensors="pt", do_rescale=False).to(device)187 for k, v in inputs.items():188 if k in ("input_ids", "attention_mask"):189 continue190 if isinstance(v, torch.Tensor) and v.is_floating_point():191 v = v.to(text_encoder.dtype)192 extra_inputs[k] = v193 else:194 inputs = tokenizer(195 [text], truncation=True, max_length=max_sequence_length + prefix_idx, return_tensors="pt"196 ).to(device)197 198 input_ids = torch.cat([inputs["input_ids"], suffix_ids], dim=1)199 attention_mask = torch.cat([inputs["attention_mask"].bool(), suffix_mask], dim=1)200 201 # mm_token_type_ids (used for M-RoPE) must cover the appended suffix tokens too; they are plain text -> type 0.202 if "mm_token_type_ids" in extra_inputs:203 tt = extra_inputs["mm_token_type_ids"]204 extra_inputs["mm_token_type_ids"] = torch.cat([tt, torch.zeros_like(suffix_ids, dtype=tt.dtype)], dim=1)205 206 outputs = text_encoder(207 input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True, **extra_inputs208 )209 hidden_states = torch.stack([outputs.hidden_states[i] for i in text_encoder_select_layers], dim=2)210 # Drop the system-prefix tokens; what remains is (image +) prompt + suffix.211 features.append(hidden_states[0, prefix_idx:])212 213 max_len = max(f.shape[0] for f in features)214 prompt_embeds = features[0].new_zeros(len(features), max_len, *features[0].shape[1:])215 prompt_embeds_mask = torch.zeros(len(features), max_len, dtype=torch.bool, device=device)216 for i, f in enumerate(features):217 prompt_embeds[i, : f.shape[0]] = f218 prompt_embeds_mask[i, : f.shape[0]] = True219 220 return prompt_embeds, prompt_embeds_mask221 222 223# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_img2img.retrieve_latents224def retrieve_latents(225 encoder_output: torch.Tensor, generator: torch.Generator | None = None, sample_mode: str = "sample"226):227 if hasattr(encoder_output, "latent_dist") and sample_mode == "sample":228 return encoder_output.latent_dist.sample(generator)229 elif hasattr(encoder_output, "latent_dist") and sample_mode == "argmax":230 return encoder_output.latent_dist.mode()231 elif hasattr(encoder_output, "latents"):232 return encoder_output.latents233 else:234 raise AttributeError("Could not access latents of provided encoder_output")235 236 237# Modified from diffusers.modular_pipelines.qwenimage.encoders.encode_vae_image238def encode_vae_image(239 image: torch.Tensor,240 vae: AutoencoderKLQwenImage,241 generator: torch.Generator,242 device: torch.device,243 dtype: torch.dtype,244 latent_channels: int = 16,245 sample_mode: str = "argmax",246):247 if not isinstance(image, torch.Tensor):248 raise ValueError(f"Expected image to be a tensor, got {type(image)}.")249 250 # preprocessed image should be a 4D tensor: batch_size, num_channels, height, width251 if image.dim() == 4:252 image = image.unsqueeze(2)253 elif image.dim() != 5:254 raise ValueError(f"Expected image dims 4 or 5, got {image.dim()}.")255 256 image = image.to(device=device, dtype=dtype)257 258 if isinstance(generator, list):259 image_latents = [260 retrieve_latents(vae.encode(image[i : i + 1]), generator=generator[i], sample_mode=sample_mode)261 for i in range(image.shape[0])262 ]263 image_latents = torch.cat(image_latents, dim=0)264 else:265 image_latents = retrieve_latents(vae.encode(image), generator=generator, sample_mode=sample_mode)266 latents_mean = (267 torch.tensor(vae.config.latents_mean)268 .view(1, latent_channels, 1, 1, 1)269 .to(image_latents.device, image_latents.dtype)270 )271 latents_std = (272 torch.tensor(vae.config.latents_std)273 .view(1, latent_channels, 1, 1, 1)274 .to(image_latents.device, image_latents.dtype)275 )276 image_latents = (image_latents - latents_mean) / latents_std277 278 return image_latents279 280 281def encode_reference_latents(282 images: list[torch.Tensor],283 vae: AutoencoderKLQwenImage,284 max_pixels: int,285 generator: torch.Generator | None,286 device: torch.device,287 vae_scale_factor: int,288 patch_size: int,289 latent_channels: int = 16,290) -> list[torch.Tensor]:291 """Encode `[0, 1]` CHW reference images to normalized VAE latents, one `(C, h, w)` tensor per image. Each image292 is downscaled (aspect-preserving, never upscaled) to fit within `max_pixels`, then snapped so the latent grid is293 patchifiable. References keep their own aspect ratio, independent of the generated output size."""294 snap = vae_scale_factor * patch_size295 vae_dtype = vae.dtype296 297 latents_mean = torch.tensor(vae.config.latents_mean).view(1, latent_channels, 1, 1, 1)298 latents_std = torch.tensor(vae.config.latents_std).view(1, latent_channels, 1, 1, 1)299 300 ref_latents = []301 for img in images:302 img = img.unsqueeze(0).to(device, dtype=vae_dtype)303 h, w = img.shape[2], img.shape[3]304 if h * w > max_pixels:305 ratio = h / w306 new_h, new_w = math.sqrt(max_pixels * ratio), math.sqrt(max_pixels / ratio)307 else:308 new_h, new_w = float(h), float(w)309 new_h = max(snap, int(round(new_h / snap)) * snap)310 new_w = max(snap, int(round(new_w / snap)) * snap)311 if (new_h, new_w) != (h, w):312 img = F.interpolate(img.float(), size=(new_h, new_w), mode="bilinear").to(vae_dtype)313 314 img = (img * 2.0 - 1.0).unsqueeze(2) # [0, 1] -> [-1, 1], add frame dim315 latent = retrieve_latents(vae.encode(img), generator=generator, sample_mode="sample")316 latent = (latent - latents_mean.to(latent.device, latent.dtype)) / latents_std.to(latent.device, latent.dtype)317 ref_latents.append(latent[:, :, 0][0]) # drop frame + batch dims -> (C, h, w)318 return ref_latents319 320 321def pack_reference_latents(322 ref_latents: list[torch.Tensor],323 pachifier: Krea2Pachifier,324 device: torch.device,325 dtype: torch.dtype,326) -> tuple[torch.Tensor, torch.Tensor]:327 """Patchify reference latents into `(1, ref_seq_len, C * p * p)` tokens and build their `(ref_seq_len, 3)` rotary328 coordinates. The i-th reference sits on frame axis `i + 1` with its own y/x grid starting at 0 (the Kontext-style329 "index" placement that marks each reference as a distinct image rather than more of the canvas)."""330 p = pachifier.config.patch_size331 tokens, position_ids = [], []332 for i, ref in enumerate(ref_latents):333 ref = ref.unsqueeze(0).to(device, dtype)334 tokens.append(pachifier.pack_latents(ref))335 _, _, h, w = ref.shape336 ids = torch.zeros(h // p, w // p, 3, device=device)337 ids[..., 0] = i + 1338 ids[..., 1] = torch.arange(h // p, device=device)[:, None]339 ids[..., 2] = torch.arange(w // p, device=device)[None, :]340 position_ids.append(ids.reshape(-1, 3))341 return torch.cat(tokens, dim=1), torch.cat(position_ids, dim=0)342 343 344# ====================345# 1. TEXT ENCODER346# ====================347 348 349class Krea2TextEncoderStep(ModularPipelineBlocks):350 model_name = "krea2"351 352 def __init__(self, text_encoder_select_layers: tuple[int, ...] | None = None):353 """Text encoder step for Krea 2.354 355 Args:356 text_encoder_select_layers (`tuple[int, ...]`, *optional*):357 Indices into the text encoder's `hidden_states` tuple (0 is the embedding output) whose states are358 stacked per token as the transformer's text conditioning. Must have359 `transformer.config.num_text_layers` entries. Defaults to the Krea 2 (Qwen3-VL-4B) taps.360 """361 if text_encoder_select_layers is None:362 text_encoder_select_layers = KREA2_TEXT_ENCODER_SELECT_LAYERS363 self.text_encoder_select_layers = tuple(text_encoder_select_layers)364 super().__init__()365 366 @property367 def description(self) -> str:368 return "Text Encoder step that generates text embeddings to guide the image generation."369 370 @property371 def expected_components(self) -> list[ComponentSpec]:372 return [373 ComponentSpec("text_encoder", Qwen3VLModel, description="The text encoder to use"),374 ComponentSpec("tokenizer", Qwen2Tokenizer, description="The tokenizer to use"),375 ComponentSpec(376 "guider",377 ClassifierFreeGuidance,378 config=FrozenDict({"guidance_scale": 4.5, "use_original_formulation": True}),379 default_creation_method="from_config",380 ),381 ]382 383 @property384 def inputs(self) -> list[InputParam]:385 return [386 InputParam.template("prompt"),387 InputParam.template("negative_prompt"),388 InputParam.template("max_sequence_length"),389 ]390 391 @property392 def intermediate_outputs(self) -> list[OutputParam]:393 return [394 OutputParam.template("prompt_embeds"),395 OutputParam.template("prompt_embeds_mask"),396 OutputParam.template("negative_prompt_embeds"),397 OutputParam.template("negative_prompt_embeds_mask"),398 ]399 400 @staticmethod401 def check_inputs(prompt, negative_prompt, max_sequence_length):402 if not isinstance(prompt, str) and not isinstance(prompt, list):403 raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")404 405 if (406 negative_prompt is not None407 and not isinstance(negative_prompt, str)408 and not isinstance(negative_prompt, list)409 ):410 raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}")411 412 if max_sequence_length is not None and max_sequence_length <= 0:413 raise ValueError(f"`max_sequence_length` must be a positive integer but is {max_sequence_length}")414 415 @torch.no_grad()416 def __call__(self, components: Krea2ModularPipeline, state: PipelineState):417 block_state = self.get_block_state(state)418 419 device = components._execution_device420 self.check_inputs(block_state.prompt, block_state.negative_prompt, block_state.max_sequence_length)421 422 block_state.prompt_embeds, block_state.prompt_embeds_mask = get_krea2_prompt_embeds(423 components.text_encoder,424 components.tokenizer,425 prompt=block_state.prompt,426 text_encoder_select_layers=self.text_encoder_select_layers,427 max_sequence_length=block_state.max_sequence_length,428 device=device,429 )430 431 block_state.negative_prompt_embeds = None432 block_state.negative_prompt_embeds_mask = None433 if components.requires_unconditional_embeds:434 negative_prompt = block_state.negative_prompt or ""435 block_state.negative_prompt_embeds, block_state.negative_prompt_embeds_mask = get_krea2_prompt_embeds(436 components.text_encoder,437 components.tokenizer,438 prompt=negative_prompt,439 text_encoder_select_layers=self.text_encoder_select_layers,440 max_sequence_length=block_state.max_sequence_length,441 device=device,442 )443 444 self.set_block_state(state, block_state)445 return components, state446 447 448class Krea2EditTextEncoderStep(ModularPipelineBlocks):449 model_name = "krea2"450 451 def __init__(self, text_encoder_select_layers: tuple[int, ...] | None = None):452 """Text encoder step for the Krea 2 edit task: encodes the prompt while embedding a coarse view of the453 reference image(s) into the conditioning through the Qwen3-VL vision tower.454 455 Args:456 text_encoder_select_layers (`tuple[int, ...]`, *optional*):457 Indices into the text encoder's `hidden_states` tuple whose states are stacked per token as the458 transformer's text conditioning. Defaults to the Krea 2 (Qwen3-VL-4B) taps.459 """460 if text_encoder_select_layers is None:461 text_encoder_select_layers = KREA2_TEXT_ENCODER_SELECT_LAYERS462 self.text_encoder_select_layers = tuple(text_encoder_select_layers)463 super().__init__()464 465 @property466 def description(self) -> str:467 return (468 "Text encoder step for the edit task. Embeds reference image(s) into the text conditioning via the "469 "Qwen3-VL vision tower, matching how the Ostris AI-Toolkit edit LoRAs are trained."470 )471 472 @property473 def expected_components(self) -> list[ComponentSpec]:474 return [475 ComponentSpec("text_encoder", Qwen3VLModel, description="The text encoder to use"),476 ComponentSpec("tokenizer", Qwen2Tokenizer, description="The tokenizer to use"),477 ComponentSpec("processor", Qwen3VLProcessor, description="The Qwen3-VL processor for reference images"),478 ComponentSpec(479 "guider",480 ClassifierFreeGuidance,481 config=FrozenDict({"guidance_scale": 4.5, "use_original_formulation": True}),482 default_creation_method="from_config",483 ),484 ]485 486 @property487 def inputs(self) -> list[InputParam]:488 return [489 InputParam.template("prompt"),490 InputParam.template("negative_prompt"),491 InputParam.template("image", required=True, note="The reference image(s) for the edit."),492 InputParam.template("max_sequence_length"),493 InputParam(494 "vl_image_max_pixels",495 type_hint=int,496 default=384 * 384,497 description="Pixel budget for the coarse Qwen3-VL view of each reference image.",498 ),499 ]500 501 @property502 def intermediate_outputs(self) -> list[OutputParam]:503 return [504 OutputParam.template("prompt_embeds"),505 OutputParam.template("prompt_embeds_mask"),506 OutputParam.template("negative_prompt_embeds"),507 OutputParam.template("negative_prompt_embeds_mask"),508 ]509 510 @staticmethod511 def check_inputs(prompt, negative_prompt, max_sequence_length):512 if not isinstance(prompt, str) and not isinstance(prompt, list):513 raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")514 if (515 negative_prompt is not None516 and not isinstance(negative_prompt, str)517 and not isinstance(negative_prompt, list)518 ):519 raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}")520 if max_sequence_length is not None and max_sequence_length <= 0:521 raise ValueError(f"`max_sequence_length` must be a positive integer but is {max_sequence_length}")522 523 @torch.no_grad()524 def __call__(self, components: Krea2ModularPipeline, state: PipelineState):525 block_state = self.get_block_state(state)526 527 device = components._execution_device528 self.check_inputs(block_state.prompt, block_state.negative_prompt, block_state.max_sequence_length)529 530 image_list = block_state.image if isinstance(block_state.image, (list, tuple)) else [block_state.image]531 ref_images = [to_chw_tensor(img).to(device) for img in image_list]532 vl_images = prep_vl_images(ref_images, block_state.vl_image_max_pixels)533 534 block_state.prompt_embeds, block_state.prompt_embeds_mask = get_krea2_edit_prompt_embeds(535 components.text_encoder,536 components.tokenizer,537 components.processor,538 prompt=block_state.prompt,539 images=vl_images,540 text_encoder_select_layers=self.text_encoder_select_layers,541 max_sequence_length=block_state.max_sequence_length,542 device=device,543 )544 545 block_state.negative_prompt_embeds = None546 block_state.negative_prompt_embeds_mask = None547 if components.requires_unconditional_embeds:548 negative_prompt = block_state.negative_prompt or ""549 block_state.negative_prompt_embeds, block_state.negative_prompt_embeds_mask = get_krea2_edit_prompt_embeds(550 components.text_encoder,551 components.tokenizer,552 components.processor,553 prompt=negative_prompt,554 images=vl_images,555 text_encoder_select_layers=self.text_encoder_select_layers,556 max_sequence_length=block_state.max_sequence_length,557 device=device,558 )559 560 self.set_block_state(state, block_state)561 return components, state562 563 564# ====================565# 2. IMAGE PREPROCESS566# ====================567 568 569class Krea2InpaintProcessImagesInputStep(ModularPipelineBlocks):570 model_name = "krea2"571 572 @property573 def description(self) -> str:574 return "Image Preprocess step for inpainting task. This processes the image and mask inputs together. Images will be resized to the given height and width."575 576 @property577 def expected_components(self) -> list[ComponentSpec]:578 return [579 ComponentSpec(580 "image_mask_processor",581 InpaintProcessor,582 config=FrozenDict({"vae_scale_factor": 16}),583 default_creation_method="from_config",584 ),585 ]586 587 @property588 def inputs(self) -> list[InputParam]:589 return [590 InputParam.template("mask_image"),591 InputParam.template("image"),592 InputParam.template("height"),593 InputParam.template("width"),594 InputParam.template("padding_mask_crop"),595 ]596 597 @property598 def intermediate_outputs(self) -> list[OutputParam]:599 return [600 OutputParam(601 name="processed_image",602 type_hint=torch.Tensor,603 description="The processed image",604 ),605 OutputParam(606 name="processed_mask_image",607 type_hint=torch.Tensor,608 description="The processed mask image",609 ),610 OutputParam(611 name="mask_overlay_kwargs",612 type_hint=dict,613 description="The kwargs for the postprocess step to apply the mask overlay",614 ),615 ]616 617 @staticmethod618 def check_inputs(height, width, vae_scale_factor):619 if height is not None and height % (vae_scale_factor * 2) != 0:620 raise ValueError(f"Height must be divisible by {vae_scale_factor * 2} but is {height}")621 622 if width is not None and width % (vae_scale_factor * 2) != 0:623 raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}")624 625 @torch.no_grad()626 def __call__(self, components: Krea2ModularPipeline, state: PipelineState):627 block_state = self.get_block_state(state)628 629 self.check_inputs(630 height=block_state.height, width=block_state.width, vae_scale_factor=components.vae_scale_factor631 )632 height = block_state.height or components.default_height633 width = block_state.width or components.default_width634 635 block_state.processed_image, block_state.processed_mask_image, block_state.mask_overlay_kwargs = (636 components.image_mask_processor.preprocess(637 image=block_state.image,638 mask=block_state.mask_image,639 height=height,640 width=width,641 padding_mask_crop=block_state.padding_mask_crop,642 )643 )644 645 self.set_block_state(state, block_state)646 return components, state647 648 649class Krea2ProcessImagesInputStep(ModularPipelineBlocks):650 model_name = "krea2"651 652 @property653 def description(self) -> str:654 return "Image Preprocess step. will resize the image to the given height and width."655 656 @property657 def expected_components(self) -> list[ComponentSpec]:658 return [659 ComponentSpec(660 "image_processor",661 VaeImageProcessor,662 config=FrozenDict({"vae_scale_factor": 16}),663 default_creation_method="from_config",664 ),665 ]666 667 @property668 def inputs(self) -> list[InputParam]:669 return [670 InputParam.template("image"),671 InputParam.template("height"),672 InputParam.template("width"),673 ]674 675 @property676 def intermediate_outputs(self) -> list[OutputParam]:677 return [678 OutputParam(679 name="processed_image",680 type_hint=torch.Tensor,681 description="The processed image",682 )683 ]684 685 @staticmethod686 def check_inputs(height, width, vae_scale_factor):687 if height is not None and height % (vae_scale_factor * 2) != 0:688 raise ValueError(f"Height must be divisible by {vae_scale_factor * 2} but is {height}")689 690 if width is not None and width % (vae_scale_factor * 2) != 0:691 raise ValueError(f"Width must be divisible by {vae_scale_factor * 2} but is {width}")692 693 @torch.no_grad()694 def __call__(self, components: Krea2ModularPipeline, state: PipelineState):695 block_state = self.get_block_state(state)696 697 self.check_inputs(698 height=block_state.height, width=block_state.width, vae_scale_factor=components.vae_scale_factor699 )700 height = block_state.height or components.default_height701 width = block_state.width or components.default_width702 703 block_state.processed_image = components.image_processor.preprocess(704 image=block_state.image,705 height=height,706 width=width,707 )708 709 self.set_block_state(state, block_state)710 return components, state711 712 713# ====================714# 3. VAE ENCODER715# ====================716 717 718class Krea2VaeEncoderStep(ModularPipelineBlocks):719 model_name = "krea2"720 721 @property722 def description(self) -> str:723 return "VAE Encoder step that converts processed_image into latent representations image_latents."724 725 @property726 def expected_components(self) -> list[ComponentSpec]:727 return [ComponentSpec("vae", AutoencoderKLQwenImage)]728 729 @property730 def inputs(self) -> list[InputParam]:731 return [732 InputParam(733 name="processed_image", required=True, type_hint=torch.Tensor, description="The image tensor to encode"734 ),735 InputParam.template("generator"),736 ]737 738 @property739 def intermediate_outputs(self) -> list[OutputParam]:740 return [OutputParam.template("image_latents")]741 742 @torch.no_grad()743 def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:744 block_state = self.get_block_state(state)745 746 device = components._execution_device747 dtype = components.vae.dtype748 749 block_state.image_latents = encode_vae_image(750 image=block_state.processed_image,751 vae=components.vae,752 generator=block_state.generator,753 device=device,754 dtype=dtype,755 latent_channels=components.num_channels_latents,756 )757 758 self.set_block_state(state, block_state)759 760 return components, state761 762 763class Krea2EditReferenceLatentsStep(ModularPipelineBlocks):764 model_name = "krea2"765 766 @property767 def description(self) -> str:768 return (769 "Reference (edit) VAE encoder step. Encodes reference image(s) to clean, normalized VAE latents and packs "770 "them into transformer tokens with their frame-axis rotary coordinates. These tokens are appended to the "771 "sequence at flow time t=0 to condition the generation on the references."772 )773 774 @property775 def expected_components(self) -> list[ComponentSpec]:776 return [777 ComponentSpec("vae", AutoencoderKLQwenImage),778 ComponentSpec("pachifier", Krea2Pachifier, default_creation_method="from_config"),779 ]780 781 @property782 def inputs(self) -> list[InputParam]:783 return [784 InputParam.template("image", required=True, note="The reference image(s) for the edit."),785 InputParam.template("generator"),786 InputParam(787 "reference_max_pixels",788 type_hint=int,789 default=1024 * 1024,790 description="Pixel budget each reference image is downscaled to fit before VAE encoding.",791 ),792 ]793 794 @property795 def intermediate_outputs(self) -> list[OutputParam]:796 return [797 OutputParam(798 name="reference_latents",799 type_hint=torch.Tensor,800 description="Packed clean reference tokens of shape (1, ref_seq_len, C * p * p), appended to the "801 "denoiser sequence at t=0.",802 ),803 OutputParam(804 name="reference_position_ids",805 type_hint=torch.Tensor,806 description="Rotary coordinates (ref_seq_len, 3) for the reference tokens; the i-th reference sits on "807 "frame axis i + 1.",808 ),809 OutputParam(810 name="ref_seq_len",811 kwargs_type="denoiser_input_fields",812 type_hint=int,813 description="Number of reference tokens appended to the denoiser sequence.",814 ),815 ]816 817 @torch.no_grad()818 def __call__(self, components: Krea2ModularPipeline, state: PipelineState) -> PipelineState:819 block_state = self.get_block_state(state)820 821 device = components._execution_device822 patch_size = components.pachifier.config.patch_size823 824 image_list = block_state.image if isinstance(block_state.image, (list, tuple)) else [block_state.image]825 ref_images = [to_chw_tensor(img) for img in image_list]826 827 ref_latents = encode_reference_latents(828 images=ref_images,829 vae=components.vae,830 max_pixels=block_state.reference_max_pixels,831 generator=block_state.generator,832 device=device,833 vae_scale_factor=components.vae_scale_factor,834 patch_size=patch_size,835 latent_channels=components.num_channels_latents,836 )837 block_state.reference_latents, block_state.reference_position_ids = pack_reference_latents(838 ref_latents, components.pachifier, device, components.vae.dtype839 )840 block_state.ref_seq_len = block_state.reference_latents.shape[1]841 842 self.set_block_state(state, block_state)843 844 return components, state845 