Add support for Krea2 (#906)
* Add support for krea2 * Update repo pointer to actual repo
This commit is contained in:
committed by
GitHub
parent
af594061ab
commit
99be3d96a2
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
|
||||
3
extensions_built_in/diffusion_models/krea2/__init__.py
Normal file
3
extensions_built_in/diffusion_models/krea2/__init__.py
Normal file
@@ -0,0 +1,3 @@
|
||||
from .krea2 import Krea2Model
|
||||
|
||||
__all__ = ["Krea2Model"]
|
||||
474
extensions_built_in/diffusion_models/krea2/krea2.py
Normal file
474
extensions_built_in/diffusion_models/krea2/krea2.py
Normal 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()
|
||||
}
|
||||
461
extensions_built_in/diffusion_models/krea2/src/mmdit.py
Normal file
461
extensions_built_in/diffusion_models/krea2/src/mmdit.py
Normal 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
|
||||
260
extensions_built_in/diffusion_models/krea2/src/pipeline.py
Normal file
260
extensions_built_in/diffusion_models/krea2/src/pipeline.py
Normal 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]
|
||||
@@ -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]
|
||||
@@ -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',
|
||||
|
||||
@@ -1 +1 @@
|
||||
VERSION = "0.10.16"
|
||||
VERSION = "0.10.17"
|
||||
|
||||
Reference in New Issue
Block a user