Files
ai-toolkit/extensions_built_in/diffusion_models/krea2/src/mmdit.py

583 lines
20 KiB
Python

"""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,
ref_span: tuple[int, int] | None = None,
kv_capture: list | None = None,
kv_cache: tuple[Tensor, 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)
if kv_capture is not None and ref_span is not None:
# Stash this block's post-RoPE ref K/V so later denoising steps can
# run without the ref tokens in the sequence (clone: drop the view
# into the full-sequence K/V so only the ref span stays alive).
kv_capture.append(
(
k[:, :, ref_span[0] : ref_span[1]].clone(),
v[:, :, ref_span[0] : ref_span[1]].clone(),
)
)
if kv_cache is not None:
# Cached ref K/V are already RoPE'd at their original positions.
k = torch.cat((k, kv_cache[0]), dim=2)
v = torch.cat((v, kv_cache[1]), dim=2)
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,
ref_span: tuple[int, int] | None = None,
kv_capture: list | None = None,
kv_cache: tuple[Tensor, Tensor] | None = None,
) -> Tensor:
attn_kwargs = dict(ref_span=ref_span, kv_capture=kv_capture, kv_cache=kv_cache)
# ``vec`` is the (B, 1, 6*features) modulation input, or a tuple
# ``(vec, refvec, split)`` for reference-image conditioning: tokens
# ``[:split]`` (text + noisy image) are modulated with ``vec`` while
# tokens ``[split:]`` (clean reference tokens) use ``refvec`` built from
# t=0 (ComfyUI Kontext "index_timestep_zero"). Applied per span rather
# than materializing a per-token (B, L, 6*features) tensor.
if isinstance(vec, tuple):
vec, refvec, split = vec
m = self.mod(vec)
r = self.mod(refvec)
def mod(h, scale, shift):
return torch.cat(
(
(1 + m[scale]) * h[:, :split] + m[shift],
(1 + r[scale]) * h[:, split:] + r[shift],
),
dim=1,
)
def gate(h, g):
return torch.cat((m[g] * h[:, :split], r[g] * h[:, split:]), dim=1)
x = x + gate(
self.attn(mod(self.prenorm(x), 0, 1), freqs, mask, **attn_kwargs), 2
)
x = x + gate(self.mlp(mod(self.postnorm(x), 3, 4)), 5)
return x
prescale, preshift, pregate, postscale, postshift, postgate = self.mod(vec)
x = x + pregate * self.attn(
(1 + prescale) * self.prenorm(x) + preshift, freqs, mask, **attn_kwargs
)
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,
reflen: int = 0,
isolate_refs: bool = False,
ref_kv_capture: list | None = None,
ref_kv_cache: tuple[list, 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))
blockvec = tvec
if reflen > 0:
# The last ``reflen`` image tokens are clean reference tokens: they
# get t=0 modulation (ComfyUI Kontext "index_timestep_zero") while
# text + noisy image tokens keep the real t. Padding tokens fall in
# the t=0 span, but they are masked from attention and sliced off
# the output, so their values never matter.
t0 = self.tmlp(
temb(
torch.zeros_like(t[:, 0, 0]),
self.config.tdim,
device=img.device,
dtype=img.dtype,
)
)
blockvec = (tvec, self.tproj(t0), txtlen + imglen - reflen)
padmask = mask # (B, L) key-padding mask, incl. the 256-alignment pad
mask = _mask(mask)
if reflen > 0 and isolate_refs:
# Asymmetric attention (OminiControl2-style "feature reuse"): ref
# queries attend only to ref keys, while text + noisy queries still
# see everything. Combined with the t=0 modulation above, ref hidden
# states become independent of t and of the noisy tokens, so their
# per-layer K/V can be computed once and cached across denoising
# steps at inference. Changes attention flow vs the base model, so
# it needs to be trained in.
split = txtlen + imglen - reflen
is_ref = torch.zeros(
combined.shape[1], dtype=torch.bool, device=combined.device
)
is_ref[split : split + reflen] = True
mask = mask & (~is_ref[:, None] | is_ref[None, :])
# Ref K/V caching (inference-only; requires isolate_refs so the cached
# features are step-invariant). Capture mode: this pass has the refs in
# the sequence and records each block's post-RoPE ref K/V. Reuse mode:
# the refs are dropped from the sequence (reflen == 0) and the cached
# K/V are appended as extra attention keys instead.
ref_span = None
if ref_kv_capture is not None and reflen > 0:
assert isolate_refs, "ref K/V capture requires isolate_refs"
split = txtlen + imglen - reflen
ref_span = (split, split + reflen)
blockcaches = [None] * len(self.blocks)
if ref_kv_cache is not None:
blockcaches, refmask = ref_kv_cache
# live queries may attend a cached ref key wherever that ref token
# is real (refmask right-pads samples with fewer ref tokens)
extra = padmask.unsqueeze(1).unsqueeze(3) & refmask.unsqueeze(1).unsqueeze(2)
mask = torch.cat((mask, extra), dim=3)
freqs = self.posemb(pos)
for block, blockkv in zip(self.blocks, blockcaches):
if self.gradient_checkpointing and torch.is_grad_enabled():
combined = checkpoint(
block,
combined,
blockvec,
freqs,
mask,
use_reentrant=False,
)
else:
combined = block(
combined,
blockvec,
freqs,
mask,
ref_span=ref_span,
kv_capture=ref_kv_capture,
kv_cache=blockkv,
)
final = self.last(combined, t)
output = final[:, txtlen : txtlen + imglen - reflen, :]
return output