From 8dcfb3b5e71e05aa921866deae1f8b2686706309 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Mon, 24 Aug 2026 12:07:26 +0800 Subject: [PATCH] [diffusion] Fuse LongCat-Image QKNorm and interleaved RoPE (#35995) --- .../jit/csrc/diffusion/qknorm_rope.cuh | 72 ++++-- python/sglang/kernels/ops/diffusion/README.md | 2 +- .../ops/diffusion/rope/qknorm_rope_jit.py | 8 +- .../runtime/models/dits/longcat_image.py | 217 +++++++++++++++--- .../benchmark/diffusion/bench_qknorm_rope.py | 42 +++- .../ops/diffusion/test_model_fast_paths.py | 58 ++++- .../kernels/ops/diffusion/test_rope.py | 48 +++- 7 files changed, 394 insertions(+), 53 deletions(-) diff --git a/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh b/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh index 40c2b4fe0..43a2b51b3 100644 --- a/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh +++ b/python/sglang/kernels/jit/csrc/diffusion/qknorm_rope.cuh @@ -163,6 +163,32 @@ SGL_DEVICE T rotary_sub(T x, T cos, T y, T sin) { #endif } +template +SGL_DEVICE T rotary_add_fp32(T x, float cos, T y, float sin) { + const float x_fp32 = device::cast(x); + const float y_fp32 = device::cast(y); +#ifdef USE_ROCM + return device::cast(x_fp32 * cos + y_fp32 * sin); +#else + const float lhs = __fmul_rn(x_fp32, cos); + const float rhs = __fmul_rn(y_fp32, sin); + return device::cast(__fadd_rn(lhs, rhs)); +#endif +} + +template +SGL_DEVICE T rotary_sub_fp32(T x, float cos, T y, float sin) { + const float x_fp32 = device::cast(x); + const float y_fp32 = device::cast(y); +#ifdef USE_ROCM + return device::cast(x_fp32 * cos - y_fp32 * sin); +#else + const float lhs = __fmul_rn(x_fp32, cos); + const float rhs = __fmul_rn(-y_fp32, sin); + return device::cast(__fadd_rn(lhs, rhs)); +#endif +} + template < int64_t kHeadDim, int64_t kRopeDim, @@ -196,8 +222,8 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParamsT __grid_c !kIsNeox || (kRotaryLanes >= 2 && kRotaryLanes % 2 == 0), "NeoX fused qknorm+rope requires an even rotary lane count"); static_assert( - !kRoundNormBeforeRope || std::is_same_v, - "Rounded QKNorm+RoPE requires cache and activation dtypes to match"); + !kRoundNormBeforeRope || std::is_same_v || std::is_same_v, + "Rounded QKNorm+RoPE requires cache and activation dtypes to match or an FP32 cache"); using Packed = packed_t; using Storage = AlignedVector; @@ -307,8 +333,13 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParamsT __grid_c (kCacheHasFullWidth ? lane_id : lane_id % kHalfRotaryLanes) * kElemsPerThread + 2 * j + i; const auto cos = load_cache_value(cos_ptr, cache_idx); const auto sin = load_cache_value(sin_ptr, cache_idx); - values[i] = lane_id < kHalfRotaryLanes ? rotary_sub(values[i], cos, partner_values[i], sin) - : rotary_add(values[i], cos, partner_values[i], sin); + if constexpr (std::is_same_v) { + values[i] = lane_id < kHalfRotaryLanes ? rotary_sub_fp32(values[i], cos, partner_values[i], sin) + : rotary_add_fp32(values[i], cos, partner_values[i], sin); + } else { + values[i] = lane_id < kHalfRotaryLanes ? rotary_sub(values[i], cos, partner_values[i], sin) + : rotary_add(values[i], cos, partner_values[i], sin); + } } } } @@ -317,13 +348,22 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParamsT __grid_c #pragma unroll for (uint32_t j = 0; j < kVecSize; ++j) { auto& values = unpack(output_vec[j]); - const auto half_idx = lane_id * kElemsPerThread / 2 + j; - const auto cos = load_cache_value(cos_ptr, half_idx); - const auto sin = load_cache_value(sin_ptr, half_idx); + const auto cache_idx_0 = + kCacheHasFullWidth ? lane_id * kElemsPerThread + 2 * j : lane_id * kElemsPerThread / 2 + j; + const auto cache_idx_1 = kCacheHasFullWidth ? cache_idx_0 + 1 : cache_idx_0; + const auto cos_0 = load_cache_value(cos_ptr, cache_idx_0); + const auto sin_0 = load_cache_value(sin_ptr, cache_idx_0); + const auto cos_1 = load_cache_value(cos_ptr, cache_idx_1); + const auto sin_1 = load_cache_value(sin_ptr, cache_idx_1); const auto x = values[0]; const auto y = values[1]; - values[0] = rotary_sub(x, cos, y, sin); - values[1] = rotary_add(y, cos, x, sin); + if constexpr (std::is_same_v) { + values[0] = rotary_sub_fp32(x, cos_0, y, sin_0); + values[1] = rotary_add_fp32(y, cos_1, x, sin_1); + } else { + values[0] = rotary_sub(x, cos_0, y, sin_0); + values[1] = rotary_add(y, cos_1, x, sin_1); + } } } } @@ -383,11 +423,15 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParamsT __grid_c for (uint32_t i = 0; i < kElemsPerThread; i += 2) { const float x = elems[i]; const float y = elems[i + 1]; - const int half_idx = static_cast(lane_id * kElemsPerThread + i) / 2; - const float cos = cast(load_cache_value(cos_ptr, half_idx)); - const float sin = cast(load_cache_value(sin_ptr, half_idx)); - elems[i] = x * cos - y * sin; - elems[i + 1] = y * cos + x * sin; + const auto cache_idx_0 = + kCacheHasFullWidth ? lane_id * kElemsPerThread + i : (lane_id * kElemsPerThread + i) / 2; + const auto cache_idx_1 = kCacheHasFullWidth ? cache_idx_0 + 1 : cache_idx_0; + const float cos_0 = cast(load_cache_value(cos_ptr, cache_idx_0)); + const float sin_0 = cast(load_cache_value(sin_ptr, cache_idx_0)); + const float cos_1 = cast(load_cache_value(cos_ptr, cache_idx_1)); + const float sin_1 = cast(load_cache_value(sin_ptr, cache_idx_1)); + elems[i] = x * cos_0 - y * sin_0; + elems[i + 1] = y * cos_1 + x * sin_1; } } } diff --git a/python/sglang/kernels/ops/diffusion/README.md b/python/sglang/kernels/ops/diffusion/README.md index 3dd406bde..e93bd831a 100644 --- a/python/sglang/kernels/ops/diffusion/README.md +++ b/python/sglang/kernels/ops/diffusion/README.md @@ -114,7 +114,7 @@ Several norms look interchangeable and are not. Start here. | Entry point | Backend | Contract | |---|---|---| -| `fused_inplace_qknorm_rope` | JIT CUDA | one bf16 rounding step vs split baseline; `round_norm_before_rope=True` makes it exact | +| `fused_inplace_qknorm_rope` | JIT CUDA | one bf16 rounding step vs split baseline; `round_norm_before_rope=True` makes it exact; supports compact and full-width NeoX/interleaved caches | | `fused_qknorm_rope_pack_kv` | JIT CUDA | as above, also packs prefix K/V | | `fused_rope_rotate_half_bitexact` | Triton | bit-exact (elementwise only) | | `fused_interleaved_rope_fp64` | JIT CUDA | bit-exact vs paired SANA-Video fp64 RoPE | diff --git a/python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py b/python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py index 32cd8ac58..7cf56d34b 100644 --- a/python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py +++ b/python/sglang/kernels/ops/diffusion/rope/qknorm_rope_jit.py @@ -96,15 +96,13 @@ def _can_use_fused_qknorm_rope( rotary_lanes, ) return False - elif cache_has_full_width: - logger.warning("Full-width cos/sin caches are only supported for NeoX RoPE") - return False if pack_kv and cache_has_full_width: logger.warning("KV packing does not support full-width cos/sin caches") return False - if round_norm_before_rope and cache_dtype != dtype: + if round_norm_before_rope and cache_dtype not in (dtype, torch.float32): logger.warning( - "Exact fused QKNorm+RoPE requires cache dtype %s to match activation dtype %s", + "Exact fused QKNorm+RoPE requires cache dtype %s to match activation " + "dtype %s or use float32", cache_dtype, dtype, ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py b/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py index 7257ade0b..9d22fe7fa 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/longcat_image.py @@ -11,11 +11,9 @@ and feeds them directly to the timestep embedder. The diffusers pipeline passes SGLang's DenoisingStage passes the raw timestep instead, so the value reaching the embedder is identical and no division is needed here. -Attention alignment: uses USPAttention (FA3/FA4 on Hopper/Blackwell) with -SGLang fused RMSNorm (apply_qk_norm). RoPE is applied separately via -diffusers apply_rotary_emb because LongCat's axes_dims_rope=[16,56,56] -sums to head_dim=128 (full rotation), which is incompatible with flashinfer's -cos_sin_cache format that requires rotary_dim <= head_dim. +Attention alignment: uses USPAttention (FA3/FA4 on Hopper/Blackwell) and the +SGLang fused QKNorm+RoPE kernel. LongCat's full-width, interleaved RoPE cache is +handled directly instead of materializing the Diffusers rotate-pair chain. """ from typing import List, Optional, Tuple @@ -34,9 +32,18 @@ from diffusers.models.normalization import ( AdaLayerNormZeroSingle, ) +from sglang.kernels.ops.diffusion import ( + BitExactFusionGate, + can_use_fused_inplace_qknorm_rope, + tensors_equal, +) 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.layernorm import RMSNorm, apply_qk_norm +from sglang.multimodal_gen.runtime.layers.layernorm import ( + RMSNorm, + apply_qk_norm, + apply_qk_norm_rope, +) from sglang.multimodal_gen.runtime.layers.linear import ( ColumnParallelLinear, RowParallelLinear, @@ -49,6 +56,124 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) +_LONGCAT_QKNORM_ROPE = BitExactFusionGate("LongCat fused QKNorm+RoPE") + + +def _longcat_qknorm_rope_reference( + q: torch.Tensor, + k: torch.Tensor, + q_norm: RMSNorm, + k_norm: RMSNorm, + head_dim: int, + image_rotary_emb: Tuple[torch.Tensor, torch.Tensor], +) -> Tuple[torch.Tensor, torch.Tensor]: + q, k = apply_qk_norm(q, k, q_norm, k_norm, head_dim) + q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1) + k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1) + return q, k + + +def _apply_longcat_qknorm_rope( + q: torch.Tensor, + k: torch.Tensor, + q_norm: RMSNorm, + k_norm: RMSNorm, + head_dim: int, + image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]], + cos_sin_cache: Optional[torch.Tensor], + positions: Optional[torch.Tensor], +) -> Tuple[torch.Tensor, torch.Tensor]: + if image_rotary_emb is None: + return apply_qk_norm(q, k, q_norm, k_norm, head_dim) + + q_eps = q_norm.variance_epsilon + k_eps = k_norm.variance_epsilon + can_fuse = ( + cos_sin_cache is not None + and positions is not None + and q.is_cuda + and not torch.compiler.is_compiling() + and q_eps == k_eps + and q.dtype in (torch.float16, torch.bfloat16) + and k.dtype == q.dtype + and q_norm.weight.dtype == q.dtype + and k_norm.weight.dtype == k.dtype + and q.is_contiguous() + and k.is_contiguous() + and can_use_fused_inplace_qknorm_rope( + head_dim=head_dim, + rope_dim=head_dim, + is_neox=False, + dtype=q.dtype, + cache_dtype=cos_sin_cache.dtype, + round_norm_before_rope=True, + cache_has_full_width=True, + ) + ) + verified = _LONGCAT_QKNORM_ROPE.verified + if ( + can_fuse + and not _LONGCAT_QKNORM_ROPE.disabled + and (verified or _LONGCAT_QKNORM_ROPE.can_attempt_once()) + ): + if q.shape[0] > 1: + positions = positions.repeat(q.shape[0]) + q_input = q.clone() if not verified else q + k_input = k.clone() if not verified else k + try: + out = apply_qk_norm_rope( + q=q, + k=k, + q_norm=q_norm, + k_norm=k_norm, + head_dim=head_dim, + cos_sin_cache=cos_sin_cache, + is_neox=False, + positions=positions, + round_norm_before_rope=True, + cache_has_full_width=True, + ) + except Exception as exc: + _LONGCAT_QKNORM_ROPE.on_exception(exc, logger=logger) + return _longcat_qknorm_rope_reference( + q_input, + k_input, + q_norm, + k_norm, + head_dim, + image_rotary_emb, + ) + else: + if verified: + return out + ref = _longcat_qknorm_rope_reference( + q_input, + k_input, + q_norm, + k_norm, + head_dim, + image_rotary_emb, + ) + return _LONGCAT_QKNORM_ROPE.accept_or_fallback( + out, + ref, + equal=tensors_equal, + logger=logger, + mismatch_msg=( + "LongCat fused QKNorm+RoPE is not bit-exact on this " + "platform; falling back to the Diffusers chain" + ), + ) + + return _longcat_qknorm_rope_reference( + q, + k, + q_norm, + k_norm, + head_dim, + image_rotary_emb, + ) + # --------------------------------------------------------------------------- # FFN @@ -108,9 +233,9 @@ class _LongCatFFN(nn.Module): class _LongCatJointAttention(nn.Module): """Double-stream (joint) attention for _TransformerBlock. - img and txt tokens are projected separately, QK-norm applied via SGLang - fused kernel, RoPE applied via diffusers apply_rotary_emb (supports full - head_dim rotation), then concatenated (txt first) before USPAttention. + img and txt tokens are projected separately, passed through fused QKNorm + and full-width interleaved RoPE, then concatenated (txt first) before + USPAttention. TP: Q/K/V and add_q/k/v use ColumnParallelLinear (heads sharded across TP ranks). Output projections use RowParallelLinear (all-reduce after matmul). @@ -200,6 +325,8 @@ class _LongCatJointAttention(nn.Module): hidden_states: torch.Tensor, encoder_hidden_states: torch.Tensor, image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + cos_sin_cache: Optional[torch.Tensor] = None, + positions: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: txt_seq_len = encoder_hidden_states.shape[1] @@ -217,10 +344,34 @@ class _LongCatJointAttention(nn.Module): ek = ek.unflatten(-1, (self.num_local_heads, self.head_dim)) ev = ev.unflatten(-1, (self.num_local_heads, self.head_dim)) - # SGLang fused QK-norm - q, k = apply_qk_norm(q, k, self.norm_q, self.norm_k, self.head_dim) - eq, ek = apply_qk_norm( - eq, ek, self.norm_added_q, self.norm_added_k, self.head_dim + if image_rotary_emb is None: + image_rotary_emb_txt = image_rotary_emb_img = None + else: + cos, sin = image_rotary_emb + image_rotary_emb_txt = (cos[:txt_seq_len], sin[:txt_seq_len]) + image_rotary_emb_img = (cos[txt_seq_len:], sin[txt_seq_len:]) + positions_txt = positions[:txt_seq_len] if positions is not None else None + positions_img = positions[txt_seq_len:] if positions is not None else None + + q, k = _apply_longcat_qknorm_rope( + q, + k, + self.norm_q, + self.norm_k, + self.head_dim, + image_rotary_emb_img, + cos_sin_cache, + positions_img, + ) + eq, ek = _apply_longcat_qknorm_rope( + eq, + ek, + self.norm_added_q, + self.norm_added_k, + self.head_dim, + image_rotary_emb_txt, + cos_sin_cache, + positions_txt, ) # Concatenate: txt first, then img (matches diffusers convention) @@ -228,12 +379,6 @@ class _LongCatJointAttention(nn.Module): k = torch.cat([ek, k], dim=1) v = torch.cat([ev, v], dim=1) - # RoPE applied after concat, over the full [txt+img] sequence. - # image_rotary_emb shape: [txt_len+img_len, head_dim] — matches q/k dim=1. - if image_rotary_emb is not None: - q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1) - k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1) - x = self.attn(q, k, v, num_replicated_prefix=txt_seq_len) x = x.flatten(2, 3).to(q.dtype) @@ -298,6 +443,8 @@ class _LongCatSingleAttention(nn.Module): self, hidden_states: torch.Tensor, image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, + cos_sin_cache: Optional[torch.Tensor] = None, + positions: Optional[torch.Tensor] = None, ) -> torch.Tensor: q, _ = self.to_q(hidden_states) k, _ = self.to_k(hidden_states) @@ -306,13 +453,16 @@ class _LongCatSingleAttention(nn.Module): k = k.unflatten(-1, (self.num_local_heads, self.head_dim)) v = v.unflatten(-1, (self.num_local_heads, self.head_dim)) - # SGLang fused QK-norm - q, k = apply_qk_norm(q, k, self.norm_q, self.norm_k, self.head_dim) - - # RoPE via diffusers (supports full head_dim rotation, sequence_dim=1) - if image_rotary_emb is not None: - q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1) - k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1) + q, k = _apply_longcat_qknorm_rope( + q, + k, + self.norm_q, + self.norm_k, + self.head_dim, + image_rotary_emb, + cos_sin_cache, + positions, + ) x = self.attn(q, k, v) return x.flatten(2, 3).to(q.dtype) @@ -400,6 +550,8 @@ class _SingleTransformerBlock(nn.Module): encoder_hidden_states: torch.Tensor, temb: torch.Tensor, image_rotary_emb=None, + cos_sin_cache=None, + positions=None, **kwargs, ): text_seq_len = encoder_hidden_states.shape[1] @@ -412,6 +564,8 @@ class _SingleTransformerBlock(nn.Module): attn_output = self.attn( hidden_states=norm_hidden_states, image_rotary_emb=image_rotary_emb, + cos_sin_cache=cos_sin_cache, + positions=positions, ) hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2) gate = gate.unsqueeze(1) @@ -462,6 +616,8 @@ class _TransformerBlock(nn.Module): encoder_hidden_states: torch.Tensor, temb: torch.Tensor, image_rotary_emb=None, + cos_sin_cache=None, + positions=None, **kwargs, ): norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1( @@ -475,6 +631,8 @@ class _TransformerBlock(nn.Module): hidden_states=norm_hidden_states, encoder_hidden_states=norm_encoder_hidden_states, image_rotary_emb=image_rotary_emb, + cos_sin_cache=cos_sin_cache, + positions=positions, ) attn_output = gate_msa.unsqueeze(1) * attn_output @@ -680,6 +838,9 @@ class LongCatImageTransformer2DModel(BaseDiT, LayerwiseOffloadableModuleMixin): image_rotary_emb = kwargs.get("image_rotary_emb") or self.pos_embed( torch.cat((txt_ids, img_ids), dim=0) ) + cos, sin = image_rotary_emb + cos_sin_cache = torch.cat((cos, sin), dim=-1).contiguous() + positions = torch.arange(cos.shape[0], device=cos.device, dtype=torch.int64) for block in self.transformer_blocks: encoder_hidden_states, hidden_states = block( @@ -687,6 +848,8 @@ class LongCatImageTransformer2DModel(BaseDiT, LayerwiseOffloadableModuleMixin): encoder_hidden_states=encoder_hidden_states, temb=temb, image_rotary_emb=image_rotary_emb, + cos_sin_cache=cos_sin_cache, + positions=positions, ) for block in self.single_transformer_blocks: @@ -695,6 +858,8 @@ class LongCatImageTransformer2DModel(BaseDiT, LayerwiseOffloadableModuleMixin): encoder_hidden_states=encoder_hidden_states, temb=temb, image_rotary_emb=image_rotary_emb, + cos_sin_cache=cos_sin_cache, + positions=positions, ) hidden_states = self.norm_out(hidden_states, temb) diff --git a/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py b/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py index 86aadcf79..3a73d39f3 100644 --- a/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py +++ b/test/registered/kernels/benchmark/diffusion/bench_qknorm_rope.py @@ -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"]), ) diff --git a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py index 5595bd4f5..f971997bf 100644 --- a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py +++ b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py @@ -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) # ------------------------------------------------------------------------- diff --git a/test/registered/kernels/ops/diffusion/test_rope.py b/test/registered/kernels/ops/diffusion/test_rope.py index 1889b1ba2..4f9bb1427 100644 --- a/test/registered/kernels/ops/diffusion/test_rope.py +++ b/test/registered/kernels/ops/diffusion/test_rope.py @@ -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,