From 5beddc8afce90ddf1662e7a8cff6f89214add9a9 Mon Sep 17 00:00:00 2001 From: Yihao Wang <42559837+AgainstEntropy@users.noreply.github.com> Date: Tue, 23 Jun 2026 22:56:51 -0700 Subject: [PATCH] [Diffusion] Add Krea 2 support (#29052) Co-authored-by: Claude Opus 4.8 --- .../sglang/multimodal_gen/apps/webui/main.py | 36 +- .../configs/models/dits/krea2.py | 77 +++ .../configs/pipeline_configs/krea2.py | 135 ++++ .../multimodal_gen/configs/sample/krea2.py | 29 + python/sglang/multimodal_gen/registry.py | 11 + .../runtime/models/dits/krea2.py | 595 ++++++++++++++++++ .../multimodal_gen/runtime/pipelines/krea2.py | 105 ++++ .../stages/model_specific_stages/krea2.py | 167 +++++ 8 files changed, 1143 insertions(+), 12 deletions(-) create mode 100644 python/sglang/multimodal_gen/configs/models/dits/krea2.py create mode 100644 python/sglang/multimodal_gen/configs/pipeline_configs/krea2.py create mode 100644 python/sglang/multimodal_gen/configs/sample/krea2.py create mode 100644 python/sglang/multimodal_gen/runtime/models/dits/krea2.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines/krea2.py create mode 100644 python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/krea2.py diff --git a/python/sglang/multimodal_gen/apps/webui/main.py b/python/sglang/multimodal_gen/apps/webui/main.py index 61cac9ab7..7dc19f013 100644 --- a/python/sglang/multimodal_gen/apps/webui/main.py +++ b/python/sglang/multimodal_gen/apps/webui/main.py @@ -50,29 +50,41 @@ def run_sgl_diffusion_webui(server_args: ServerArgs): validate_repo_id(candidate) # let it raise if invalid return candidate - repo_id = resolve_model_repo_id(server_args.model_path) - if envs.SGLANG_USE_MODELSCOPE.get(): - from modelscope.hub.api import HubApi + # Prefer the hub pipeline tag for Hub models; fall back to the loaded pipeline's + # own task_type for local checkpoints (e.g. a single .safetensors path), which + # have no hub repo to query. + task_name = None + try: + repo_id = resolve_model_repo_id(server_args.model_path) + if envs.SGLANG_USE_MODELSCOPE.get(): + from modelscope.hub.api import HubApi - api = HubApi() - model_info_obj = api.model_info(repo_id) - task_name = model_info_obj.tasks[0]["Name"].replace("-synthesis", "") - else: - from huggingface_hub import model_info + api = HubApi() + model_info_obj = api.model_info(repo_id) + task_name = model_info_obj.tasks[0]["Name"].replace("-synthesis", "") + else: + from huggingface_hub import model_info - task_name = model_info(repo_id).pipeline_tag + task_name = model_info(repo_id).pipeline_tag + except Exception as e: + logger.info( + "Could not resolve task from the model hub (%s); using the loaded " + "pipeline's task_type.", + e, + ) # init client sync_scheduler_client.initialize(server_args) if task_name in ("text-to-video", "image-to-video", "video-to-video"): task_type = "video" - elif task_name in ["text-to-image", "image-to-image"]: + elif task_name in ("text-to-image", "image-to-image"): task_type = "image" else: - raise ValueError( - f"The task name {task_name} of model {server_args.model_path} is not a valid task name. Please check the model path." + task_type = ( + "image" if server_args.pipeline_config.task_type.is_image_gen() else "video" ) + task_name = task_name or server_args.pipeline_config.task_type.name video_visible_only = task_type == "video" image_visible_only = task_type == "image" diff --git a/python/sglang/multimodal_gen/configs/models/dits/krea2.py b/python/sglang/multimodal_gen/configs/models/dits/krea2.py new file mode 100644 index 000000000..99a83f97e --- /dev/null +++ b/python/sglang/multimodal_gen/configs/models/dits/krea2.py @@ -0,0 +1,77 @@ +# Krea-2 (K2) single-stream MMDiT architecture config. +# +# Parameter names follow the released K2 checkpoint, so the MMDiT safetensors load +# without remapping (identity `param_names_mapping`). +from dataclasses import dataclass, field +from typing import Tuple + +from sglang.multimodal_gen.configs.models.dits.base import DiTArchConfig, DiTConfig + + +@dataclass +class Krea2ArchConfig(DiTArchConfig): + features: int = 6144 # hidden dim + tdim: int = 256 # timestep embedding dim + txtdim: int = 2560 # text-encoder hidden dim (Qwen3-VL-4B hidden_size) + heads: int = 48 + kvheads: int = 12 # GQA 4:1 + multiplier: int = 4 # SwiGLU expansion multiplier + layers: int = 28 + patch: int = 2 + channels: int = 16 # VAE latent channels + bias: bool = False + theta: float = 1e3 # RoPE theta + txtlayers: int = 12 # number of text-encoder hidden-state layers fused by txtfusion + txtheads: int = 20 + txtkvheads: int = 20 + + # 3-axis RoPE split over head_dim=128: [global, h, w] = (32, 48, 48). + axes_dims: Tuple[int, int, int] = (32, 48, 48) + + # Joint (text+image) sequence is padded to a multiple of this many tokens. + seq_multiple_of: int = 256 + + # Packed patch-token width (channels * patch**2); used by the VAE unpack path. + in_channels: int = 64 + + # BaseDiT-required instance attrs (overwritten in __post_init__). + hidden_size: int = 6144 + num_attention_heads: int = 48 + num_channels_latents: int = 16 + + # Module/parameter names match the released checkpoint, so weights load with an + # identity mapping. + param_names_mapping: dict = field(default_factory=dict) + + # Diffusers LoRA checkpoints prefix every DiT key with the pipeline component name + # (e.g. transformer.transformer_blocks.0.attn.to_q.lora_A). Strip that prefix so the + # LoRA keys line up with this model's module names. Only the LoRA path uses this; the + # main checkpoint still loads with the identity param_names_mapping above. + lora_param_names_mapping: dict = field( + default_factory=lambda: {r"^transformer\.": ""} + ) + + def __post_init__(self) -> None: + super().__post_init__() + self.hidden_size = self.features + self.num_attention_heads = self.heads + self.num_channels_latents = self.channels + assert self.features % self.heads == 0 + assert ( + sum(self.axes_dims) == self.features // self.heads + ), f"sum(axes_dims)={sum(self.axes_dims)} != head_dim={self.features // self.heads}" + + @property + def head_dim(self) -> int: + return self.features // self.heads + + @property + def in_features_packed(self) -> int: + """Patch-embed input width: channels * patch**2.""" + return self.channels * self.patch**2 + + +@dataclass +class Krea2DitConfig(DiTConfig): + arch_config: Krea2ArchConfig = field(default_factory=Krea2ArchConfig) + prefix: str = "k2" diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/krea2.py b/python/sglang/multimodal_gen/configs/pipeline_configs/krea2.py new file mode 100644 index 000000000..ad5029858 --- /dev/null +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/krea2.py @@ -0,0 +1,135 @@ +# Krea-2 (K2) pipeline config. +from dataclasses import dataclass, field + +import torch +from diffusers.image_processor import VaeImageProcessor + +from sglang.multimodal_gen.configs.models import DiTConfig, VAEConfig +from sglang.multimodal_gen.configs.models.dits.krea2 import Krea2DitConfig +from sglang.multimodal_gen.configs.models.vaes.qwenimage import QwenImageVAEConfig +from sglang.multimodal_gen.configs.pipeline_configs.base import ( + ImagePipelineConfig, + ModelTaskType, +) + +# Resolution-interpolation endpoints for the time-shift `mu` (reference sampler): +# mu is linear in image-token count between (min_res, 0.5) and (max_res, 1.15). +_MU_MIN_RES = 256 +_MU_MAX_RES = 1280 +_MU_Y1 = 0.5 +_MU_Y2 = 1.15 + + +@dataclass +class Krea2PipelineConfig(ImagePipelineConfig): + """Krea-2 single-stream MMDiT, text-to-image. + + Reuses the Qwen-Image VAE (same checkpoint) and its latent pack/unpack + + decode de-normalization. K2-specific pieces are the 3-axis joint-stream RoPE + positions and key-padding mask (``prepare_pos_cond_kwargs``), the 4-D + layer-stacked text conditioning (``get_pos_prompt_embeds``), and the flow + time-shift ``mu``. + """ + + task_type: ModelTaskType = ModelTaskType.T2I + should_use_guidance: bool = False + enable_autocast: bool = False + vae_tiling: bool = False + vae_sp: bool = False + vae_precision: str = "bf16" + # The released Qwen3-VL text encoder is bf16 (text_encoder/config.json dtype); + # the base loader defaults to fp32, which perturbs the conditioning embeddings. + text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16",)) + + dit_config: DiTConfig = field(default_factory=Krea2DitConfig) + vae_config: VAEConfig = field(default_factory=QwenImageVAEConfig) + + # Pinned time-shift mu (distilled `oss_turbo`); set None to derive from resolution. + pinned_mu: float | None = 1.15 + + def __post_init__(self): + self.vae_scale_factor = self.vae_config.get_vae_scale_factor() + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor) + + def get_vae_scale_factor(self): + return self.vae_config.get_vae_scale_factor() + + # --- text conditioning: K2 feeds the raw layer-stacked encoder states --- + def get_pos_prompt_embeds(self, batch): + return batch.prompt_embeds[0] + + def get_neg_prompt_embeds(self, batch): + return batch.negative_prompt_embeds[0] + + # --- joint-stream RoPE positions + key-padding mask --- + def _build_pos_and_mask(self, context, text_mask, batch, device): + b, txt_len = context.shape[0], context.shape[1] + patch = self.dit_config.arch_config.patch + vsf = self.get_vae_scale_factor() + h_tok = int(batch.height) // vsf // patch + w_tok = int(batch.width) // vsf // patch + + img_ids = torch.zeros(h_tok, w_tok, 3, device=device) + img_ids[..., 1] = torch.arange(h_tok, device=device)[:, None] + img_ids[..., 2] = torch.arange(w_tok, device=device)[None, :] + img_pos = img_ids.reshape(h_tok * w_tok, 3).unsqueeze(0).expand(b, -1, -1) + txt_pos = torch.zeros(b, txt_len, 3, device=device) + pos = torch.cat([txt_pos, img_pos], dim=1) + + img_mask = torch.ones(b, h_tok * w_tok, dtype=torch.bool, device=device) + if text_mask is None: + txt_mask = torch.ones(b, txt_len, dtype=torch.bool, device=device) + else: + txt_mask = text_mask.to(device=device).bool() + mask = torch.cat([txt_mask, img_mask], dim=1) + return {"pos": pos, "mask": mask} + + def prepare_pos_cond_kwargs(self, batch, device, rotary_emb, dtype): + text_mask = batch.prompt_embeds_mask[0] if batch.prompt_embeds_mask else None + return self._build_pos_and_mask( + batch.prompt_embeds[0], text_mask, batch, device + ) + + def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype): + text_mask = ( + batch.negative_prompt_embeds_mask[0] + if batch.negative_prompt_embeds_mask + else None + ) + return self._build_pos_and_mask( + batch.negative_prompt_embeds[0], text_mask, batch, device + ) + + # --- timestep shift --- + def compute_mu(self, image_seq_len: int) -> float: + if self.pinned_mu is not None: + return self.pinned_mu + patch = self.dit_config.arch_config.patch + vsf = self.get_vae_scale_factor() + x1 = (_MU_MIN_RES // (vsf * patch)) ** 2 + x2 = (_MU_MAX_RES // (vsf * patch)) ** 2 + slope = (_MU_Y2 - _MU_Y1) / (x2 - x1) + return slope * image_seq_len + (_MU_Y1 - slope * x1) + + def prepare_sigmas(self, sigmas, num_inference_steps): + return self._prepare_sigmas(sigmas, num_inference_steps) + + # --- VAE decode (same as Qwen-Image: latents * std + mean) --- + def get_decode_scale_and_shift(self, device, dtype, vae): + vae_arch = self.vae_config.arch_config + scaling_factor = 1.0 / torch.tensor(vae_arch.latents_std, device=device).view( + 1, vae_arch.z_dim, 1, 1, 1 + ).to(device, dtype) + shift_factor = ( + torch.tensor(vae_arch.latents_mean) + .view(1, vae_arch.z_dim, 1, 1, 1) + .to(device, dtype) + ) + return scaling_factor, shift_factor + + def post_denoising_loop(self, latents, batch): + latents, batch_size, channels, height, width = self._unpad_and_unpack_latents( + latents, batch + ) + latents = latents.reshape(batch_size, channels // (2 * 2), 1, height, width) + return latents diff --git a/python/sglang/multimodal_gen/configs/sample/krea2.py b/python/sglang/multimodal_gen/configs/sample/krea2.py new file mode 100644 index 000000000..8b11fca6e --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/krea2.py @@ -0,0 +1,29 @@ +# Krea-2 (K2) sampling defaults. +# +# `guidance_scale` is the SGLang classifier-free-guidance scale, which equals the +# K2 reference `cfg + 1` (SGLang combines `uncond + scale*(cond-uncond)`, the K2 +# reference uses `cond + cfg*(cond-uncond)`). So K2 cfg=0 -> guidance_scale=1.0 +# (no CFG), K2 cfg=3.5 -> guidance_scale=4.5. +from dataclasses import dataclass + +from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams + + +@dataclass +class Krea2SamplingParams(SamplingParams): + """Distilled `oss_turbo` defaults: 8 steps, CFG disabled.""" + + negative_prompt: str = "" + num_frames: int = 1 + height: int = 1024 + width: int = 1024 + guidance_scale: float = 1.0 + num_inference_steps: int = 8 + + +@dataclass +class Krea2RawSamplingParams(Krea2SamplingParams): + """Base `oss_raw` defaults: full sampler with CFG.""" + + guidance_scale: float = 4.5 + num_inference_steps: int = 52 diff --git a/python/sglang/multimodal_gen/registry.py b/python/sglang/multimodal_gen/registry.py index d42e636af..1072cd5e8 100644 --- a/python/sglang/multimodal_gen/registry.py +++ b/python/sglang/multimodal_gen/registry.py @@ -63,6 +63,7 @@ from sglang.multimodal_gen.configs.pipeline_configs.ideogram import ( from sglang.multimodal_gen.configs.pipeline_configs.joy_image import ( JoyImageEditPipelineConfig, ) +from sglang.multimodal_gen.configs.pipeline_configs.krea2 import Krea2PipelineConfig from sglang.multimodal_gen.configs.pipeline_configs.ltx_2 import ( LTX2PipelineConfig, LTX23PipelineConfig, @@ -115,6 +116,9 @@ from sglang.multimodal_gen.configs.sample.ideogram import Ideogram4SamplingParam from sglang.multimodal_gen.configs.sample.joy_image import ( JoyImageEditSamplingParams, ) +from sglang.multimodal_gen.configs.sample.krea2 import ( + Krea2SamplingParams, +) from sglang.multimodal_gen.configs.sample.lingbot_world import ( LingBotWorldSamplingParams, ) @@ -843,6 +847,13 @@ def _register_configs(): lambda hf_id: "z-image" in hf_id.lower() and "turbo" not in hf_id.lower() ], ) + # Krea-2 (K2) + register_configs( + sampling_param_cls=Krea2SamplingParams, + pipeline_config_cls=Krea2PipelineConfig, + hf_model_paths=["krea/Krea-2"], + model_detectors=[lambda hf_id: "krea-2" in hf_id.lower()], + ) # Qwen-Image register_configs( sampling_param_cls=QwenImageSamplingParams, diff --git a/python/sglang/multimodal_gen/runtime/models/dits/krea2.py b/python/sglang/multimodal_gen/runtime/models/dits/krea2.py new file mode 100644 index 000000000..01049e4b0 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/models/dits/krea2.py @@ -0,0 +1,595 @@ +"""Krea-2 (K2) single-stream MMDiT. + +Text and image tokens are concatenated into a single joint-attention stream. The +model uses GQA attention with a sigmoid output gate, 6-way shared adaLN +modulation, a text-fusion transformer that fuses the selected text-encoder +hidden-state layers into one, and interleaved 3-axis RoPE. Module and parameter +names follow the released K2 checkpoint, so weights load without remapping. +""" + +import math +import os +from typing import Any, Optional + +import torch +import torch.nn as nn +import torch.nn.functional as F +from einops import rearrange +from torch import Tensor + +from sglang.multimodal_gen.configs.models.dits.krea2 import Krea2DitConfig +from sglang.multimodal_gen.runtime.distributed import get_tp_world_size +from sglang.multimodal_gen.runtime.layers.attention import USPAttention +from sglang.multimodal_gen.runtime.layers.attention.layer import build_varlen_mask_meta +from sglang.multimodal_gen.runtime.layers.linear import ( + ColumnParallelLinear, + RowParallelLinear, +) +from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( + LayerwiseOffloadableModuleMixin, +) +from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + + +# --------------------------------------------------------------------------- # +# Functional helpers +# --------------------------------------------------------------------------- # +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 _fused_qknorm_rope_enabled() -> bool: + return os.getenv("SGLANG_ENABLE_FUSED_QKNORM_ROPE", "1").lower() not in ( + "0", + "false", + "off", + "no", + ) + + +def _can_use_fused_qknorm_rope(head_dim: int, dtype: torch.dtype) -> bool: + from sglang.jit_kernel.diffusion.qknorm_rope import ( + can_use_fused_inplace_qknorm_rope, + ) + + return can_use_fused_inplace_qknorm_rope(head_dim, head_dim, False, dtype) + + +def _qknorm_rope_cos_sin_cache(freqs: Tensor) -> Tensor: + """``[num_tokens, head_dim]`` cos|sin cache for the fused QKNorm+RoPE kernel. + + K2's ``rope`` packs each token's rotation as ``[[cos, -sin], [sin, cos]]`` in a + ``[B, N, head_dim//2, 2, 2]`` tensor; the kernel wants the per-token cosines then + sines concatenated. Positions come from the image grid (batch-invariant), so the + first batch row is representative. + """ + return torch.cat([freqs[0, :, :, 0, 0], freqs[0, :, :, 1, 0]], dim=-1).float() + + +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 + ) + 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) + + +def norm_scale_shift( + x: Tensor, weight: Tensor, scale: Tensor, shift: Tensor, eps: float +) -> Tensor: + """Fused RMSNorm + modulation: ``rms_norm(x) * weight * (1 + scale) + shift``. + + ``weight`` is the effective RMSNorm weight (K2 stores ``scale``, so callers + pass ``scale + 1``), kept off the checkpoint so the identity load is unaffected. + """ + if x.is_cuda and x.shape[-1] % 256 == 0: + from sglang.jit_kernel.diffusion.cutedsl.scale_residual_norm_scale_shift import ( + fused_norm_scale_shift, + ) + + return fused_norm_scale_shift( + x.contiguous(), + weight.contiguous(), + None, + scale.contiguous(), + shift.contiguous(), + "rms", + eps, + ) + normed = F.rms_norm(x.float(), (x.shape[-1],), weight=weight.float(), eps=eps) + return (normed.to(x.dtype) * (1 + scale) + shift).to(x.dtype) + + +# --------------------------------------------------------------------------- # +# Submodules +# --------------------------------------------------------------------------- # +class TimeEmbed(nn.Module): + """Timestep embedding MLP: linear_1 -> gelu(tanh) -> linear_2.""" + + def __init__(self, in_dim: int, dim: int): + super().__init__() + self.linear_1 = nn.Linear(in_dim, dim) + self.linear_2 = nn.Linear(dim, dim) + + def forward(self, x: Tensor) -> Tensor: + return self.linear_2(F.gelu(self.linear_1(x), approximate="tanh")) + + +class TxtIn(nn.Module): + """Text-context projection: rms-norm -> linear_1 -> gelu(tanh) -> linear_2.""" + + def __init__(self, txt_dim: int, dim: int): + super().__init__() + self.norm = RMSNorm(txt_dim) + self.linear_1 = nn.Linear(txt_dim, dim) + self.linear_2 = nn.Linear(dim, dim) + + def forward(self, x: Tensor) -> Tensor: + return self.linear_2(F.gelu(self.linear_1(self.norm(x)), approximate="tanh")) + + +class PositionalEncoding(nn.Module): + def __init__(self, dim, axdims: list[int], theta: float = 1e2, ntk: float = 1.0): + super().__init__() + self.axdims = axdims + 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 RMSNorm(nn.Module): + """RMSNorm with effective scale ``weight + 1`` (``weight`` initialized to 0), + computed in fp32. The parameter is named ``weight`` to match the released + checkpoint; the ``+ 1`` is applied in the forward.""" + + def __init__(self, features: int, eps: float = 1e-05, device: torch.device = None): + super().__init__() + self.features = features + self.eps = eps + self.weight = 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.weight.float() + 1.0) + ) + return t.to(dtype) + + +class SwiGLU(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) + # Tensor-parallel: gate/up shard the hidden dim by column, down all-reduces. + self.gate = ColumnParallelLinear( + features, mlpdim, bias=bias, gather_output=False + ) + self.up = ColumnParallelLinear(features, mlpdim, bias=bias, gather_output=False) + self.down = RowParallelLinear( + mlpdim, features, bias=bias, input_is_parallel=True + ) + + def forward(self, x: Tensor) -> Tensor: + gate, _ = self.gate(x) + up, _ = self.up(x) + out, _ = self.down(F.silu(gate) * up) + return out + + +class Attention(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 + + # Tensor-parallel: q/k/v/gate shard heads by column, to_out all-reduces. + # Parameter names match the released checkpoint (to_q/to_k/to_v/to_gate, + # norm_q/norm_k, to_out.0) so the checkpoint loads with an identity mapping. + tp = get_tp_world_size() + assert ( + self.heads % tp == 0 and self.kvheads % tp == 0 + ), f"heads={self.heads}, kvheads={self.kvheads} must be divisible by tp={tp}" + self.local_heads = self.heads // tp + self.local_kvheads = self.kvheads // tp + + self.to_q = ColumnParallelLinear( + dim, self.headdim * self.heads, bias=bias, gather_output=False + ) + self.to_k = ColumnParallelLinear( + dim, self.headdim * self.kvheads, bias=bias, gather_output=False + ) + self.to_v = ColumnParallelLinear( + dim, self.headdim * self.kvheads, bias=bias, gather_output=False + ) + self.to_gate = ColumnParallelLinear(dim, dim, bias=bias, gather_output=False) + self.norm_q = RMSNorm(self.headdim) + self.norm_k = RMSNorm(self.headdim) + # to_out is a ModuleList ([linear]) so the param is to_out.0.weight, matching + # the diffusers Attention layout in the released checkpoint. + self.to_out = nn.ModuleList( + [RowParallelLinear(dim, dim, bias=bias, input_is_parallel=True)] + ) + # Native GQA flash via the platform backend; parameterless. + self.attn = USPAttention( + num_heads=self.local_heads, + head_size=self.headdim, + num_kv_heads=self.local_kvheads, + dropout_rate=0, + softmax_scale=None, + causal=False, + ) + + def forward( + self, + qkv: Tensor, + freqs: Tensor | None = None, + key_mask: Tensor | None = None, + mask_meta: dict | None = None, + ) -> Tensor: + q, _ = self.to_q(qkv) + k, _ = self.to_k(qkv) + v, _ = self.to_v(qkv) + gate, _ = self.to_gate(qkv) + + hd = self.headdim + # Fast path: fuse RMSNorm(q), RMSNorm(k) and RoPE into one in-place kernel on + # the [B, S, H, D] layout USPAttention consumes (also skips the [B, H, L, D] + # transpose round-trip the eager path needs). Eager fallback below preserves + # parity off CUDA / for unsupported dtypes. + if ( + freqs is not None + and q.is_cuda + and q.dtype in (torch.float16, torch.bfloat16) + and _fused_qknorm_rope_enabled() + and _can_use_fused_qknorm_rope(hd, q.dtype) + ): + from sglang.jit_kernel.diffusion.qknorm_rope import ( + fused_inplace_qknorm_rope, + ) + + b, s = qkv.shape[0], qkv.shape[1] + q = q.view(b, s, self.local_heads, hd) + k = k.view(b, s, self.local_kvheads, hd) + v = v.view(b, s, self.local_kvheads, hd) + positions = torch.arange(s, device=q.device, dtype=torch.long) + if b > 1: + positions = positions.repeat(b) + fused_inplace_qknorm_rope( + q.reshape(-1, self.local_heads, hd), + k.reshape(-1, self.local_kvheads, hd), + (self.norm_q.weight.float() + 1.0).to(q.dtype), + (self.norm_k.weight.float() + 1.0).to(k.dtype), + _qknorm_rope_cos_sin_cache(freqs), + positions, + is_neox=False, + eps=self.norm_q.eps, + head_dim=hd, + rope_dim=hd, + ) + out = self.attn( + q, k, v, attn_mask=key_mask, attn_mask_meta=mask_meta + ).flatten(2) + else: + q, k, v = ( + rearrange(q, "B L (H D) -> B H L D", H=self.local_heads), + rearrange(k, "B L (H D) -> B H L D", H=self.local_kvheads), + rearrange(v, "B L (H D) -> B H L D", H=self.local_kvheads), + ) + q, k = self.norm_q(q), self.norm_k(k) + if freqs is not None: + q, k = ropeapply(q, k, freqs) + # USPAttention expects [B, S, H, D]; a [B, S] key mask + varlen metadata + # routes a ragged batch through the FA varlen fast path, else maskless. + out = self.attn( + q.transpose(1, 2).contiguous(), + k.transpose(1, 2).contiguous(), + v.transpose(1, 2).contiguous(), + attn_mask=key_mask, + attn_mask_meta=mask_meta, + ).flatten(2) + out, _ = self.to_out[0](out * F.sigmoid(gate)) + return out + + +class LastLayer(nn.Module): + def __init__(self, features: int, patch: int, channels: int): + super().__init__() + self.norm = RMSNorm(features) + self.linear = nn.Linear(features, patch * patch * channels, bias=True) + self.scale_shift_table = nn.Parameter(torch.zeros(2, features)) + + def forward(self, x: Tensor, tvec: Tensor) -> Tensor: + mod = tvec + rearrange(self.scale_shift_table, "two d -> 1 two d") + scale, shift = mod.chunk(2, dim=1) + x = norm_scale_shift(x, self.norm.weight + 1, scale, shift, self.norm.eps) + x = self.linear(x) + return x + + +class TextFusionBlock(nn.Module): + def __init__( + self, + features: int, + heads: int, + multiplier: int, + bias: bool = False, + kvheads: int = None, + ): + super().__init__() + self.norm1 = RMSNorm(features) + self.norm2 = RMSNorm(features) + self.attn = Attention(dim=features, heads=heads, bias=bias, kvheads=kvheads) + self.ff = SwiGLU(features, multiplier, bias) + + def forward( + self, + x: Tensor, + key_mask: Tensor | None = None, + mask_meta: dict | None = None, + ) -> Tensor: + x = x + self.attn(self.norm1(x), key_mask=key_mask, mask_meta=mask_meta) + x = x + self.ff(self.norm2(x)) + return x + + +class TextFusionTransformer(nn.Module): + """Fuses `num_txt_layers` selected encoder hidden-state layers into one. + + Depth is fixed at 2 layerwise + 2 refiner blocks; `num_txt_layers` is the + projector input width (the layer axis), NOT the transformer depth. + """ + + 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 = nn.ModuleList( + [ + TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads) + for _ in range(2) + ] + ) + self.projector = nn.Linear(num_txt_layers, 1, bias=False) + self.refiner_blocks = nn.ModuleList( + [ + TextFusionBlock(txt_dim, heads, multiplier, bias, kvheads) + for _ in range(2) + ] + ) + + def forward( + self, + x: Tensor, + key_mask: Tensor | None = None, + mask_meta: dict | 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()) + x = rearrange(x, "(b l) n d -> b l d n", b=b, l=l) + x = self.projector(x) + x = x.squeeze(-1) + for block in self.refiner_blocks: + x = block(x, key_mask=key_mask, mask_meta=mask_meta) + return x + + +class SingleStreamBlock(nn.Module): + def __init__( + self, + features: int, + heads: int, + multiplier: int, + bias: bool = False, + kvheads: int = None, + ): + super().__init__() + # (6, features) modulation table added to the timestep projection (AdaLN-single), + # stored directly on the block to match the released checkpoint. + self.scale_shift_table = nn.Parameter(torch.zeros(6, features)) + self.norm1 = RMSNorm(features) + self.norm2 = RMSNorm(features) + self.attn = Attention(dim=features, heads=heads, bias=bias, kvheads=kvheads) + self.ff = SwiGLU(features, multiplier, bias) + + def forward( + self, + x: Tensor, + vec: Tensor, + freqs: Tensor, + key_mask: Tensor | None = None, + mask_meta: dict | None = None, + ) -> Tensor: + mod = vec + self.scale_shift_table.reshape(-1) + prescale, preshift, pregate, postscale, postshift, postgate = mod.chunk( + 6, dim=-1 + ) + x = x + pregate * self.attn( + norm_scale_shift( + x, self.norm1.weight + 1, prescale, preshift, self.norm1.eps + ), + freqs, + key_mask, + mask_meta, + ) + x = x + postgate * self.ff( + norm_scale_shift( + x, self.norm2.weight + 1, postscale, postshift, self.norm2.eps + ) + ) + return x + + +# --------------------------------------------------------------------------- # +# Top-level model +# --------------------------------------------------------------------------- # +class Krea2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): + """K2 single-stream MMDiT for the SGLang diffusion runtime. + + Attribute names follow the released K2 checkpoint, so weights load with an + identity ``param_names_mapping``. + """ + + _fsdp_shard_conditions = [] + _compile_conditions = [] + param_names_mapping = Krea2DitConfig().arch_config.param_names_mapping + reverse_param_names_mapping = {} + + def __init__( + self, + config: Krea2DitConfig, + hf_config: dict[str, Any], + quant_config: Optional[Any] = None, + ) -> None: + super().__init__(config=config, hf_config=hf_config) + ac = config.arch_config + self.arch_config = ac + + self.hidden_size = ac.features + self.num_attention_heads = ac.heads + self.num_channels_latents = ac.channels + self.patch = ac.patch + self.channels = ac.channels + self.tdim = ac.tdim + + head_dim = ac.features // ac.heads + axes = list(ac.axes_dims) + assert sum(axes) == head_dim, f"sum(axes)={sum(axes)}, head_dim={head_dim}" + assert all(a % 2 == 0 for a in axes), f"axes={axes}" + + self.posemb = PositionalEncoding(ac.features, axes, theta=ac.theta, ntk=1.0) + self.img_in = nn.Linear(ac.channels * ac.patch**2, ac.features, bias=True) + self.transformer_blocks = nn.ModuleList( + [ + SingleStreamBlock( + ac.features, ac.heads, ac.multiplier, ac.bias, ac.kvheads + ) + for _ in range(ac.layers) + ] + ) + self.time_embed = TimeEmbed(ac.tdim, ac.features) + self.text_fusion = TextFusionTransformer( + ac.txtlayers, + ac.txtdim, + ac.txtheads, + ac.multiplier, + ac.bias, + ac.txtkvheads, + ) + self.txt_in = TxtIn(ac.txtdim, ac.features) + self.final_layer = LastLayer(ac.features, ac.patch, ac.channels) + # GELU(tanh) is applied in the forward; the linear matches time_mod_proj.weight. + self.time_mod_proj = nn.Linear(ac.features, ac.features * 6) + self.seq_multiple_of = ac.seq_multiple_of + # The 28 single-stream blocks (the ~24GB bulk) are streamed layer-by-layer + # under --dit-layerwise-offload, keeping only a small working set resident. + self.layer_names = ["transformer_blocks"] + + def _forward_impl( + self, + img: Tensor, + context: Tensor, + t: Tensor, + pos: Tensor, + mask: Tensor | None = None, + ) -> Tensor: + img = self.img_in(img) + t = self.time_embed(temb(t, self.tdim, device=img.device, dtype=img.dtype)) + tvec = self.time_mod_proj(F.gelu(t, approximate="tanh")) + + # A single or same-prompt batch has no padding, so attention runs maskless + # (native-GQA flash). A ragged batch builds varlen metadata from the + # key mask and takes the FA varlen path instead. + txt_key = txt_meta = joint_key = joint_meta = None + if mask is not None and not bool(mask.all()): + txt_key = mask[:, : context.shape[1]] + txt_meta = build_varlen_mask_meta(txt_key) + joint_key = mask + joint_meta = build_varlen_mask_meta(mask) + + context = self.text_fusion(context, key_mask=txt_key, mask_meta=txt_meta) + context = self.txt_in(context) + + txtlen, imglen = context.shape[1], img.shape[1] + combined = torch.cat((context, img), dim=1) + freqs = self.posemb(pos) + + for block in self.transformer_blocks: + combined = block(combined, tvec, freqs, joint_key, joint_meta) + + final = self.final_layer(combined, t) + output = final[:, txtlen : txtlen + imglen, :] + return output + + def forward( + self, + hidden_states: Tensor, + encoder_hidden_states: Tensor, + timestep: Tensor, + encoder_hidden_states_image=None, + guidance=None, + pos: Tensor = None, + mask: Tensor = None, + **kwargs, + ) -> Tensor: + return self._forward_impl( + img=hidden_states, + context=encoder_hidden_states, + t=timestep, + pos=pos, + mask=mask, + ) + + +EntryClass = [Krea2Transformer2DModel] diff --git a/python/sglang/multimodal_gen/runtime/pipelines/krea2.py b/python/sglang/multimodal_gen/runtime/pipelines/krea2.py new file mode 100644 index 000000000..948fbcc2c --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines/krea2.py @@ -0,0 +1,105 @@ +"""Krea-2 text-to-image pipeline (native diffusers layout). + +The released repo is diffusers-style (``model_index.json`` + ``transformer/``, +``text_encoder/``, ``vae/``, ``tokenizer/``, ``scheduler/`` subfolders), so the +base ``load_modules`` loads every component from it (the MMDiT via +``Krea2Transformer2DModel``, the Qwen3-VL text encoder, the Qwen-Image VAE, and +the ``FlowMatchEulerDiscreteScheduler``). This pipeline only adds the two +K2-specific touches the base loader can't infer: dropping the unused Qwen3-VL +vision tower (K2 conditions on text only) and building the assistant-suffix +tokenizer (``processor``), which has no ``model_index.json`` entry. The stage +chain is Krea2BeforeDenoisingStage -> DenoisingStage -> DecodingStage. +""" + +from typing import Any + +import torch +from transformers import Qwen2TokenizerFast + +from sglang.multimodal_gen.runtime.pipelines_core import LoRAPipeline +from sglang.multimodal_gen.runtime.pipelines_core.composed_pipeline_base import ( + ComposedPipelineBase, +) +from sglang.multimodal_gen.runtime.pipelines_core.stages import DenoisingStage +from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.krea2 import ( + Krea2BeforeDenoisingStage, +) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +_TEXT_MAX_LENGTH = 512 + + +class Krea2Pipeline(LoRAPipeline, ComposedPipelineBase): + pipeline_name = "Krea2Pipeline" + + from sglang.multimodal_gen.configs.pipeline_configs.krea2 import Krea2PipelineConfig + from sglang.multimodal_gen.configs.sample.krea2 import Krea2SamplingParams + + pipeline_config_cls = Krea2PipelineConfig + sampling_params_cls = Krea2SamplingParams + + # Every entry is a diffusers/transformers component declared in model_index.json, + # so the base loader handles them. "processor" is added in load_modules (no entry). + _required_config_modules = [ + "text_encoder", + "tokenizer", + "vae", + "transformer", + "scheduler", + ] + + def initialize_pipeline(self, server_args: ServerArgs): + vae_config = server_args.pipeline_config.vae_config + if hasattr(vae_config, "post_init"): + vae_config.post_init() + + def load_modules( + self, + server_args: ServerArgs, + loaded_modules: dict[str, torch.nn.Module] | None = None, + ) -> dict[str, Any]: + modules = super().load_modules(server_args, loaded_modules) + + # K2 conditions on text only: drop the unused Qwen3-VL vision tower that the + # base loader brings in with the full Qwen3VLModel (frees its weights and + # shrinks the encoder's CPU<->GPU page). It sits on the encoder or under .model. + text_encoder = modules.get("text_encoder") + if text_encoder is not None: + for owner in (text_encoder, getattr(text_encoder, "model", None)): + if owner is not None and getattr(owner, "visual", None) is not None: + del owner.visual + break + + # The conditioner appends a fixed assistant suffix, tokenized separately; + # model_index.json has no "processor" entry, so build one from tokenizer/. + tok_path = self._resolve_component_path(server_args, "tokenizer", "tokenizer") + modules["processor"] = Qwen2TokenizerFast.from_pretrained( + tok_path, max_length=_TEXT_MAX_LENGTH + ) + return modules + + def create_pipeline_stages(self, server_args: ServerArgs): + self.add_stage( + Krea2BeforeDenoisingStage( + vae=self.get_module("vae"), + text_encoder=self.get_module("text_encoder"), + tokenizer=self.get_module("tokenizer"), + processor=self.get_module("processor"), + transformer=self.get_module("transformer"), + scheduler=self.get_module("scheduler"), + ), + "k2_before_denoising_stage", + ) + self.add_stage( + DenoisingStage( + transformer=self.get_module("transformer"), + scheduler=self.get_module("scheduler"), + ), + ) + self.add_standard_decoding_stage() + + +EntryClass = [Krea2Pipeline] diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/krea2.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/krea2.py new file mode 100644 index 000000000..daab5da95 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/krea2.py @@ -0,0 +1,167 @@ +"""Krea-2 (K2) pre-processing stage: text encode + latent/timestep preparation. + +Consolidates everything before the denoising loop: Qwen3-VL text encoding with +the K2 system-prompt template and 12-layer hidden-state stacking, initial noise +latent packing, and the rectified-flow timestep schedule. Produces a batch the +standard DenoisingStage consumes unchanged. +""" + +import numpy as np +import torch +from einops import rearrange + +from sglang.multimodal_gen.runtime.distributed import get_local_torch_device +from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager import ( + ComponentUse, +) +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage +from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + +# K2 text-encode template (Qwen3-VL, text only). The system prefix is stripped +# from the hidden states after encoding (drop_idx tokens); a fixed assistant +# suffix is appended so the final token positions match training. +_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" +) +_SUFFIX = "<|im_end|>\n<|im_start|>assistant\n" +_DROP_IDX = 34 +_SUFFIX_START_IDX = 5 +_MAX_LENGTH = 512 +_SELECT_LAYERS = (2, 5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 35) + + +class Krea2BeforeDenoisingStage(PipelineStage): + def __init__(self, vae, text_encoder, tokenizer, processor, transformer, scheduler): + super().__init__() + self.vae = vae + self.text_encoder = text_encoder + self.tokenizer = tokenizer + self.processor = processor + self.transformer = transformer + self.scheduler = scheduler + + def component_uses(self, server_args: ServerArgs, stage_name: str | None = None): + # Declare the text encoder so the residency manager can CPU-offload it + # (with --text-encoder-cpu-offload) for the denoise loop, where it is idle. + return [ + ComponentUse( + self._component_stage_name(stage_name), + "text_encoder", + target_dtype=torch.bfloat16, + ) + ] + + @torch.no_grad() + def _encode(self, prompts, device, dtype): + """Reproduce the K2 conditioner: template + 12-layer hidden-state stack.""" + text = [_PREFIX + p for p in prompts] + suffix_inputs = self.processor( + text=[_SUFFIX] * len(text), return_tensors="pt" + ).to(device) + suffix_ids = suffix_inputs["input_ids"] + suffix_mask = suffix_inputs["attention_mask"].bool() + + # Pad to the batch's longest sequence (no padding for a single prompt) so + # the joint stream carries only valid tokens and attention needs no mask. + # The Qwen3-VL encoder is causal with right padding, so valid-token hidden + # states are identical to fixed max-length padding. + inputs = self.tokenizer( + text, + truncation=True, + padding="longest", + max_length=_MAX_LENGTH + _DROP_IDX - _SUFFIX_START_IDX, + return_tensors="pt", + ).to(device) + input_ids = torch.cat([inputs["input_ids"], suffix_ids], dim=1) + mask = torch.cat([inputs["attention_mask"].bool(), suffix_mask], dim=1) + + states = self.text_encoder( + input_ids=input_ids, attention_mask=mask, output_hidden_states=True + ) + hiddens = torch.stack([states.hidden_states[i] for i in _SELECT_LAYERS], dim=2) + hiddens = hiddens[:, _DROP_IDX:] + mask = mask[:, _DROP_IDX:] + return hiddens.to(dtype), mask + + @torch.no_grad() + def forward(self, batch: Req, server_args: ServerArgs) -> Req: + device = get_local_torch_device() + pipeline_config = server_args.pipeline_config + arch = pipeline_config.dit_config.arch_config + dtype = torch.bfloat16 + + prompts = batch.prompt if isinstance(batch.prompt, list) else [batch.prompt] + n = len(prompts) + height, width = int(batch.height), int(batch.width) + patch = arch.patch + vsf = pipeline_config.get_vae_scale_factor() + + # Text conditioning (positive + negative for CFG). The residency manager + # pages the encoder to GPU for this block only, then offloads it for the + # denoise loop (frees ~8GB) when --text-encoder-cpu-offload is set. + neg_prompts = ( + batch.negative_prompt + if isinstance(batch.negative_prompt, list) + else [batch.negative_prompt or ""] * n + ) + with self.use_declared_component( + component_name="text_encoder", module=self.text_encoder + ) as text_encoder: + self.text_encoder = text_encoder + prompt_embeds, prompt_mask = self._encode(prompts, device, dtype) + neg_embeds, neg_mask = self._encode(neg_prompts, device, dtype) + + # Initial noise latents, packed to [B, S_img, channels*patch**2]. + seed = batch.seed if batch.seed is not None else 0 + lat_h, lat_w = height // vsf, width // vsf + noise = torch.cat( + [ + torch.randn( + 1, + arch.channels, + lat_h, + lat_w, + device=device, + dtype=dtype, + generator=torch.Generator(device=device).manual_seed(seed + i), + ) + for i in range(n) + ], + dim=0, + ) + latents = rearrange( + noise, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=patch, pw=patch + ) + + # Rectified-flow timestep schedule with resolution-dependent time-shift mu. + # The repo ships a FlowMatchEulerDiscreteScheduler; drive it with the Flux + # sigma grid linspace(1, 1/n, n) so its shifted sigmas match the K2 sampler. + num_inference_steps = batch.num_inference_steps + image_seq_len = (lat_h // patch) * (lat_w // patch) + mu = pipeline_config.compute_mu(image_seq_len) + scheduler = self.scheduler + sigmas = np.linspace(1.0, 1.0 / num_inference_steps, num_inference_steps) + scheduler.set_timesteps( + num_inference_steps, device=device, mu=mu, sigmas=sigmas + ) + + batch.prompt_embeds = [prompt_embeds] + batch.prompt_embeds_mask = [prompt_mask] + batch.negative_prompt_embeds = [neg_embeds] + batch.negative_prompt_embeds_mask = [neg_mask] + batch.latents = latents + batch.raw_latent_shape = latents.shape + batch.scheduler = scheduler + # The DiT's TimeEmbed applies its own 1000x, so feed it the [0,1] sigmas + # (the diffusers scheduler reports timesteps on the 0..1000 scale instead). + batch.timesteps = scheduler.sigmas[:num_inference_steps] + batch.num_inference_steps = num_inference_steps + batch.sigmas = None + batch.generator = torch.Generator(device=device).manual_seed(seed) + return batch