Add support for Krea2 (#906)

* Add support for krea2

* Update repo pointer to actual repo
This commit is contained in:
Jaret Burkett (Ostris)
2026-06-23 09:18:14 -06:00
committed by GitHub
parent af594061ab
commit 99be3d96a2
10 changed files with 1307 additions and 1 deletions

View File

@@ -46,6 +46,7 @@ AI Toolkit is an easy to use all in one training suite for diffusion models. I t
- [Wan-AI/Wan2.2-TI2V-5B-Diffusers](https://huggingface.co/Wan-AI/Wan2.2-TI2V-5B-Diffusers) (Wan 2.2 TI2V 5B)
- [Lightricks/LTX-2](https://huggingface.co/Lightricks/LTX-2) (LTX-2)
- [Lightricks/LTX-2.3](https://huggingface.co/Lightricks/LTX-2.3) (LTX-2.3)
- [krea/Krea-2-Raw](https://huggingface.co/krea/Krea-2-Raw) (Krea 2)
### Audio
- [ACE-Step/Ace-Step1.5](https://huggingface.co/ACE-Step/Ace-Step1.5) (Ace Step 1.5)

View File

@@ -15,6 +15,7 @@ from .hidream.hidream_o1_model import HidreamO1Model
from .z_image.z_image_l2p_model import ZImageL2PModel
from .ideogram4 import Ideogram4Model
from .prx_pixel_t2i import PRXPixelT2IModel
from .krea2 import Krea2Model
from .boogu_image import BooguImageModel, BooguImageEditModel
AI_TOOLKIT_MODELS = [
@@ -45,6 +46,7 @@ AI_TOOLKIT_MODELS = [
ZImageL2PModel,
Ideogram4Model,
PRXPixelT2IModel,
Krea2Model,
BooguImageModel,
BooguImageEditModel,
]

View File

@@ -0,0 +1,3 @@
from .krea2 import Krea2Model
__all__ = ["Krea2Model"]

View File

@@ -0,0 +1,474 @@
"""Krea 2 (K2) for ai-toolkit.
Krea 2 is a single-stream MMDiT text-to-image model:
- text encoder: Qwen3-VL-4B-Instruct (a stack of 12 hidden-state layers is fed
in; ``src/text_encoder.py``),
- autoencoder: the Qwen-Image VAE (f8, 16 latent channels, the same VAE the
``qwen_image`` arch uses),
- denoiser: ``SingleStreamDiT`` (``src/mmdit.py``), which fuses the text layers
with a small ``TextFusionTransformer`` and runs the packed [text | image]
sequence through ``SingleStreamBlock`` layers.
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 os
from typing import List, Optional
import torch
from safetensors.torch import load_file, save_file
import huggingface_hub
from huggingface_hub.errors import EntryNotFoundError
from diffusers import AutoencoderKLQwenImage
from transformers import (
AutoTokenizer,
Qwen2TokenizerFast,
Qwen3VLForConditionalGeneration,
)
from optimum.quanto import freeze, QTensor
from toolkit.config_modules import GenerateImageConfig, ModelConfig
from toolkit.models.base_model import BaseModel
from toolkit.basic import flush
from toolkit.advanced_prompt_embeds import AdvancedPromptEmbeds
from toolkit.samplers.custom_flowmatch_sampler import (
CustomFlowMatchEulerDiscreteScheduler,
)
from toolkit.accelerator import unwrap_model
from toolkit.metadata import get_meta_for_safetensors
from toolkit.util.quantize import quantize, get_qtype, quantize_model
from .src.mmdit import SingleStreamDiT, SingleMMDiTConfig
from .src.text_encoder import encode_krea_prompt, SELECT_LAYERS
from .src.pipeline import Krea2Pipeline, pad_text_features, predict_velocity
# The reference "single_mmdit_large_wide" architecture (oss_raw / oss_turbo share it).
KREA2_MMDIT_CONFIG = dict(
features=6144,
tdim=256,
txtdim=2560,
heads=48,
kvheads=12,
multiplier=4,
layers=28,
patch=2,
channels=16,
txtheads=20,
txtkvheads=20,
txtlayers=12,
)
# Krea 2's mu schedule is exponential time-shifting whose mu is linearly
# interpolated in image-token count between (256-res -> 0.5) and (1280-res ->
# 1.15) -- exactly what CustomFlowMatchEulerDiscreteScheduler's dynamic shifting
# does, so we mirror those endpoints here for the training timestep distribution.
# x1 = (256 // (8*2))**2 = 256
# x2 = (1280 // (8*2))**2 = 6400
scheduler_config = {
"base_image_seq_len": 256,
"max_image_seq_len": 6400,
"base_shift": 0.5,
"max_shift": 1.15,
"num_train_timesteps": 1000,
"shift": 1.0,
"use_dynamic_shifting": True,
"time_shift_type": "exponential",
}
# Defaults; both overridable via model.model_kwargs.
QWEN3_VL_PATH = "Qwen/Qwen3-VL-4B-Instruct"
QWEN_IMAGE_VAE_PATH = "Qwen/Qwen-Image"
HF_TOKEN = os.getenv("HF_TOKEN", None)
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.
``name_or_path`` may be: a ``.safetensors`` file, a directory containing one
(``filename`` or the lone ``.safetensors`` in it), or a hub repo id (the
file ``filename`` is downloaded, defaulting to ``model.safetensors``).
"""
if name_or_path.endswith(".safetensors") and os.path.isfile(name_or_path):
return load_file(name_or_path)
if os.path.isdir(name_or_path):
if filename is not None:
return load_file(os.path.join(name_or_path, filename))
candidates = [f for f in os.listdir(name_or_path) if f.endswith(".safetensors")]
if len(candidates) == 1:
return load_file(os.path.join(name_or_path, candidates[0]))
raise FileNotFoundError(
f"Could not pick an MMDiT checkpoint in {name_or_path}: found "
f"{candidates}. Set model.model_kwargs.checkpoint_filename."
)
# Treat as a hub repo id. When no filename is given, derive it from the repo
# name's trailing segment (e.g. "krea/Krea-2-Raw" -> "raw.safetensors",
# "krea/Krea-2-Turbo" -> "turbo.safetensors").
fname = filename or (name_or_path.split("/")[-1].split("-")[-1].lower() + ".safetensors")
try:
path = huggingface_hub.hf_hub_download(
repo_id=name_or_path, filename=fname, token=HF_TOKEN
)
except EntryNotFoundError as e:
raise FileNotFoundError(
f"Could not find {fname!r} in hub repo {name_or_path!r}. Set "
"model.model_kwargs.checkpoint_filename to the weight file name."
) from e
return load_file(path)
class Krea2Model(BaseModel):
arch = "krea2"
def __init__(
self,
device,
model_config: ModelConfig,
dtype="bf16",
custom_pipeline=None,
noise_scheduler=None,
**kwargs,
):
super().__init__(
device, model_config, dtype, custom_pipeline, noise_scheduler, **kwargs
)
self.is_flow_matching = True
self.is_transformer = True
self.target_lora_modules = ["SingleStreamDiT"]
self.patch_size = KREA2_MMDIT_CONFIG["patch"]
self.vae_scale_factor = 8 # Qwen-Image VAE is f8
# Safety cap on prompt token length (truncation only); embeds are stored
# per-sample at natural length and padded to the batch max at the model call.
self.max_text_length = int(
self.model_config.model_kwargs.get("max_text_length", 512)
)
# Qwen2TokenizerFast used to tokenize the assistant suffix (matches the
# reference's separate processor pass).
self.processor = None
@staticmethod
def get_train_scheduler():
return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
def get_bucket_divisibility(self):
# 8 for the VAE downsample, 2 for the patch size.
return self.vae_scale_factor * self.patch_size
# ------------------------------------------------------------------
# Loading
# ------------------------------------------------------------------
def _load_transformer(self):
dtype = self.torch_dtype
self.print_and_status_update("Loading transformer (SingleStreamDiT)")
mmdit_kwargs = dict(KREA2_MMDIT_CONFIG)
mmdit_kwargs.update(self.model_config.model_kwargs.get("mmdit_config", {}))
config = SingleMMDiTConfig(**mmdit_kwargs)
# Build on meta, then materialize straight from the checkpoint.
with torch.device("meta"):
transformer = SingleStreamDiT(config)
self.print_and_status_update(" - fetching transformer weights")
state_dict = _load_mmdit_state_dict(
self.model_config.name_or_path,
self.model_config.model_kwargs.get("checkpoint_filename", None),
)
state_dict = {
k: (v.to(dtype) if v.is_floating_point() else v)
for k, v in state_dict.items()
}
self.print_and_status_update(" - loading transformer state dict")
transformer.load_state_dict(state_dict, strict=True, assign=True)
del state_dict
flush()
return transformer
def _load_text_encoder(self):
dtype = self.torch_dtype
te_path = self.model_config.model_kwargs.get("text_encoder_path", QWEN3_VL_PATH)
self.print_and_status_update(f"Loading Qwen3-VL text encoder from {te_path}")
tokenizer = AutoTokenizer.from_pretrained(
te_path, max_length=self.max_text_length, token=HF_TOKEN
)
processor = Qwen2TokenizerFast.from_pretrained(
te_path, max_length=self.max_text_length, token=HF_TOKEN
)
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
text_encoder.eval()
text_encoder.requires_grad_(False)
flush()
return tokenizer, processor, text_encoder
def _load_vae(self):
vae_path = self.model_config.model_kwargs.get("vae_path", QWEN_IMAGE_VAE_PATH)
self.print_and_status_update(f"Loading Qwen-Image VAE from {vae_path}")
vae = AutoencoderKLQwenImage.from_pretrained(
vae_path, subfolder="vae", torch_dtype=self.vae_torch_dtype, token=HF_TOKEN
)
vae.eval()
vae.requires_grad_(False)
return vae
def load_model(self):
dtype = self.torch_dtype
self.print_and_status_update("Loading Krea 2 model")
transformer = self._load_transformer()
if self.model_config.quantize:
self.print_and_status_update("Quantizing transformer")
quantize_model(self, transformer)
flush()
if self.model_config.low_vram:
self.print_and_status_update("Moving transformer to CPU")
transformer.to("cpu")
else:
transformer.to(self.device_torch, dtype=dtype)
flush()
tokenizer, 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)
quantize(text_encoder, weights=get_qtype(self.model_config.qtype_te))
freeze(text_encoder)
flush()
if self.model_config.low_vram:
self.print_and_status_update("Moving text encoder to CPU")
text_encoder.to("cpu")
else:
text_encoder.to(self.device_torch)
flush()
vae = self._load_vae()
vae.to(self.vae_device_torch, dtype=self.vae_torch_dtype)
self.noise_scheduler = Krea2Model.get_train_scheduler()
self.vae = vae
self.text_encoder = text_encoder
self.tokenizer = tokenizer
self.processor = processor
self.model = transformer
self.pipeline = Krea2Pipeline(self)
self.print_and_status_update("Model Loaded")
# ------------------------------------------------------------------
# Generation (training previews)
# ------------------------------------------------------------------
def get_generation_pipeline(self):
return Krea2Pipeline(self)
def generate_single_image(
self,
pipeline: Krea2Pipeline,
gen_config: GenerateImageConfig,
conditional_embeds: AdvancedPromptEmbeds,
unconditional_embeds: AdvancedPromptEmbeds,
generator: torch.Generator,
extra: dict,
):
if self.model.device == torch.device("cpu"):
self.model.to(self.device_torch)
sc = self.get_bucket_divisibility()
gen_config.width = int(gen_config.width // sc * sc)
gen_config.height = int(gen_config.height // sc * sc)
img = pipeline(
conditional_embeds=conditional_embeds,
unconditional_embeds=unconditional_embeds,
height=gen_config.height,
width=gen_config.width,
num_inference_steps=gen_config.num_inference_steps,
guidance_scale=gen_config.guidance_scale,
latents=gen_config.latents,
generator=generator,
)[0]
return img
# ------------------------------------------------------------------
# Training hooks
# ------------------------------------------------------------------
def get_noise_prediction(
self,
latent_model_input: torch.Tensor, # (B, 16, h, w)
timestep: torch.Tensor, # 0..1000 scale
text_embeddings: AdvancedPromptEmbeds,
**kwargs,
):
if self.model.device == torch.device("cpu"):
self.model.to(self.device_torch)
# 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
if t.dim() == 0:
t = t.unsqueeze(0)
if t.shape[0] != latent_model_input.shape[0]:
t = t.expand(latent_model_input.shape[0])
context, text_mask = pad_text_features(
text_embeddings.text_embeds, self.device_torch, self.torch_dtype
)
pred = predict_velocity(
self.transformer,
latent_model_input.to(self.device_torch, self.torch_dtype),
t,
context,
text_mask,
)
return pred
def get_prompt_embeds(self, prompt) -> AdvancedPromptEmbeds:
if isinstance(prompt, str):
prompt = [prompt]
if self.text_encoder.device == torch.device("cpu"):
self.text_encoder.to(self.device_torch)
# 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
# batch size; predict_velocity restores the layer axis. Padding to the
# 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:
features = encode_krea_prompt(
self.text_encoder,
self.tokenizer,
self.processor,
p,
max_length=self.max_text_length,
select_layers=SELECT_LAYERS,
)
# (L, n, d) -> (L, n*d)
features = features.reshape(features.shape[0], -1)
features_list.append(features.to(self.torch_dtype))
return AdvancedPromptEmbeds(text_embeds=features_list)
def get_loss_target(self, *args, **kwargs):
# Flow-matching velocity target: noise - clean.
noise = kwargs.get("noise")
batch = kwargs.get("batch")
return (noise - batch.latents).detach()
def get_model_has_grad(self):
return False
def get_te_has_grad(self):
return False
# ------------------------------------------------------------------
# VAE (Qwen-Image AutoencoderKLQwenImage -- same handling as qwen_image arch)
# ------------------------------------------------------------------
def encode_images(self, image_list: List[torch.Tensor], device=None, dtype=None):
if device is None:
device = self.vae_device_torch
if dtype is None:
dtype = self.vae_torch_dtype
if self.vae.device == torch.device("cpu"):
self.vae.to(device)
self.vae.eval()
self.vae.requires_grad_(False)
image_list = [image.to(device, dtype=dtype) for image in image_list]
images = torch.stack(image_list).to(device, dtype=dtype)
# AutoencoderKLQwenImage is a video VAE: add a frame dim.
images = images.unsqueeze(2)
latents = self.vae.encode(images).latent_dist.sample()
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(
1, self.vae.config.z_dim, 1, 1, 1
).to(latents.device, latents.dtype)
latents = (latents - latents_mean) * latents_std
latents = latents.squeeze(2) # drop frame dim
return latents.to(device, dtype=dtype)
def decode_latents(self, latents: torch.Tensor, device=None, dtype=None):
if device is None:
device = self.vae_device_torch
if dtype is None:
dtype = self.vae_torch_dtype
if self.vae.device == torch.device("cpu"):
self.vae.to(device)
latents = latents.to(device, dtype=dtype)
latents = latents.unsqueeze(2) # add frame dim
latents_mean = (
torch.tensor(self.vae.config.latents_mean)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents_std = (
torch.tensor(self.vae.config.latents_std)
.view(1, self.vae.config.z_dim, 1, 1, 1)
.to(latents.device, latents.dtype)
)
latents = latents * latents_std + latents_mean
images = self.vae.decode(latents).sample
images = images.squeeze(2) # drop frame dim
return images.to(device, dtype=dtype)
# ------------------------------------------------------------------
# Saving / bookkeeping
# ------------------------------------------------------------------
def save_model(self, output_path, meta, save_dtype):
if not output_path.endswith(".safetensors"):
output_path = output_path + ".safetensors"
transformer: SingleStreamDiT = unwrap_model(self.model)
state_dict = transformer.state_dict()
save_dict = {}
for k, v in state_dict.items():
if isinstance(v, QTensor):
v = v.dequantize()
save_dict[k] = v.clone().to("cpu", dtype=save_dtype)
meta = get_meta_for_safetensors(meta, name="krea2")
save_file(save_dict, output_path, metadata=meta)
def get_base_model_version(self):
return "krea2"
def get_transformer_block_names(self) -> Optional[List[str]]:
return ["blocks"]
def convert_lora_weights_before_save(self, state_dict):
return {
k.replace("transformer.", "diffusion_model."): v
for k, v in state_dict.items()
}
def convert_lora_weights_before_load(self, state_dict):
return {
k.replace("diffusion_model.", "transformer."): v
for k, v in state_dict.items()
}

View File

@@ -0,0 +1,461 @@
"""Krea 2 (K2) single-stream MMDiT backbone.
Vendored from the reference ``mmdit.py`` for ai-toolkit. This is a single-stream
MMDiT: Qwen3-VL text features are fused by a small ``TextFusionTransformer`` and
then concatenated with the patchified image latent tokens into one sequence that
flows through ``SingleStreamBlock`` layers. The model predicts the flow-matching
velocity on the image tokens.
Differences from the reference (all training-driven, numerically equivalent):
- ``torch.compile`` decorators are dropped (they fight gradient checkpointing,
LoRA module swapping and variable shapes during training).
- Attention uses a plain ``F.scaled_dot_product_attention`` instead of forcing
the cuDNN SDPA backend, so it works across dtypes / masks / backward.
- ``enable_gradient_checkpointing`` / ``disable_gradient_checkpointing`` and a
per-block ``torch.utils.checkpoint`` wrapper are added (gated on
``torch.is_grad_enabled()`` so eval/sampling never pays for it).
"""
import math
from dataclasses import dataclass
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor
from torch.nn.attention import SDPBackend, sdpa_kernel
from torch.utils.checkpoint import checkpoint
def rope(pos: Tensor, dim: int, theta: float = 1e4, ntk: float = 1.0) -> Tensor:
scale = torch.arange(0, dim, 2, dtype=torch.float64, device=pos.device) / dim
omega = 1.0 / ((theta * ntk) ** scale)
out = torch.einsum("...n,d->...nd", pos, omega)
out = torch.stack(
[torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1
)
out = rearrange(out, "b n d (i j) -> b n d i j", i=2, j=2)
return out.float()
def ropeapply(xq: Tensor, xk: Tensor, freqs: Tensor) -> tuple[Tensor, Tensor]:
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
freqs = freqs[:, None, :, :, :]
xq_ = freqs[..., 0] * xq_[..., 0] + freqs[..., 1] * xq_[..., 1]
xk_ = freqs[..., 0] * xk_[..., 0] + freqs[..., 1] * xk_[..., 1]
return xq_.reshape(*xq.shape).to(xq.dtype), xk_.reshape(*xk.shape).to(xk.dtype)
def attention(
q: Tensor,
k: Tensor,
v: Tensor,
mask: Tensor | None = None,
scale: float | None = None,
gqa: bool = False,
) -> Tensor:
with sdpa_kernel(SDPBackend.CUDNN_ATTENTION):
x = F.scaled_dot_product_attention(
q, k, v, attn_mask=mask, scale=scale, enable_gqa=gqa
)
return rearrange(x, "B H L D -> B L (H D)")
def _mask(mask: Tensor) -> Tensor:
"""Expand a (B, L) key-padding mask into a (B, 1, L, L) attention mask."""
return mask.unsqueeze(1).unsqueeze(2) * mask.unsqueeze(1).unsqueeze(3)
def temb(
t: Tensor,
dim: int,
period: float = 1e4,
tfactor: float = 1e3,
device: torch.device = None,
dtype: torch.dtype = None,
) -> Tensor:
half = dim // 2
freqs = torch.exp(
-math.log(period)
* torch.arange(half, dtype=torch.float32, device=device)
/ half
)
# t: (B,) -> args: (B, 1, half), so the embedding broadcasts as a per-sample vec.
args = (t.float() * tfactor)[:, None, None] * freqs
sin, cos = torch.sin(args), torch.cos(args)
return torch.cat((cos, sin), dim=-1).to(dtype=dtype)
@dataclass
class SingleMMDiTConfig:
features: int
tdim: int
txtdim: int
heads: int
multiplier: int
layers: int
patch: int
channels: int
bias: bool = False
theta: float = 1e3
kvheads: int | None = None
txtlayers: int = 1
txtheads: int = 20
txtkvheads: int = 20
class SimpleModulation(torch.nn.Module):
def __init__(self, dim: int):
super().__init__()
self.lin = torch.nn.Parameter(torch.zeros(2, dim))
self.multiplier = 2
# vec (b d)
def forward(self, vec: Tensor):
out = vec + rearrange(self.lin, "two d -> 1 two d")
scale, shift = out.chunk(self.multiplier, dim=1)
return scale, shift
class DoubleSharedModulation(torch.nn.Module):
def __init__(self, dim: int):
super().__init__()
self.lin = torch.nn.Parameter(torch.zeros(6 * dim))
# vec (b (6 d))
def forward(self, vec: Tensor):
out = vec + self.lin
prescale, preshift, pregate, postscale, postshift, postgate = out.chunk(
6, dim=-1
)
return prescale, preshift, pregate, postscale, postshift, postgate
class PositionalEncoding(torch.nn.Module):
def __init__(self, dim, axdims: list[int], theta: float = 1e2, ntk: float = 1.0):
super().__init__()
self.axdims = axdims # how to split the head dimension across the position axes
self.theta = theta
self.ntk = ntk
def forward(self, pos: Tensor) -> Tensor:
return torch.cat(
[
rope(pos[..., i], d, self.theta, self.ntk)
for i, d in enumerate(self.axdims)
],
dim=-3,
)
class QKNorm(torch.nn.Module):
def __init__(self, dim: int):
super().__init__()
self.qnorm = RMSNorm(dim)
self.knorm = RMSNorm(dim)
def forward(self, q: Tensor, k: Tensor, v: Tensor) -> tuple[Tensor, Tensor, Tensor]:
return self.qnorm(q), self.knorm(k), v
class RMSNorm(torch.nn.Module):
def __init__(self, features: int, eps: float = 1e-05, device: torch.device = None):
super().__init__()
self.features = features
self.eps = eps
self.scale = torch.nn.Parameter(
torch.zeros(features, device=device, dtype=torch.float32)
)
def forward(self, x: Tensor) -> Tensor:
t, dtype = x.float(), x.dtype
t = F.rms_norm(
t, (self.features,), eps=self.eps, weight=(self.scale.float() + 1.0)
)
return t.to(dtype)
class SwiGLU(torch.nn.Module):
def __init__(
self, features: int, multiplier: int, bias: bool = False, multiple: int = 128
):
super().__init__()
mlpdim = int(2 * features / 3) * multiplier
mlpdim = multiple * ((mlpdim + multiple - 1) // multiple)
self.gate = torch.nn.Linear(features, mlpdim, bias=bias)
self.up = torch.nn.Linear(features, mlpdim, bias=bias)
self.down = torch.nn.Linear(mlpdim, features, bias=bias)
def forward(self, x: Tensor) -> Tensor:
return self.down(F.silu(self.gate(x)) * self.up(x))
class Attention(torch.nn.Module):
def __init__(self, dim: int, heads: int, kvheads: int = None, bias: bool = False):
super().__init__()
self.heads = heads
self.kvheads = kvheads if kvheads is not None else heads
self.headdim = dim // self.heads
self.wq = torch.nn.Linear(dim, self.headdim * self.heads, bias=bias)
self.wk = torch.nn.Linear(dim, self.headdim * self.kvheads, bias=bias)
self.wv = torch.nn.Linear(dim, self.headdim * self.kvheads, bias=bias)
self.gate = torch.nn.Linear(dim, dim, bias=bias)
self.qknorm = QKNorm(self.headdim)
self.gqa = self.heads != self.kvheads
self.wo = torch.nn.Linear(dim, dim, bias=bias)
def forward(
self, qkv: Tensor, freqs: Tensor | None = None, mask: Tensor | None = None
) -> Tensor:
q, k, v, gate = self.wq(qkv), self.wk(qkv), self.wv(qkv), self.gate(qkv)
q, k, v = (
rearrange(q, "B L (H D) -> B H L D", H=self.heads),
rearrange(k, "B L (H D) -> B H L D", H=self.kvheads),
rearrange(v, "B L (H D) -> B H L D", H=self.kvheads),
)
q, k, v = self.qknorm(q, k, v)
if freqs is not None:
q, k = ropeapply(q, k, freqs)
out = self.wo(attention(q, k, v, mask=mask, gqa=self.gqa) * F.sigmoid(gate))
return out
class LastLayer(torch.nn.Module):
def __init__(self, features: int, patch: int, channels: int):
super().__init__()
self.norm = RMSNorm(features)
self.linear = torch.nn.Linear(features, patch * patch * channels, bias=True)
self.modulation = SimpleModulation(features)
def forward(self, x: Tensor, tvec: Tensor) -> Tensor:
scale, shift = self.modulation(tvec)
x = (1 + scale) * self.norm(x) + shift
x = self.linear(x)
return x
class TextFusionBlock(torch.nn.Module):
def __init__(
self,
features: int,
heads: int,
multiplier: int,
bias: bool = False,
kvheads: int = None,
):
super().__init__()
self.prenorm = RMSNorm(features)
self.postnorm = RMSNorm(features)
self.attn = Attention(dim=features, heads=heads, bias=bias, kvheads=kvheads)
self.mlp = SwiGLU(features, multiplier, bias)
def forward(self, x: Tensor, mask: Tensor | None = None) -> Tensor:
x = x + self.attn(self.prenorm(x), mask=mask)
x = x + self.mlp(self.postnorm(x))
return x
class TextFusionTransformer(torch.nn.Module):
# num_txt_layers is the number of selected encoder hidden-state layers fed in
# (projected down to 1), NOT the transformer depth — that's fixed at 2 + 2 blocks.
def __init__(
self,
num_txt_layers: int,
txt_dim: int,
heads: int,
multiplier: int,
bias: bool = False,
kvheads: int = None,
):
super().__init__()
self.layerwise_blocks = torch.nn.ModuleList(
[
TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads)
for _ in range(2)
]
)
self.projector = torch.nn.Linear(num_txt_layers, 1, bias=False)
self.refiner_blocks = torch.nn.ModuleList(
[
TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads)
for _ in range(2)
]
)
def forward(self, x: Tensor, mask: Tensor | None = None) -> Tensor:
b, l, n, d = x.shape
x = x.reshape(b * l, n, d)
for block in self.layerwise_blocks:
x = block(x.contiguous(), mask=None)
x = rearrange(x, "(b l) n d -> b l d n", b=b, l=l)
# Collapse to 3D for the projector: a quantized (quanto) Linear's matmul
# kernel only accepts 2D/3D activations, and this layer-axis projection
# (n -> 1) otherwise feeds it a 4D (b, l, d, n) tensor.
x = self.projector(x.reshape(b * l, d, n))
x = x.reshape(b, l, d)
for block in self.refiner_blocks:
x = block(x, mask=mask)
return x
class SingleStreamBlock(nn.Module):
def __init__(
self,
features: int,
heads: int,
multiplier: int,
bias: bool = False,
kvheads: int = None,
):
super().__init__()
self.mod = DoubleSharedModulation(features)
self.prenorm = RMSNorm(features)
self.postnorm = RMSNorm(features)
self.attn = Attention(dim=features, heads=heads, bias=bias, kvheads=kvheads)
self.mlp = SwiGLU(features, multiplier, bias)
def forward(
self, x: Tensor, vec: Tensor, freqs: Tensor, mask: Tensor | None = None
) -> Tensor:
prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec)
x = x + pregate * self.attn(
(1 + prescale) * self.prenorm(x) + preshift, freqs, mask
)
x = x + postgate * self.mlp((1 + postscale) * self.postnorm(x) + postshift)
return x
class SingleStreamDiT(nn.Module):
def __init__(self, config: SingleMMDiTConfig):
super().__init__()
self.config = config
self.gradient_checkpointing = False
headdim = config.features // config.heads
axes = [
headdim - 12 * (headdim // 16),
6 * (headdim // 16),
6 * (headdim // 16),
]
assert sum(axes) == headdim, f"sum(axes) = {sum(axes)}, headdim = {headdim}"
assert all(a % 2 == 0 for a in axes), f"axes = {axes}"
self.posemb = PositionalEncoding(
config.features, axes, theta=config.theta, ntk=1.0
)
self.first = nn.Linear(
config.channels * config.patch**2, config.features, bias=True
)
self.blocks = nn.ModuleList(
[
SingleStreamBlock(
config.features,
config.heads,
config.multiplier,
config.bias,
config.kvheads,
)
for _ in range(config.layers)
]
)
self.tmlp = nn.Sequential(
nn.Linear(config.tdim, config.features),
nn.GELU(approximate="tanh"),
nn.Linear(config.features, config.features),
)
self.txtfusion = TextFusionTransformer(
config.txtlayers,
config.txtdim,
config.txtheads,
config.multiplier,
config.bias,
config.txtkvheads,
)
self.txtmlp = nn.Sequential(
RMSNorm(config.txtdim),
nn.Linear(config.txtdim, config.features),
nn.GELU(approximate="tanh"),
nn.Linear(config.features, config.features),
)
self.last = LastLayer(config.features, config.patch, config.channels)
self.tproj = nn.Sequential(
nn.GELU(approximate="tanh"), nn.Linear(config.features, config.features * 6)
)
@property
def device(self) -> torch.device:
return next(self.parameters()).device
@property
def dtype(self) -> torch.dtype:
return next(self.parameters()).dtype
def enable_gradient_checkpointing(self):
self.gradient_checkpointing = True
def disable_gradient_checkpointing(self):
self.gradient_checkpointing = False
def forward(
self,
img: Tensor,
context: Tensor,
t: Tensor,
pos: Tensor,
mask: Tensor | None = None,
) -> Tensor:
img = self.first(img)
t = self.tmlp(temb(t, self.config.tdim, device=img.device, dtype=img.dtype))
tvec = self.tproj(t)
txtmask = _mask(mask[:, : context.shape[1]])
context = self.txtfusion(context, mask=txtmask)
context = self.txtmlp(context)
txtlen, imglen = context.shape[1], img.shape[1]
combined = torch.cat((context, img), dim=1)
# Pad combined sequence to a multiple of 256 to stabilize compiled kernel shapes.
fulllen = combined.shape[1]
_padlen = (-fulllen) % 256
if _padlen > 0:
combined = F.pad(combined, (0, 0, 0, _padlen))
mask = F.pad(mask, (0, _padlen), value=False)
pos = F.pad(pos, (0, 0, 0, _padlen))
mask = _mask(mask)
freqs = self.posemb(pos)
for block in self.blocks:
if self.gradient_checkpointing and torch.is_grad_enabled():
combined = checkpoint(
block,
combined,
tvec,
freqs,
mask,
use_reentrant=False,
)
else:
combined = block(combined, tvec, freqs, mask)
final = self.last(combined, t)
output = final[:, txtlen : txtlen + imglen, :]
return output

View File

@@ -0,0 +1,260 @@
"""Packing / sampling helpers for Krea 2.
Turns image latents + stacked Qwen3-VL text features into the single sequence the
``SingleStreamDiT`` consumes, and provides a minimal flow-matching sampler used to
render preview images during training.
Time convention: Krea 2 is a plain flow-matching model whose time runs ``t=1``
(pure noise) -> ``t=0`` (clean), the velocity it predicts is ``noise - clean``,
and ``x_t = (1 - t) * clean + t * noise``. This is *identical* to ai-toolkit's
convention, so unlike ideogram4 there is no flipping or negation -- the toolkit
``timestep / 1000`` flows straight through as ``t``.
"""
from __future__ import annotations
import math
from typing import List, Optional
import torch
from einops import rearrange, repeat
from PIL import Image
from diffusers.utils.torch_utils import randn_tensor
from .mmdit import SingleStreamDiT
# ---------------------------------------------------------------------------
# Text feature padding.
# ---------------------------------------------------------------------------
def pad_text_features(
features_list: List[torch.Tensor],
device: torch.device,
dtype: torch.dtype,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Right-pad a list of per-sample ``(Lt_i, F)`` features into a batch.
Each caption is stored 2D at its natural length -- the 12 stacked Qwen3-VL
hidden-state layers are flattened into the feature axis ``F = n * d`` so the
ai-toolkit batching machinery treats the list length as the batch size (it
only special-cases 2D per-sample tensors). The layer axis is restored in
``predict_velocity`` right before the MMDiT call. Padding to the batch max is
deferred to here. Returns ``(features (B, Lt, F), mask (B, Lt))``; the mask is
1 for real text tokens and 0 for padding.
"""
lengths = [f.shape[0] for f in features_list]
max_len = max(lengths)
dim = features_list[0].shape[-1]
batch_size = len(features_list)
features = torch.zeros(batch_size, max_len, dim, device=device, dtype=dtype)
mask = torch.zeros(batch_size, max_len, dtype=torch.long, device=device)
for i, f in enumerate(features_list):
ln = f.shape[0]
features[i, :ln] = f.to(device, dtype)
mask[i, :ln] = 1
return features, mask
# ---------------------------------------------------------------------------
# Latent <-> token packing and combined position / mask construction.
# ---------------------------------------------------------------------------
def prepare(
img: torch.Tensor, txtlen: int, patch: int, txtmask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Patchify the latent and build the combined text+image position / mask.
in: img (B, C, h, w) image latent
txtlen number of text tokens
patch transformer patch size
txtmask (B, txtlen) long/bool mask, 1 for real text tokens
out: (img_tokens (B, h/p*w/p, C*p*p), pos (B, txtlen+imglen, 3),
mask (B, txtlen+imglen))
"""
b, _, h, w = img.shape
h_, w_ = h // patch, w // patch
imgids = torch.zeros((h_, w_, 3), device=img.device)
imgids[..., 1] = torch.arange(h_, device=img.device)[:, None]
imgids[..., 2] = torch.arange(w_, device=img.device)[None, :]
imgpos = repeat(imgids, "h w three -> b (h w) three", b=b, three=3)
imgmask = torch.ones(b, h_ * w_, device=img.device, dtype=torch.bool)
img = rearrange(img, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch)
txtpos = torch.zeros(b, txtlen, 3, device=img.device)
mask = torch.cat((txtmask.to(img.device).bool(), imgmask), dim=1)
pos = torch.cat((txtpos, imgpos), dim=1)
return img, 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
) -> torch.Tensor:
"""Run the MMDiT on the packed [text | image] 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.
"""
patch = model.config.patch
b, c, h, w = latents.shape
# Restore the stacked-layer axis flattened in pad_text_features: F -> (n, d).
n = model.config.txtlayers
context = context.reshape(
context.shape[0], context.shape[1], n, context.shape[-1] // n
)
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)
# (B, imglen, c*p*p) -> (B, c, h, w)
velocity = rearrange(
out,
"b (h w) (c ph pw) -> b c (h ph) (w pw)",
ph=patch,
pw=patch,
h=h // patch,
w=w // patch,
)
return velocity
# ---------------------------------------------------------------------------
# Resolution-aware flow-matching timestep schedule.
# ---------------------------------------------------------------------------
def timesteps(
seq_len: int,
steps: int,
x1: float,
x2: float,
y1: float = 0.5,
y2: float = 1.15,
sigma: float = 1.0,
mu: Optional[float] = None,
) -> List[float]:
"""Resolution-aware flow-matching timestep schedule (t: 1 -> 0).
``mu`` is interpolated linearly in image-sequence length between (x1, y1) and
(x2, y2), then used to time-shift a uniform 1->0 grid. Pass an explicit ``mu``
to pin a constant shift regardless of resolution (the distilled turbo
checkpoint was trained at a fixed mu=1.15).
"""
ts = torch.linspace(1, 0, steps + 1)
if mu is None:
slope = (y2 - y1) / (x2 - x1)
mu = slope * seq_len + (y1 - slope * x1)
ts = math.exp(mu) / (math.exp(mu) + (1.0 / ts - 1.0) ** sigma)
return ts.tolist()
# ---------------------------------------------------------------------------
# Minimal sampling pipeline (for training previews).
# ---------------------------------------------------------------------------
class Krea2Pipeline:
"""Lightweight flow-matching sampler used by ai-toolkit's preview generation."""
def __init__(self, model):
# ``model`` is the Krea2Model so we can reuse its encode/decode and config.
self.model = model
@property
def device(self):
return self.model.device_torch
def to(self, *args, **kwargs):
return self
def set_progress_bar_config(self, **kwargs):
pass
@torch.no_grad()
def __call__(
self,
conditional_embeds,
unconditional_embeds,
height: int = 1024,
width: int = 1024,
num_inference_steps: int = 28,
guidance_scale: float = 4.5,
latents: Optional[torch.Tensor] = None,
generator: Optional[torch.Generator] = None,
**kwargs,
) -> List[Image.Image]:
model = self.model
device = model.device_torch
dtype = model.torch_dtype
transformer: SingleStreamDiT = model.transformer
patch = model.patch_size
ae_scale = model.vae_scale_factor # 8
mkw = model.model_config.model_kwargs
y1 = float(mkw.get("schedule_y1", 0.5))
y2 = float(mkw.get("schedule_y2", 1.15))
minres = int(mkw.get("schedule_min_res", 256))
maxres = int(mkw.get("schedule_max_res", 1280))
mu = mkw.get("schedule_mu", None)
mu = float(mu) if mu is not None else None
do_cfg = guidance_scale > 0 and unconditional_embeds is not None
gh = height // (ae_scale * patch)
gw = width // (ae_scale * patch)
latent_channels = transformer.config.channels
# Starting gaussian noise in the (B, C, h8, w8) latent layout.
if latents is None:
shape = (1, latent_channels, height // ae_scale, width // ae_scale)
latents = randn_tensor(
shape, generator=generator, device=device, dtype=torch.float32
)
latents = latents.to(device, dtype=torch.float32)
cond_feats, cond_mask = pad_text_features(
conditional_embeds.text_embeds, device, dtype
)
if do_cfg:
uncond_feats, uncond_mask = pad_text_features(
unconditional_embeds.text_embeds, device, dtype
)
# min_res / max_res define the (x1,y1)-(x2,y2) interpolation endpoints for mu.
align = ae_scale * patch
x1 = (minres // align) ** 2
x2 = (maxres // align) ** 2
ts = timesteps(gh * gw, num_inference_steps, x1, x2, y1=y1, y2=y2, mu=mu)
# Euler integration of the flow ODE (with optional CFG).
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
)
if do_cfg:
v_uncond = predict_velocity(
transformer, latents.to(dtype), t, uncond_feats, uncond_mask
)
v = v_cond + guidance_scale * (v_cond - v_uncond)
else:
v = v_cond
latents = latents + (tprev - tcurr) * v.to(torch.float32)
images = model.decode_latents(latents, device=device, dtype=dtype)
images = images.float().clamp(-1.0, 1.0)
images = ((images + 1.0) * 127.5).round().to(torch.uint8)
images = images.permute(0, 2, 3, 1).cpu().numpy()
return [Image.fromarray(arr) for arr in images]

View File

@@ -0,0 +1,84 @@
"""Qwen3-VL text conditioning for Krea 2.
Vendored / adapted from the reference ``encoder.py``. Krea 2 conditions on a
*stack* of hidden states pulled from several layers of Qwen3-VL-4B-Instruct
(``SELECT_LAYERS``), wrapped in a fixed instruction template. The MMDiT's
``TextFusionTransformer`` later collapses that layer axis down to one.
The reference encodes a whole batch padded to ``max_length``; here we encode one
prompt at a time at its natural length (the ai-toolkit pattern -- caches stay
small, any prompts can share a batch, and per-sample padding is deferred to the
model call). The fixed instruction prefix is fed through the model as context but
its hidden states are sliced off the returned features, exactly like the
reference.
"""
import torch
from torch import Tensor
# Layers of Qwen3-VL whose hidden states are stacked and fed to the MMDiT (12).
SELECT_LAYERS: tuple[int, ...] = (2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35)
# Fixed instruction template wrapped around every prompt. The prefix is fed
# through the model as context but its hidden states are dropped from the output
# (the assistant only ever sees the prompt + suffix tokens as conditioning).
PROMPT_TEMPLATE_ENCODE_PREFIX = (
"<|im_start|>system\nDescribe the image by detailing the color, shape, size, "
"texture, quantity, text, spatial relationships of the objects and "
"background:<|im_end|>\n<|im_start|>user\n"
)
PROMPT_TEMPLATE_ENCODE_SUFFIX = "<|im_end|>\n<|im_start|>assistant\n"
# Number of leading tokens (the system prefix) sliced off the encoded features.
PROMPT_TEMPLATE_ENCODE_START_IDX = 34
@torch.no_grad()
def encode_krea_prompt(
qwen,
tokenizer,
processor,
prompt: str,
max_length: int = 512,
select_layers: tuple[int, ...] = SELECT_LAYERS,
prefix_idx: int = PROMPT_TEMPLATE_ENCODE_START_IDX,
) -> Tensor:
"""Encode a single prompt into stacked Qwen3-VL hidden states.
Returns a ``(L, num_select_layers, hidden)`` float tensor (in the encoder's
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.
"""
device = qwen.device
# The suffix ("...assistant\n") is tokenized without the BOS/template extras
# the main tokenizer adds, matching the reference's separate processor pass.
suffix_inputs = processor(
text=[PROMPT_TEMPLATE_ENCODE_SUFFIX], return_tensors="pt"
).to(device, non_blocking=True)
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)
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)
# (1, L, num_layers, hidden)
hiddens = torch.stack([states.hidden_states[i] for i in select_layers], dim=2)
# Drop the system-prefix tokens; what remains is prompt + suffix conditioning.
hiddens = hiddens[:, prefix_idx:]
return hiddens[0]

View File

@@ -1039,6 +1039,27 @@ export const modelArchs: ModelArch[] = [
'model.layer_offloading',
],
},
{
name: 'krea2',
label: 'Krea 2 (K2)',
group: 'image',
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],
},
disableSections: [
'network.conv',
],
additionalSections: [
'model.low_vram',
'model.layer_offloading',
],
},
{
name: 'boogu_image',
label: 'Boogu Image',

View File

@@ -1 +1 @@
VERSION = "0.10.16"
VERSION = "0.10.17"