Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions docs/docs/sglang-diffusion/fused_kernels.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -172,10 +172,27 @@ Every kernel here only moves values (plus zero fill, plus at most one same-order
| `pack_qkv_destination_major` | Triton | Ulysses destination-major QKV pack |
| `varlen_pack_qkv`, `varlen_scatter_to_padded` | Triton | varlen gather/scatter around the masked attention path |
| `varlen_pack_segmented_qkv` | Triton | varlen gather from a virtual prefix/main Q/K/V sequence |
| `joint_qkv_cat` | Triton | Joy Image Edit's three image/text Q/K/V concatenations, including strided packed-V inputs |
| `causal_conv3d_cat_pad` | KDA (JIT CUDA) / Triton | causal Conv3d `cat` + `pad` |
| `cat_pad_channels_last_3d` | Triton | Wan causal VAE `cat + F.pad + contiguous` (three passes plus cache bookkeeping) in one pass |
| `dup_up3d_add` | Triton | `repeat_interleave + permute().contiguous() + add` |

Joy Image Edit uses the joint-copy path for eligible CUDA FP16/BF16 image
tensors of at least 32 MiB. It preserves image-first token order, copies
values without arithmetic, and verifies each new shape/stride signature
against the native concatenations before enabling it. Small inputs,
unsupported layouts, gradient-bearing inputs, and unverified signatures
during graph capture use the native path. This does not enable model-level
BCG support.

For large Hopper BF16 image Q/K with 32 heads of width 128, Joy also uses
the existing out-of-place QK-Norm + RoPE kernel to read the packed projection
directly. This removes two input copies while retaining the original CUDA
arithmetic and contiguous outputs. Each new signature is checked bitwise;
inputs remain intact if the operation fails. Other shapes and platforms,
compilation, unverified capture, and `SGLANG_ENABLE_FUSED_QKNORM_ROPE=0`
retain the original helper.

### Quantized layout producers

These kernels preserve the quantized checkpoint path's selected reference operation. They are not a claim that FP8 or NVFP4 is equivalent to an unquantized BF16 checkpoint.
Expand All @@ -197,6 +214,7 @@ Kernels are written against a specific eager chain in a specific model, so cover
| Qwen-Image | linear+GELU, select-0/1 LN modulation, added-QKV fusion, QK RMSNorm+RoPE+joint QKV writes, residual norm/modulate+NVFP4 producer |
| GLM-Image | LN+modulate, per-head qk LN, residual-gate add, linear+GELU |
| ERNIE-Image | RMSNorm+scale/shift, residual-gated variant, rotate-half RoPE, residual-gate add |
| Joy Image Edit | Bitwise image/text QKV concatenation; strided image QK-Norm + RoPE on Hopper |
| Z-Image | BF16-native RMSNorm scale / tanh-residual, per-head QK RMSNorm |
| Ideogram 4 | gate RMSNorm, SwiGLU, rotate-half RoPE, modulate, residual-gate add |
| LTX-2 | QK-norm + split RoPE, ada-values split, RMSNorm+modulate, modulate, residual-gate add, linear+GELU |
Expand Down
9 changes: 9 additions & 0 deletions python/sglang/kernels/ops/diffusion/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -294,6 +294,13 @@
_CUDA,
"HunyuanVideo QKV pack + RoPE.",
),
(
"diffusion.joint_qkv_cat",
KernelBackend.TRITON,
"layout.joint_qkv_cat_triton:joint_qkv_cat",
_CUDA,
"Concatenate image/text QKV views into joint attention inputs.",
),
(
"diffusion.rmsnorm_preserve_reduction",
KernelBackend.TRITON,
Expand Down Expand Up @@ -589,6 +596,8 @@
# Rotary embeddings and the QK-norm chains fused around them
"try_fused_flux2_qkv_epilogue": "sglang.kernels.kda_kernels.flux2_qkv_epilogue_jit",
"hunyuan_qkv_rope_pack": "rope.hunyuan_qkv_pack_triton",
"can_use_joint_qkv_cat": "layout.joint_qkv_cat_triton",
"joint_qkv_cat": "layout.joint_qkv_cat_triton",
"can_use_ltx2_qknorm_split_rope_cuda": "sglang.kernels.kda_kernels.ltx2_qknorm_split_rope_jit",
"ltx2_qknorm_split_rope_cuda": "sglang.kernels.kda_kernels.ltx2_qknorm_split_rope_jit",
"apply_ltx2_split_rotary_emb": "rope.ltx2_rotary_triton",
Expand Down
124 changes: 124 additions & 0 deletions python/sglang/kernels/ops/diffusion/layout/joint_qkv_cat_triton.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,124 @@
"""Copy image/text QKV views into contiguous joint attention inputs."""

import torch
import triton
import triton.language as tl


@triton.jit
def _joint_qkv_cat_kernel(
IQ,
IK,
IV,
TQ,
TK,
TV,
OUT,
IMAGE_TOKENS: tl.constexpr,
TEXT_TOKENS: tl.constexpr,
HIDDEN: tl.constexpr,
BATCH: tl.constexpr,
IMAGE_STRIDES: tl.constexpr,
TEXT_STRIDES: tl.constexpr,
BLOCK: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
component = tl.program_id(1)
sequence = IMAGE_TOKENS + TEXT_TOKENS
batch = row // sequence
token = row % sequence
cols = tl.arange(0, BLOCK)
mask = cols < HIDDEN
# Each program copies one Q, K or V row. Separate strides preserve packed
# V views when Q/K have already been normalized into contiguous tensors.
if component == 0:
image, text = IQ, TQ
ib, it = IMAGE_STRIDES[0], IMAGE_STRIDES[1]
tb, tt = TEXT_STRIDES[0], TEXT_STRIDES[1]
elif component == 1:
image, text = IK, TK
ib, it = IMAGE_STRIDES[2], IMAGE_STRIDES[3]
tb, tt = TEXT_STRIDES[2], TEXT_STRIDES[3]
else:
image, text = IV, TV
ib, it = IMAGE_STRIDES[4], IMAGE_STRIDES[5]
tb, tt = TEXT_STRIDES[4], TEXT_STRIDES[5]
if token < IMAGE_TOKENS:
value = tl.load(image + batch * ib + token * it + cols, mask, other=0)
else:
value = tl.load(
text + batch * tb + (token - IMAGE_TOKENS) * tt + cols,
mask,
other=0,
)
# No floating-point arithmetic: preserve signed zeros and NaN payloads.
tl.store(OUT + (component * BATCH * sequence + row) * HIDDEN + cols, value, mask)


def can_use_joint_qkv_cat(*inputs: torch.Tensor) -> bool:
if len(inputs) != 6 or torch.compiler.is_compiling() or torch.version.hip:
return False
first = inputs[0]
if first.ndim != 4 or not first.is_cuda:
return False
if first.dtype not in (torch.float16, torch.bfloat16):
return False
batch, tokens, heads, dim = first.shape
if min(batch, tokens, heads, dim) <= 0 or heads * dim > 8192:
return False
for index, value in enumerate(inputs):
expected = first.shape if index < 3 else inputs[3].shape
if (
value.ndim != 4
or value.shape != expected
or value.shape[0] != batch
or value.shape[2:] != (heads, dim)
or value.shape[1] <= 0
or value.dtype != first.dtype
or value.device != first.device
or value.requires_grad
or value.stride(-1) != 1
or value.stride(-2) != dim
):
return False
return True


def joint_qkv_cat(
img_q: torch.Tensor,
img_k: torch.Tensor,
img_v: torch.Tensor,
txt_q: torch.Tensor,
txt_k: torch.Tensor,
txt_v: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Concatenate three pairs of ``[B, S, H, D]`` tensors, image first.

Batch/token strides may differ across inputs; each head row is contiguous.
The three outputs occupy disjoint, contiguous regions of one allocation.
"""
inputs = (img_q, img_k, img_v, txt_q, txt_k, txt_v)
assert can_use_joint_qkv_cat(*inputs)
batch, image_tokens, heads, dim = img_q.shape
text_tokens = txt_q.shape[1]
output = torch.empty(
(3, batch, image_tokens + text_tokens, heads, dim),
device=img_q.device,
dtype=img_q.dtype,
)
image_strides = tuple(s for x in inputs[:3] for s in x.stride()[:2])
text_strides = tuple(s for x in inputs[3:] for s in x.stride()[:2])
with torch.cuda.device(img_q.device):
_joint_qkv_cat_kernel[(batch * (image_tokens + text_tokens), 3)](
*inputs,
output,
image_tokens,
text_tokens,
heads * dim,
batch,
image_strides,
text_strides,
BLOCK=triton.next_power_of_2(heads * dim),
num_warps=4,
)
return output.unbind(0)
181 changes: 166 additions & 15 deletions python/sglang/multimodal_gen/runtime/models/dits/joy_image.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,16 @@
# SPDX-License-Identifier: Apache-2.0

import math
import os
from functools import lru_cache
from typing import Any, Optional, Tuple

import torch
import torch.nn as nn
from einops import rearrange

from sglang.kernels.ops import diffusion as diffusion_kernels
from sglang.kernels.ops.diffusion.sites.bitexact_gate import BitExactFusionGate
from sglang.multimodal_gen.configs.models.dits.joy_image import JoyImageDiTConfig
from sglang.multimodal_gen.configs.models.fsdp import is_blocks_or_double_blocks
from sglang.multimodal_gen.runtime.distributed import (
Expand Down Expand Up @@ -47,6 +50,159 @@

logger = init_logger(__name__)
_MODULATION_FACTOR = 6
_JOY_QKV_CAT = BitExactFusionGate(
"Joy image/text QKV concatenation", per_signature=True
)
_JOY_IMAGE_QK_ROPE = BitExactFusionGate("Joy strided image QK RoPE", per_signature=True)


def _joy_image_qk_rope(
q: torch.Tensor,
k: torch.Tensor,
q_norm: RMSNorm,
k_norm: RMSNorm,
cache: torch.Tensor,
complex_freqs: Optional[torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
def reference():
return apply_qk_norm_with_optional_rope(
q=q.contiguous(),
k=k.contiguous(),
q_norm=q_norm,
k_norm=k_norm,
head_dim=q.shape[-1],
cos_sin_cache=cache,
freqs_complex=complex_freqs,
is_neox=False,
allow_inplace=True,
)

# Read the packed projection directly, using the same CUDA arithmetic as
# the original contiguous in-place operation. Keep unmeasured paths native.
if (
_JOY_IMAGE_QK_ROPE.disabled
or not q.is_cuda
or torch.version.hip
or torch.compiler.is_compiling()
or q.ndim != 4
or q.shape != k.shape
or q.shape[2:] != (32, 128)
or q.numel() < 16 * 1024 * 1024
or q.dtype != torch.bfloat16
or k.dtype != q.dtype
or k.device != q.device
or torch.cuda.get_device_capability(q.device) != (9, 0)
or any(x.stride() != (q.shape[1] * 12288, 12288, 128, 1) for x in (q, k))
or cache.ndim != 2
or cache.shape[1] != 128
or cache.shape[0] < q.shape[1]
or cache.dtype != torch.float32
or cache.device != q.device
or not cache.is_contiguous()
or q_norm.variance_epsilon != k_norm.variance_epsilon
or any(
norm.weight.shape != (128,)
or norm.weight.dtype != q.dtype
or norm.weight.device != q.device
or not norm.weight.is_contiguous()
for norm in (q_norm, k_norm)
)
or (
torch.is_grad_enabled()
and any(
x.requires_grad for x in (q, k, q_norm.weight, k_norm.weight, cache)
)
)
or os.getenv("SGLANG_ENABLE_FUSED_QKNORM_ROPE", "1").lower()
in {"0", "false", "off", "no"}
):
return reference()
sig = (q.device, tuple(q.shape), tuple(q.stride()), q_norm.variance_epsilon)
verified = _JOY_IMAGE_QK_ROPE.is_verified(sig)
if not verified and torch.cuda.is_current_stream_capturing():
return reference()
if not diffusion_kernels.can_use_fused_inplace_qknorm_rope(
128, 128, False, q.dtype, cache.dtype
):
return reference()
try:
q_out = torch.empty(q.shape, device=q.device, dtype=q.dtype)
k_out = torch.empty_like(q_out)
positions = torch.arange(q.shape[1], device=q.device, dtype=torch.int64)
if q.shape[0] != 1:
positions = positions.repeat(q.shape[0])
diffusion_kernels.fused_qknorm_rope_out_of_place(
q.view(-1, 32, 128),
k.view(-1, 32, 128),
q_out.view(-1, 32, 128),
k_out.view(-1, 32, 128),
q_norm.weight,
k_norm.weight,
cache,
positions,
is_neox=False,
eps=q_norm.variance_epsilon,
head_dim=128,
rope_dim=128,
)
except Exception as exc:
# The out-of-place operation leaves packed Q/K/V pristine, including
# when it has written part of an output before raising.
_JOY_IMAGE_QK_ROPE.on_exception(exc, logger=logger)
return reference()
out = (q_out, k_out)
if verified:
return out
return _JOY_IMAGE_QK_ROPE.accept_or_fallback(
out,
reference(),
sig=sig,
equal=lambda actual, expected: all(
torch.equal(a.view(torch.int16), b.view(torch.int16))
for a, b in zip(actual, expected, strict=True)
),
logger=logger,
)


def _joy_joint_qkv(*inputs: torch.Tensor) -> tuple[torch.Tensor, ...]:
def reference():
return tuple(torch.cat((inputs[i], inputs[i + 3]), dim=1) for i in range(3))

# Below 32 MiB per image tensor, eager dispatch costs more than the copy
# saves. Keep small resolutions and short sequence-parallel shards native.
if (
_JOY_QKV_CAT.disabled
or inputs[0].numel() < 16 * 1024 * 1024
or not inputs[0].is_cuda
or torch.version.hip
or torch.compiler.is_compiling()
or not diffusion_kernels.can_use_joint_qkv_cat(*inputs)
):
return reference()
sig = (inputs[0].device, inputs[0].dtype) + tuple(
(tuple(x.shape), tuple(x.stride())) for x in inputs
)
verified = _JOY_QKV_CAT.is_verified(sig)
if not verified and torch.cuda.is_current_stream_capturing():
return reference()
try:
out = diffusion_kernels.joint_qkv_cat(*inputs)
except Exception as exc:
_JOY_QKV_CAT.on_exception(exc, logger=logger)
return reference()
if verified:
return out
return _JOY_QKV_CAT.accept_or_fallback(
out,
reference(),
sig=sig,
equal=lambda actual, expected: all(
torch.equal(a.view(torch.int16), b.view(torch.int16))
for a, b in zip(actual, expected, strict=True)
),
logger=logger,
)


def fused_add_gate(
Expand Down Expand Up @@ -271,18 +427,13 @@ def forward(
raise ValueError(
f"Fused QK-Norm + RoPE kernel only supports float16/bfloat16, but got {img_q.dtype}"
)
img_q = img_q.contiguous()
img_k = img_k.contiguous()
img_q, img_k = apply_qk_norm_with_optional_rope(
q=img_q,
k=img_k,
q_norm=self.img_attn_q_norm,
k_norm=self.img_attn_k_norm,
head_dim=img_q.shape[-1],
cos_sin_cache=vis_freqs_cis,
freqs_complex=vis_complex_freqs,
is_neox=False,
allow_inplace=True,
img_q, img_k = _joy_image_qk_rope(
img_q,
img_k,
self.img_attn_q_norm,
self.img_attn_k_norm,
vis_freqs_cis,
vis_complex_freqs,
)
img_q, img_k = img_q.to(img_v), img_k.to(img_v)

Expand Down Expand Up @@ -315,9 +466,9 @@ def forward(
txt_q, txt_k = txt_q.to(txt_v), txt_k.to(txt_v)

# Attention
joint_query = torch.cat([img_q, txt_q], dim=1)
joint_key = torch.cat([img_k, txt_k], dim=1)
joint_value = torch.cat([img_v, txt_v], dim=1)
joint_query, joint_key, joint_value = _joy_joint_qkv(
img_q, img_k, img_v, txt_q, txt_k, txt_v
)
attn = self.attn(
joint_query,
joint_key,
Expand Down
Loading
Loading