diff --git a/extensions_built_in/diffusion_models/krea2/krea2.py b/extensions_built_in/diffusion_models/krea2/krea2.py index ea7fc37..75b6ae5 100644 --- a/extensions_built_in/diffusion_models/krea2/krea2.py +++ b/extensions_built_in/diffusion_models/krea2/krea2.py @@ -13,16 +13,21 @@ Flow-matching convention matches ai-toolkit exactly (t=1 noise -> t=0 clean, target = noise - clean), so ``get_noise_prediction`` does no time flip / negation. """ +import math import os -from typing import List, Optional +from typing import TYPE_CHECKING, List, Optional import torch +import torch.nn.functional as F +from PIL import Image +from torchvision.transforms.functional import to_tensor from safetensors.torch import load_file, save_file import huggingface_hub from huggingface_hub.errors import EntryNotFoundError from diffusers import AutoencoderKLQwenImage from transformers import ( + AutoProcessor, AutoTokenizer, Qwen2TokenizerFast, Qwen3VLForConditionalGeneration, @@ -51,6 +56,9 @@ from .src.mmdit import ( from .src.text_encoder import encode_krea_prompt, SELECT_LAYERS from .src.pipeline import Krea2Pipeline, pad_text_features, predict_velocity +if TYPE_CHECKING: + from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO + # The reference "single_mmdit_large_wide" architecture (oss_raw / oss_turbo share it). KREA2_MMDIT_CONFIG = dict( @@ -92,6 +100,30 @@ QWEN_IMAGE_VAE_PATH = "Qwen/Qwen-Image" HF_TOKEN = os.getenv("HF_TOKEN", None) +def patch_qwen_vl_patch_embed(model): + """Qwen-VL's vision patch_embed is a Conv3d whose kernel == stride, i.e. a plain + linear projection of each flattened patch. bf16 Conv3d has no fast cuDNN kernel and + falls back to a slow, GPU-underutilizing path. Swap it for the equivalent F.linear + (a GEMM). The weight is read lazily so this survives later .to(device)/dtype moves. + Returns the number of patch_embed modules patched. (Same patch as the + Qwen3VLCaptioner extension.)""" + patched = 0 + for module in model.modules(): + proj = getattr(module, "proj", None) + if isinstance(proj, torch.nn.Conv3d) and tuple(proj.kernel_size) == tuple( + proj.stride + ): + + def fast_forward(hidden_states, _proj=proj): + w = _proj.weight.reshape(_proj.weight.shape[0], -1) + x = hidden_states.view(-1, w.shape[1]).to(w.dtype) + return F.linear(x, w, _proj.bias) + + module.forward = fast_forward + patched += 1 + return patched + + def _load_mmdit_state_dict(name_or_path: str, filename: Optional[str]) -> dict: """Load the MMDiT weights from a local safetensors file/dir or the HF hub. @@ -160,8 +192,23 @@ class Krea2Model(BaseModel): # Qwen2TokenizerFast used to tokenize the assistant suffix (matches the # reference's separate processor pass). self.processor = None + # Qwen3-VL AutoProcessor for encoding reference images into the prompt. + self.vl_processor = None self.use_old_lokr_format = False + # Optional reference-image (edit) conditioning, enabled with + # model_kwargs.edit = true. Control images feed the model in two places: + # through the Qwen3-VL encoder alongside the prompt (edit-plus style, so + # the text embeddings see them) and as clean VAE latents appended to the + # image sequence at t=0 (ComfyUI Kontext "index_timestep_zero"). Runs in + # ComfyUI with the ComfyUI-Krea2-Ostris-Edit custom nodes. With edit off + # (the default) all of it is skipped and this is the plain T2I model. + self.is_edit = bool(self.model_config.model_kwargs.get("edit", False)) + self.encode_control_in_text_embeddings = self.is_edit + self.has_multiple_control_images = self.is_edit + # Reference images keep their own aspect/size (not resized to the target). + self.use_raw_control_images = self.is_edit + @staticmethod def get_train_scheduler(): return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config) @@ -214,14 +261,22 @@ class Krea2Model(BaseModel): text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( te_path, torch_dtype=dtype, token=HF_TOKEN ) - # We only ever encode text, so the vision tower is dead weight -- drop it to - # free VRAM and skip loading its (bf16-slow) Conv3d patch_embed onto the GPU. - if getattr(text_encoder.model, "visual", None) is not None: - text_encoder.model.visual = None + vl_processor = None + if self.is_edit: + # Edit mode: reference images are encoded into the text embeddings, + # so the vision tower stays. Swap its Conv3d patch_embed for an + # equivalent GEMM (bf16 Conv3d has no fast cuDNN kernel). + vl_processor = AutoProcessor.from_pretrained(te_path, token=HF_TOKEN) + patch_qwen_vl_patch_embed(text_encoder) + else: + # We only ever encode text, so the vision tower is dead weight -- drop it to + # free VRAM and skip loading its (bf16-slow) Conv3d patch_embed onto the GPU. + if getattr(text_encoder.model, "visual", None) is not None: + text_encoder.model.visual = None text_encoder.eval() text_encoder.requires_grad_(False) flush() - return tokenizer, processor, text_encoder + return tokenizer, processor, vl_processor, text_encoder def _load_vae(self): vae_path = self.model_config.model_kwargs.get("vae_path", QWEN_IMAGE_VAE_PATH) @@ -355,7 +410,7 @@ class Krea2Model(BaseModel): transformer.to(self.device_torch, dtype=dtype) flush() - tokenizer, processor, text_encoder = self._load_text_encoder() + tokenizer, processor, vl_processor, text_encoder = self._load_text_encoder() if self.model_config.quantize_te: self.print_and_status_update("Quantizing text encoder") text_encoder.to(self.device_torch) @@ -388,6 +443,7 @@ class Krea2Model(BaseModel): self.text_encoder = text_encoder self.tokenizer = tokenizer self.processor = processor + self.vl_processor = vl_processor self.model = transformer self.pipeline = Krea2Pipeline(self) self.print_and_status_update("Model Loaded") @@ -414,6 +470,31 @@ class Krea2Model(BaseModel): gen_config.width = int(gen_config.width // sc * sc) gen_config.height = int(gen_config.height // sc * sc) + # Reference image(s) -> clean VAE latents for the t=0 sequence tokens. + # The Qwen3-VL side already saw them (baked into the prompt embeds). + # ctrl_img_1 mirrors ctrl_img when unset, so use one or the other. + ctrl_paths = [] + if self.is_edit: + if gen_config.ctrl_img is not None: + ctrl_paths.append(gen_config.ctrl_img) + elif gen_config.ctrl_img_1 is not None: + ctrl_paths.append(gen_config.ctrl_img_1) + if gen_config.ctrl_img_2 is not None: + ctrl_paths.append(gen_config.ctrl_img_2) + if gen_config.ctrl_img_3 is not None: + ctrl_paths.append(gen_config.ctrl_img_3) + + ref_latents = None + if ctrl_paths: + ctrl_tensors = [ + to_tensor(Image.open(path).convert("RGB")) for path in ctrl_paths + ] + target_pixels = gen_config.width * gen_config.height + # one batch item (preview batch size is 1) -> List[List[(16, h, w)]] + ref_latents = [ + self._encode_ref_latents(ctrl_tensors, target_pixels=target_pixels) + ] + img = pipeline( conditional_embeds=conditional_embeds, unconditional_embeds=unconditional_embeds, @@ -423,9 +504,95 @@ class Krea2Model(BaseModel): guidance_scale=gen_config.guidance_scale, latents=gen_config.latents, generator=generator, + ref_latents=ref_latents, )[0] return img + # ------------------------------------------------------------------ + # Reference-image helpers + # ------------------------------------------------------------------ + def _ref_target_pixels(self, target_pixels: Optional[int]) -> int: + """Pixel budget each reference image is resized to fit within. + + - default: ``control_image_max_pixels`` model_kwarg (1 MP) -- a hard cap + so raw, full-size control images don't blow up the token count / VRAM. + - ``match_target_res`` model_kwarg: use the target generation area instead. + """ + max_pixels = int( + self.model_config.model_kwargs.get("control_image_max_pixels", 1024 * 1024) + ) + if ( + self.model_config.model_kwargs.get("match_target_res", False) + and target_pixels + ): + return int(target_pixels) + return max_pixels + + def _encode_ref_latents( + self, control_tensors, target_pixels: Optional[int] = None + ) -> List[torch.Tensor]: + """Encode ``[0, 1]`` reference image tensors to VAE latents. + + Returns a list of ``(16, h, w)`` latents (one per reference image). Each + control image is resized so its area fits within the pixel budget (see + ``_ref_target_pixels``) -- preserving aspect ratio -- then snapped so the + latent grid is divisible by the patch size. ``control_tensors`` is a list + of ``(C, H, W)`` or ``(1, C, H, W)`` tensors in ``[0, 1]``. + """ + sc = self.get_bucket_divisibility() # 16: VAE(8) * patch(2) + budget = self._ref_target_pixels(target_pixels) + match = self.model_config.model_kwargs.get("match_target_res", False) + + latents = [] + for img in control_tensors: + if img.dim() == 3: + img = img.unsqueeze(0) + img = img.to(self.device_torch, dtype=self.torch_dtype) + + h, w = img.shape[2], img.shape[3] + # match_target_res: scale area *to* the budget; otherwise only scale + # *down* when the image is larger than the budget. + area = h * w + if match or area > budget: + ratio = h / w + new_h = math.sqrt(budget * ratio) + new_w = new_h / ratio + else: + new_h, new_w = float(h), float(w) + + # snap to a multiple of the bucket divisibility so the VAE latent grid + # is patchifiable (the transformer rearranges 2x2 latent patches). + new_h = max(sc, int(round(new_h / sc)) * sc) + new_w = max(sc, int(round(new_w / sc)) * sc) + if (new_h, new_w) != (h, w): + img = F.interpolate(img, size=(new_h, new_w), mode="bilinear") + + # encode_images expects [-1, 1]; control tensors arrive in [0, 1]. + latent = self.encode_images( + img * 2 - 1, device=self.device_torch, dtype=self.torch_dtype + ) + latents.append(latent[0]) # drop batch dim -> (16, h, w) + return latents + + def _batch_ref_latents_from_batch( + self, + batch: "DataLoaderBatchDTO", + batch_size: int, + target_pixels: Optional[int] = None, + ) -> Optional[List[List[torch.Tensor]]]: + """Build predict_velocity's ``ref_latents`` from a train batch.""" + control_list = batch.control_tensor_list + if control_list is None and batch.control_tensor is not None: + control_list = [batch.control_tensor[b : b + 1] for b in range(batch_size)] + if control_list is None: + return None + if len(control_list) != batch_size: + raise ValueError("Control tensor list length does not match batch size") + return [ + self._encode_ref_latents(controls, target_pixels=target_pixels) + for controls in control_list + ] + # ------------------------------------------------------------------ # Training hooks # ------------------------------------------------------------------ @@ -434,11 +601,25 @@ class Krea2Model(BaseModel): latent_model_input: torch.Tensor, # (B, 16, h, w) timestep: torch.Tensor, # 0..1000 scale text_embeddings: AdvancedPromptEmbeds, + batch: "DataLoaderBatchDTO" = None, **kwargs, ): if self.model.device == torch.device("cpu"): self.model.to(self.device_torch) + # Clean reference latents from the batch's control images (if any); they + # ride along in the sequence at t=0 and are never noised. + ref_latents = None + if batch is not None and self.is_edit: + with torch.no_grad(): + _, _, lh, lw = latent_model_input.shape + target_pixels = (lh * self.vae_scale_factor) * ( + lw * self.vae_scale_factor + ) + ref_latents = self._batch_ref_latents_from_batch( + batch, latent_model_input.shape[0], target_pixels=target_pixels + ) + # toolkit timestep (0..1000, 1000 = pure noise) -> Krea flow time t in # [0, 1] with t=1 = pure noise. Same convention -> straight divide. t = timestep.to(self.device_torch, dtype=torch.float32) / 1000.0 @@ -457,16 +638,69 @@ class Krea2Model(BaseModel): t, context, text_mask, + ref_latents=ref_latents, ) return pred - def get_prompt_embeds(self, prompt) -> AdvancedPromptEmbeds: + def _prep_vlm_images(self, ctrl: List[torch.Tensor]) -> List[torch.Tensor]: + """Resize reference images for the Qwen3-VL pass. + + Downscaled (aspect-preserved, never upscaled) to fit ``vlm_max_pixels`` + total area (384^2 by default, the boogu_image_edit / ComfyUI + TextEncodeQwenImageEditPlus budget) -- the MLLM only needs a coarse + understanding of the reference; high-res detail flows through the VAE + ref latents. + """ + target = int(self.model_config.model_kwargs.get("vlm_max_pixels", 384 * 384)) + images = [] + for img in ctrl: + if img.dim() == 4: + img = img[0] + img = img.to(self.device_torch) + h, w = img.shape[1], img.shape[2] + scale = min(1.0, math.sqrt(target / (h * w))) + nh, nw = max(round(h * scale), 28), max(round(w * scale), 28) + if (nh, nw) != (h, w): + img = ( + F.interpolate( + img.unsqueeze(0).float(), + size=(nh, nw), + mode="bicubic", + antialias=True, + ) + .squeeze(0) + .clamp(0, 1) + ) + images.append(img.float()) + return images + + def get_prompt_embeds(self, prompt, control_images=None) -> AdvancedPromptEmbeds: if isinstance(prompt, str): prompt = [prompt] if self.text_encoder.device == torch.device("cpu"): self.text_encoder.to(self.device_torch) + # Normalize control images to a per-prompt list (List[List[Tensor]]). + # They arrive as a (B, C, H, W) batch tensor (control_tensor), a list of + # per-sample lists (control_tensor_list), or a flat list of (1, C, H, W) + # tensors for a single prompt (sampling / blank-embed caching). + if control_images is not None: + if isinstance(control_images, torch.Tensor): + control_images = [ + [control_images[i]] for i in range(control_images.shape[0]) + ] + elif len(control_images) > 0 and not isinstance(control_images[0], list): + control_images = [control_images] + if len(control_images) == 1 and len(prompt) > 1: + control_images = control_images * len(prompt) + if len(control_images) != len(prompt): + raise ValueError( + "Number of prompts must match number of control image sets" + ) + else: + control_images = [None] * len(prompt) + # Encode each prompt at its natural length and store one (L, 12*2560) # tensor per batch item. The (L, 12, 2560) stack is flattened to 2D so the # toolkit's batching reads the list length (not the seq length) as the @@ -474,7 +708,8 @@ class Krea2Model(BaseModel): # batch max is deferred to the model call so caches stay small and any # prompts can share a batch. features_list = [] - for p in prompt: + for p, ctrl in zip(prompt, control_images): + images = self._prep_vlm_images(ctrl) if ctrl is not None else None features = encode_krea_prompt( self.text_encoder, self.tokenizer, @@ -482,6 +717,9 @@ class Krea2Model(BaseModel): p, max_length=self.max_text_length, select_layers=SELECT_LAYERS, + images=images, + vl_processor=self.vl_processor, + dtype=self.torch_dtype, ) # (L, n, d) -> (L, n*d) features = features.reshape(features.shape[0], -1) @@ -577,6 +815,7 @@ class Krea2Model(BaseModel): # ------------------------------------------------------------------ def save_model(self, output_path, meta, save_dtype): from toolkit.util.quantize import dequantize_if_quantized + if not output_path.endswith(".safetensors"): output_path = output_path + ".safetensors" transformer: SingleStreamDiT = unwrap_model(self.model) @@ -584,7 +823,9 @@ class Krea2Model(BaseModel): save_dict = {} for k, v in state_dict.items(): # dequantize any quantized (e.g. quanto/torchao) weights so we save plain full precision tensors - save_dict[k] = dequantize_if_quantized(v).clone().to("cpu", dtype=save_dtype) + save_dict[k] = ( + dequantize_if_quantized(v).clone().to("cpu", dtype=save_dtype) + ) meta = get_meta_for_safetensors(meta, name="krea2") save_file(save_dict, output_path, metadata=meta) diff --git a/extensions_built_in/diffusion_models/krea2/src/mmdit.py b/extensions_built_in/diffusion_models/krea2/src/mmdit.py index e40710f..014919b 100644 --- a/extensions_built_in/diffusion_models/krea2/src/mmdit.py +++ b/extensions_built_in/diffusion_models/krea2/src/mmdit.py @@ -328,6 +328,33 @@ class SingleStreamBlock(nn.Module): def forward( self, x: Tensor, vec: Tensor, freqs: Tensor, mask: Tensor | None = None ) -> Tensor: + # ``vec`` is the (B, 1, 6*features) modulation input, or a tuple + # ``(vec, refvec, split)`` for reference-image conditioning: tokens + # ``[:split]`` (text + noisy image) are modulated with ``vec`` while + # tokens ``[split:]`` (clean reference tokens) use ``refvec`` built from + # t=0 (ComfyUI Kontext "index_timestep_zero"). Applied per span rather + # than materializing a per-token (B, L, 6*features) tensor. + if isinstance(vec, tuple): + vec, refvec, split = vec + m = self.mod(vec) + r = self.mod(refvec) + + def mod(h, scale, shift): + return torch.cat( + ( + (1 + m[scale]) * h[:, :split] + m[shift], + (1 + r[scale]) * h[:, split:] + r[shift], + ), + dim=1, + ) + + def gate(h, g): + return torch.cat((m[g] * h[:, :split], r[g] * h[:, split:]), dim=1) + + x = x + gate(self.attn(mod(self.prenorm(x), 0, 1), freqs, mask), 2) + x = x + gate(self.mlp(mod(self.postnorm(x), 3, 4)), 5) + return x + prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec) x = x + pregate * self.attn( (1 + prescale) * self.prenorm(x) + preshift, freqs, mask @@ -417,6 +444,7 @@ class SingleStreamDiT(nn.Module): t: Tensor, pos: Tensor, mask: Tensor | None = None, + reflen: int = 0, ) -> Tensor: img = self.first(img) t = self.tmlp(temb(t, self.config.tdim, device=img.device, dtype=img.dtype)) @@ -438,6 +466,23 @@ class SingleStreamDiT(nn.Module): mask = F.pad(mask, (0, _padlen), value=False) pos = F.pad(pos, (0, 0, 0, _padlen)) + blockvec = tvec + if reflen > 0: + # The last ``reflen`` image tokens are clean reference tokens: they + # get t=0 modulation (ComfyUI Kontext "index_timestep_zero") while + # text + noisy image tokens keep the real t. Padding tokens fall in + # the t=0 span, but they are masked from attention and sliced off + # the output, so their values never matter. + t0 = self.tmlp( + temb( + torch.zeros_like(t[:, 0, 0]), + self.config.tdim, + device=img.device, + dtype=img.dtype, + ) + ) + blockvec = (tvec, self.tproj(t0), txtlen + imglen - reflen) + mask = _mask(mask) freqs = self.posemb(pos) @@ -447,15 +492,15 @@ class SingleStreamDiT(nn.Module): combined = checkpoint( block, combined, - tvec, + blockvec, freqs, mask, use_reentrant=False, ) else: - combined = block(combined, tvec, freqs, mask) + combined = block(combined, blockvec, freqs, mask) final = self.last(combined, t) - output = final[:, txtlen : txtlen + imglen, :] + output = final[:, txtlen : txtlen + imglen - reflen, :] return output diff --git a/extensions_built_in/diffusion_models/krea2/src/pipeline.py b/extensions_built_in/diffusion_models/krea2/src/pipeline.py index 0345c58..6566d2b 100644 --- a/extensions_built_in/diffusion_models/krea2/src/pipeline.py +++ b/extensions_built_in/diffusion_models/krea2/src/pipeline.py @@ -90,20 +90,78 @@ def prepare( return img, pos, mask +def pack_ref_latents( + ref_latents: List[List[torch.Tensor]], patch: int, device, dtype +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Patchify per-sample reference latents into padded ref tokens / pos / mask. + + ``ref_latents`` is a list (one entry per batch item) of lists of ``(C, h, w)`` + reference latents. The i-th reference of a sample is placed on RoPE axis 0 at + index ``i + 1`` with its own y/x grid starting at 0 -- the ComfyUI Kontext + "index" placement (axis 0 is otherwise always 0, so the base weights see the + references as a new "frame" axis). Samples with fewer reference tokens are + right-padded and masked out. Returns ``(tokens (B, Lr, C*p*p), + pos (B, Lr, 3), mask (B, Lr))``. + """ + token_dim = None + seqs, ids = [], [] + for refs in ref_latents: + toks, rpos = [], [] + for i, ref in enumerate(refs): + ref = ref.to(device, dtype) + _, h, w = ref.shape + h_, w_ = h // patch, w // patch + toks.append( + rearrange(ref, "c (h ph) (w pw) -> (h w) (c ph pw)", ph=patch, pw=patch) + ) + token_dim = toks[-1].shape[-1] + refids = torch.zeros((h_, w_, 3), device=device) + refids[..., 0] = i + 1 + refids[..., 1] = torch.arange(h_, device=device)[:, None] + refids[..., 2] = torch.arange(w_, device=device)[None, :] + rpos.append(refids.reshape(-1, 3)) + seqs.append(toks) + ids.append(rpos) + + b = len(seqs) + # a sample may have no references (its span is fully masked/padded) + seqs = [ + torch.cat(t, dim=0) + if t + else torch.zeros(0, token_dim, device=device, dtype=dtype) + for t in seqs + ] + ids = [torch.cat(p, dim=0) if p else torch.zeros(0, 3, device=device) for p in ids] + max_len = max(s.shape[0] for s in seqs) + tokens = torch.zeros(b, max_len, token_dim, device=device, dtype=dtype) + pos = torch.zeros(b, max_len, 3, device=device) + mask = torch.zeros(b, max_len, device=device, dtype=torch.bool) + for i, (s, p) in enumerate(zip(seqs, ids)): + ln = s.shape[0] + tokens[i, :ln] = s + pos[i, :ln] = p + mask[i, :ln] = True + return tokens, pos, mask + + def predict_velocity( model: SingleStreamDiT, latents: torch.Tensor, # (B, C, h, w) t: torch.Tensor, # (B,) flow time in [0, 1] (1 = pure noise) context: torch.Tensor, # (B, Lt, n*d) flattened stacked Qwen3-VL features text_mask: torch.Tensor, # (B, Lt) 1 for real text tokens + ref_latents: Optional[List[List[torch.Tensor]]] = None, # per-sample (C, h, w) refs ) -> torch.Tensor: - """Run the MMDiT on the packed [text | image] sequence. + """Run the MMDiT on the packed [text | image | refs] sequence. ``latents`` stay in the unpacked ``(B, C, h, w)`` latent layout; image-token packing is internal to this function. ``context`` arrives 2D-per-sample flattened ``(B, Lt, n*d)`` and is restored to ``(B, Lt, n, d)`` for the MMDiT. - Returns the velocity ``noise - clean`` reshaped back to ``(B, C, h, w)``. No - time flip / negation: Krea's convention matches toolkit's. + ``ref_latents`` (optional) are clean reference latents appended after the + image tokens and conditioned at t=0 ("index_timestep_zero"); the prediction + only ever covers the noisy target tokens. Returns the velocity + ``noise - clean`` reshaped back to ``(B, C, h, w)``. No time flip / negation: + Krea's convention matches toolkit's. """ patch = model.config.patch b, c, h, w = latents.shape @@ -116,7 +174,17 @@ def predict_velocity( img_tokens, pos, mask = prepare(latents, context.shape[1], patch, text_mask) - out = model(img=img_tokens, context=context, t=t, pos=pos, mask=mask) + reflen = 0 + if ref_latents is not None and any(len(r) > 0 for r in ref_latents): + ref_tokens, ref_pos, ref_mask = pack_ref_latents( + ref_latents, patch, img_tokens.device, img_tokens.dtype + ) + reflen = ref_tokens.shape[1] + img_tokens = torch.cat((img_tokens, ref_tokens), dim=1) + pos = torch.cat((pos, ref_pos), dim=1) + mask = torch.cat((mask, ref_mask), dim=1) + + out = model(img=img_tokens, context=context, t=t, pos=pos, mask=mask, reflen=reflen) # (B, imglen, c*p*p) -> (B, c, h, w) velocity = rearrange( @@ -193,6 +261,7 @@ class Krea2Pipeline: guidance_scale: float = 4.5, latents: Optional[torch.Tensor] = None, generator: Optional[torch.Generator] = None, + ref_latents: Optional[List[List[torch.Tensor]]] = None, **kwargs, ) -> List[Image.Image]: model = self.model @@ -242,11 +311,21 @@ class Krea2Pipeline: for tcurr, tprev in zip(ts[:-1], ts[1:]): t = torch.full((latents.shape[0],), tcurr, dtype=dtype, device=device) v_cond = predict_velocity( - transformer, latents.to(dtype), t, cond_feats, cond_mask + transformer, + latents.to(dtype), + t, + cond_feats, + cond_mask, + ref_latents=ref_latents, ) if do_cfg: v_uncond = predict_velocity( - transformer, latents.to(dtype), t, uncond_feats, uncond_mask + transformer, + latents.to(dtype), + t, + uncond_feats, + uncond_mask, + ref_latents=ref_latents, ) v = v_cond + guidance_scale * (v_cond - v_uncond) else: diff --git a/extensions_built_in/diffusion_models/krea2/src/text_encoder.py b/extensions_built_in/diffusion_models/krea2/src/text_encoder.py index 3baa445..70230aa 100644 --- a/extensions_built_in/diffusion_models/krea2/src/text_encoder.py +++ b/extensions_built_in/diffusion_models/krea2/src/text_encoder.py @@ -13,6 +13,8 @@ its hidden states are sliced off the returned features, exactly like the reference. """ +from typing import List, Optional + import torch from torch import Tensor @@ -43,6 +45,9 @@ def encode_krea_prompt( max_length: int = 512, select_layers: tuple[int, ...] = SELECT_LAYERS, prefix_idx: int = PROMPT_TEMPLATE_ENCODE_START_IDX, + images: Optional[List[Tensor]] = None, + vl_processor=None, + dtype: Optional[torch.dtype] = None, ) -> Tensor: """Encode a single prompt into stacked Qwen3-VL hidden states. @@ -50,6 +55,15 @@ def encode_krea_prompt( dtype) holding the prompt + suffix token features -- the system prefix has been sliced off. ``L`` is the natural (unpadded) length so the caller stores one tensor per prompt and pads to the batch max at the model call. + + ``images`` (optional) are reference images -- ``(C, H, W)`` tensors in + ``[0, 1]`` -- embedded in the user message ahead of the prompt via named + vision placeholders (``Picture 1: <|vision_start|><|image_pad|><|vision_end|>``, + the ComfyUI ``TextEncodeQwenImageEditPlus`` layout). The ``vl_processor`` + (Qwen3-VL AutoProcessor) expands each ``<|image_pad|>`` to the image's token + grid, so the returned features carry the vision tokens as extra conditioning. + The system prefix is unchanged, so ``prefix_idx`` slicing stays valid and the + image + prompt tokens all survive the slice. """ device = qwen.device @@ -61,21 +75,59 @@ def encode_krea_prompt( suffix_ids = suffix_inputs["input_ids"] suffix_mask = suffix_inputs["attention_mask"].bool() - # Prefix + prompt at natural length (no padding); truncate very long prompts. - text = PROMPT_TEMPLATE_ENCODE_PREFIX + prompt - inputs = tokenizer( - [text], - truncation=True, - return_length=False, - return_overflowing_tokens=False, - max_length=max_length + prefix_idx, - return_tensors="pt", - ).to(device, non_blocking=True) + extra_inputs = {} + if images is not None and len(images) > 0: + image_prompt = "".join( + f"Picture {i + 1}: <|vision_start|><|image_pad|><|vision_end|>" + for i in range(len(images)) + ) + text = PROMPT_TEMPLATE_ENCODE_PREFIX + image_prompt + prompt + # No truncation here: the expanded image-pad runs must stay intact. + inputs = vl_processor( + text=[text], + images=list(images), + return_tensors="pt", + do_rescale=False, + ).to(device) + for k, v in inputs.items(): + if k in ("input_ids", "attention_mask"): + continue + if ( + isinstance(v, torch.Tensor) + and v.is_floating_point() + and dtype is not None + ): + v = v.to(dtype) + extra_inputs[k] = v + else: + # Prefix + prompt at natural length (no padding); truncate very long prompts. + text = PROMPT_TEMPLATE_ENCODE_PREFIX + prompt + inputs = tokenizer( + [text], + truncation=True, + return_length=False, + return_overflowing_tokens=False, + max_length=max_length + prefix_idx, + return_tensors="pt", + ).to(device, non_blocking=True) input_ids = torch.cat([inputs["input_ids"], suffix_ids], dim=1) mask = torch.cat([inputs["attention_mask"].bool(), suffix_mask], dim=1) - states = qwen(input_ids=input_ids, attention_mask=mask, output_hidden_states=True) + # mm_token_type_ids (used for M-RoPE) must cover the appended suffix tokens + # too; they are plain text -> type 0. + if "mm_token_type_ids" in extra_inputs: + tt = extra_inputs["mm_token_type_ids"] + extra_inputs["mm_token_type_ids"] = torch.cat( + [tt, torch.zeros_like(suffix_ids, dtype=tt.dtype)], dim=1 + ) + + states = qwen( + input_ids=input_ids, + attention_mask=mask, + output_hidden_states=True, + **extra_inputs, + ) # (1, L, num_layers, hidden) hiddens = torch.stack([states.hidden_states[i] for i in select_layers], dim=2) diff --git a/toolkit/dataloader_mixins.py b/toolkit/dataloader_mixins.py index f878e94..32861f3 100644 --- a/toolkit/dataloader_mixins.py +++ b/toolkit/dataloader_mixins.py @@ -1891,9 +1891,7 @@ class TextEmbeddingCachingMixin: self.sd.set_device_state_preset('cache_text_encoder') did_move = True - if file_item.encode_control_in_text_embeddings: - if file_item.control_path is None: - raise Exception(f"Could not find a control image for {file_item.path} which is needed for this model") + if file_item.encode_control_in_text_embeddings and file_item.control_path is not None: ctrl_img_list = [] control_path_list = file_item.control_path if not isinstance(file_item.control_path, list): diff --git a/ui/src/app/jobs/new/options.ts b/ui/src/app/jobs/new/options.ts index 74c7335..c28ae30 100644 --- a/ui/src/app/jobs/new/options.ts +++ b/ui/src/app/jobs/new/options.ts @@ -1085,7 +1085,76 @@ export const modelArchs: ModelArch[] = [ additionalSections: [ 'model.low_vram', 'model.layer_offloading', - 'model.assistant_lora_path' + 'model.assistant_lora_path', + ], + }, + { + name: 'krea2:o_edit', + label: 'Krea 2 (raw) [Edit Training]', + group: 'experimental', + defaults: { + 'config.process[0].model.name_or_path': ['krea/Krea-2-Raw', defaultNameOrPath], + 'config.process[0].model.quantize': [true, false], + 'config.process[0].model.quantize_te': [true, false], + 'config.process[0].train.timestep_type': ['linear', 'sigmoid'], + 'config.process[0].network.conv': [undefined, 16], + 'config.process[0].network.conv_alpha': [undefined, 16], + 'config.process[0].model.low_vram': [true, false], + 'config.process[0].model.model_kwargs': [ + { + edit: true, + match_target_res: false, + }, + {}, + ], + }, + disableSections: [ + 'network.conv', + ], + additionalSections: [ + 'datasets.multi_control_paths', + 'sample.multi_ctrl_imgs', + 'model.low_vram', + 'model.layer_offloading', + 'model.qie.match_target_res', + ], + }, + { + name: 'krea2:o_edit_turbo', + label: 'Krea 2 Turbo (w/ Training Adapter) [Edit Training]', + group: 'experimental', + defaults: { + 'config.process[0].model.name_or_path': ['krea/Krea-2-Turbo', defaultNameOrPath], + 'config.process[0].model.quantize': [true, false], + 'config.process[0].model.quantize_te': [true, false], + 'config.process[0].train.timestep_type': ['linear', 'sigmoid'], + 'config.process[0].network.conv': [undefined, 16], + 'config.process[0].network.conv_alpha': [undefined, 16], + 'config.process[0].model.low_vram': [true, false], + 'config.process[0].model.assistant_lora_path': [ + 'ostris/krea2_turbo_training_adapter/krea2_turbo_training_adapter_v1.safetensors', + undefined, + ], + 'config.process[0].sample.guidance_scale': [1, 4], + 'config.process[0].sample.sample_steps': [8, 25], + 'config.process[0].model.model_kwargs': [ + { + edit: true, + match_target_res: false, + }, + {}, + ], + }, + disableSections: [ + 'network.conv', + ], + additionalSections: [ + 'datasets.multi_control_paths', + 'sample.multi_ctrl_imgs', + 'model.low_vram', + 'model.layer_offloading', + 'model.assistant_lora_path', + 'model.qie.match_target_res', ], }, { diff --git a/version.py b/version.py index 0b4f246..dba99a5 100644 --- a/version.py +++ b/version.py @@ -1 +1 @@ -VERSION = "0.10.19" +VERSION = "0.10.20"