[diffusion] Fuse LongCat-Image QKNorm and interleaved RoPE (#35995)

This commit is contained in:
Xiaoyu Zhang
2026-08-24 12:07:26 +08:00
committed by GitHub
parent 09592f5889
commit 8dcfb3b5e7
7 changed files with 394 additions and 53 deletions
@@ -14,7 +14,7 @@ from sglang.kernels.jit.benchmark.utils import (
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(
est_time=13, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
est_time=15, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
)
MAX_SEQ_LEN = 131072
@@ -30,6 +30,8 @@ class CaseSpec:
head_dim: int
rope_dim: int
is_neox: bool
cache_has_full_width: bool = False
round_norm_before_rope: bool = False
BENCH_CASES = (
@@ -38,6 +40,7 @@ BENCH_CASES = (
CaseSpec("qwen_image_partial", 1, 4096, 32, 128, 64, False),
# Z-Image-Turbo default 1024x1024 config: dim=3840, num_heads=30 -> head_dim=128.
CaseSpec("zimage_1024", 1, 4096, 30, 128, 128, False),
CaseSpec("longcat_1024", 1, 4608, 24, 128, 128, False, True, True),
CaseSpec("batch2_medium", 2, 2048, 24, 128, 128, False),
)
CASE_BY_NAME = {case.name: case for case in BENCH_CASES}
@@ -46,7 +49,7 @@ CASE_NAMES = get_benchmark_range(
ci_range=[case.name for case in BENCH_CASES],
)
LINE_VALS = ["split", "fused"]
LINE_NAMES = ["JIT QKNorm + FlashInfer RoPE", "SGL JIT Fused QKNorm+RoPE"]
LINE_NAMES = ["Split QKNorm + RoPE", "SGL JIT Fused QKNorm+RoPE"]
STYLES = [("red", "-"), ("blue", "--")]
@@ -77,6 +80,13 @@ def make_inputs(case: CaseSpec) -> dict[str, torch.Tensor | bool]:
)
generator = torch.Generator(device=DEFAULT_DEVICE)
generator.manual_seed(seed)
cos_sin_cache = create_cos_sin_cache(case.rope_dim)
if case.cache_has_full_width:
cos, sin = cos_sin_cache.chunk(2, dim=-1)
cos_sin_cache = torch.cat(
(cos.repeat_interleave(2, dim=-1), sin.repeat_interleave(2, dim=-1)),
dim=-1,
).contiguous()
return {
"q": torch.randn(
case.batch_size * case.num_tokens,
@@ -114,8 +124,10 @@ def make_inputs(case: CaseSpec) -> dict[str, torch.Tensor | bool]:
dtype=torch.int64,
generator=generator,
),
"cos_sin_cache": create_cos_sin_cache(case.rope_dim),
"cos_sin_cache": cos_sin_cache,
"is_neox": case.is_neox,
"cache_has_full_width": case.cache_has_full_width,
"round_norm_before_rope": case.round_norm_before_rope,
}
@@ -128,7 +140,9 @@ def clone_inputs(
return out
def split_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None:
def split_qknorm_rope(
inputs: dict[str, torch.Tensor | bool],
) -> tuple[torch.Tensor, torch.Tensor] | None:
from flashinfer.rope import apply_rope_with_cos_sin_cache_inplace
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
@@ -142,6 +156,18 @@ def split_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None:
is_neox = bool(inputs["is_neox"])
fused_inplace_qknorm(q, k, q_weight, k_weight)
if inputs["cache_has_full_width"]:
cos, sin = cos_sin_cache.chunk(2, dim=-1)
cos = cos[positions]
sin = sin[positions]
def apply_interleaved(x: torch.Tensor) -> torch.Tensor:
x_real, x_imag = x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
x_rotated = torch.stack((-x_imag, x_real), dim=-1).flatten(-2)
return (x.float() * cos[:, None] + x_rotated * sin[:, None]).to(x.dtype)
return apply_interleaved(q), apply_interleaved(k)
apply_rope_with_cos_sin_cache_inplace(
positions=positions,
query=q.view(q.shape[0], -1),
@@ -163,7 +189,13 @@ def fused_qknorm_rope(inputs: dict[str, torch.Tensor | bool]) -> None:
inputs["cos_sin_cache"],
inputs["positions"],
is_neox=bool(inputs["is_neox"]),
rope_dim=inputs["cos_sin_cache"].shape[-1],
rope_dim=(
inputs["cos_sin_cache"].shape[-1] // 2
if inputs["cache_has_full_width"]
else inputs["cos_sin_cache"].shape[-1]
),
round_norm_before_rope=bool(inputs["round_norm_before_rope"]),
cache_has_full_width=bool(inputs["cache_has_full_width"]),
)
@@ -34,6 +34,7 @@ import sglang.multimodal_gen.runtime.models.dits.ernie_image as ernie_image
import sglang.multimodal_gen.runtime.models.dits.flux as flux
import sglang.multimodal_gen.runtime.models.dits.flux_2 as flux2
import sglang.multimodal_gen.runtime.models.dits.glm_image as glm_image
import sglang.multimodal_gen.runtime.models.dits.longcat_image as longcat_image
import sglang.multimodal_gen.runtime.models.dits.ltx_2 as ltx2_module
import sglang.multimodal_gen.runtime.models.dits.sana as sana
from sglang.kernels.ops.diffusion import (
@@ -56,7 +57,11 @@ from sglang.kernels.ops.diffusion.common.platform import is_cuda
from sglang.multimodal_gen.configs.models.vaes.stablediffusion3 import (
StableDiffusion3VAEConfig,
)
from sglang.multimodal_gen.runtime.layers.layernorm import RMSNorm, RMSNormNoWeight
from sglang.multimodal_gen.runtime.layers.layernorm import (
RMSNorm,
RMSNormNoWeight,
apply_qk_norm,
)
from sglang.multimodal_gen.runtime.layers.rotary_embedding.utils import (
_apply_rotary_emb,
)
@@ -85,6 +90,9 @@ from sglang.multimodal_gen.runtime.models.dits.hunyuanvideo import (
_hunyuan_pack_qkv,
_hunyuan_qknorm,
)
from sglang.multimodal_gen.runtime.models.dits.longcat_image import (
_apply_longcat_qknorm_rope,
)
from sglang.multimodal_gen.runtime.models.dits.ltx_2 import _ltx2_rms_norm_modulate
from sglang.multimodal_gen.runtime.models.dits.sana import (
_eager_ln_modulate as _sana_eager_ln_modulate,
@@ -490,6 +498,54 @@ def test_ernie_qknorm_rope_first_attempt_exception_uses_pristine_inputs():
assert ernie_image._ERNIE_QKNORM_ROPE.disabled
# -------------------------------------------------------------------------
# LongCat-Image -- full-width interleaved QKNorm + RoPE
# -------------------------------------------------------------------------
@requires_inline_ptx
def test_longcat_qknorm_rope_is_bit_exact():
torch.manual_seed(3)
batch, seq, heads, head_dim = 2, 17, 24, 128
offset = 11
q = torch.randn(batch, seq, heads, head_dim, device="cuda", dtype=torch.bfloat16)
k = torch.randn_like(q)
q_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16)
k_norm = RMSNorm(head_dim, eps=1e-6).to(device="cuda", dtype=torch.bfloat16)
with torch.no_grad():
q_norm.weight.copy_(torch.randn_like(q_norm.weight))
k_norm.weight.copy_(torch.randn_like(k_norm.weight))
cos = torch.randn(offset + seq, head_dim, device="cuda")
sin = torch.randn_like(cos)
image_rotary_emb = (cos[offset:], sin[offset:])
cache = torch.cat((cos, sin), dim=-1).contiguous()
positions = torch.arange(offset, offset + seq, device="cuda", dtype=torch.int64)
q_ref, k_ref = apply_qk_norm(q.clone(), k.clone(), q_norm, k_norm, head_dim)
q_ref = longcat_image.apply_rotary_emb(q_ref, image_rotary_emb, sequence_dim=1)
k_ref = longcat_image.apply_rotary_emb(k_ref, image_rotary_emb, sequence_dim=1)
q_fused, k_fused = q.clone(), k.clone()
q_out, k_out = _apply_longcat_qknorm_rope(
q_fused,
k_fused,
q_norm,
k_norm,
head_dim,
image_rotary_emb,
cache,
positions,
)
assert q_out.data_ptr() == q_fused.data_ptr()
assert k_out.data_ptr() == k_fused.data_ptr()
assert torch.equal(q_out, q_ref)
assert torch.equal(k_out, k_ref)
assert longcat_image._LONGCAT_QKNORM_ROPE.verified
assert not longcat_image._LONGCAT_QKNORM_ROPE.disabled
# -------------------------------------------------------------------------
# LTX-2 -- weightless RMSNorm + modulate (quality-gated)
# -------------------------------------------------------------------------
@@ -7,7 +7,8 @@ Two families with different oracles:
sgl_kernel RoPE). In the default mode the two differ by about one bf16
rounding step, so those cases use a tolerance; with
``round_norm_before_rope=True`` the fused kernel reproduces the split
rounding exactly and ``torch.equal`` applies.
rounding exactly and ``torch.equal`` applies. Full-width interleaved caches
use the Diffusers float32 RoPE chain as their oracle.
The LTX-2 split-RoPE kernel lives in ``test_rope_ltx2.py``: it is validated on
B200 and registered on that lane alone, which the cases here cannot share --
their oracle is the *split* baseline (a separate qknorm kernel plus sgl_kernel
@@ -270,6 +271,51 @@ def test_qknorm_rope_preserves_full_width_neox_cache() -> None:
assert torch.equal(k, k_ref)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_qknorm_rope_preserves_full_width_interleaved_cache(
dtype: torch.dtype,
) -> None:
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm
num_tokens, num_heads, head_dim = 257, 24, 128
q = torch.randn(num_tokens, num_heads, head_dim, device=DEVICE, dtype=dtype)
k = torch.randn_like(q)
q_weight = torch.randn(head_dim, device=DEVICE, dtype=dtype)
k_weight = torch.randn(head_dim, device=DEVICE, dtype=dtype)
positions = torch.randperm(num_tokens, device=DEVICE, dtype=torch.int64)
cos = torch.randn(num_tokens, head_dim, device=DEVICE)
sin = torch.randn_like(cos)
cache = torch.cat((cos, sin), dim=-1).contiguous()
def apply_interleaved_rope(x: torch.Tensor) -> torch.Tensor:
x_real, x_imag = x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
x_rotated = torch.stack((-x_imag, x_real), dim=-1).flatten(-2)
selected_cos = cos[positions, None]
selected_sin = sin[positions, None]
return (x.float() * selected_cos + x_rotated * selected_sin).to(dtype)
q_ref, k_ref = q.clone(), k.clone()
fused_inplace_qknorm(q_ref, k_ref, q_weight, k_weight, eps=1e-6)
q_ref = apply_interleaved_rope(q_ref)
k_ref = apply_interleaved_rope(k_ref)
fused_inplace_qknorm_rope(
q,
k,
q_weight,
k_weight,
cache,
positions,
is_neox=False,
eps=1e-6,
round_norm_before_rope=True,
cache_has_full_width=True,
)
assert torch.equal(q, q_ref)
assert torch.equal(k, k_ref)
def test_qknorm_rope_requires_opt_in_for_strided_packed_gqa() -> None:
from sglang.multimodal_gen.runtime.layers.layernorm import (
RMSNorm,