Team Ai
Modelpublic

diffusers-modular/krea2-edit

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