[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
@@ -163,6 +163,32 @@ SGL_DEVICE T rotary_sub(T x, T cos, T y, T sin) {
#endif
}
template <typename T>
SGL_DEVICE T rotary_add_fp32(T x, float cos, T y, float sin) {
const float x_fp32 = device::cast<fp32_t>(x);
const float y_fp32 = device::cast<fp32_t>(y);
#ifdef USE_ROCM
return device::cast<T>(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<T>(__fadd_rn(lhs, rhs));
#endif
}
template <typename T>
SGL_DEVICE T rotary_sub_fp32(T x, float cos, T y, float sin) {
const float x_fp32 = device::cast<fp32_t>(x);
const float y_fp32 = device::cast<fp32_t>(y);
#ifdef USE_ROCM
return device::cast<T>(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<T>(__fadd_rn(lhs, rhs));
#endif
}
template <
int64_t kHeadDim,
int64_t kRopeDim,
@@ -196,8 +222,8 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParamsT<kPackKV> __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<DType, CacheDType>,
"Rounded QKNorm+RoPE requires cache and activation dtypes to match");
!kRoundNormBeforeRope || std::is_same_v<DType, CacheDType> || std::is_same_v<CacheDType, fp32_t>,
"Rounded QKNorm+RoPE requires cache and activation dtypes to match or an FP32 cache");
using Packed = packed_t<DType>;
using Storage = AlignedVector<Packed, kVecSize>;
@@ -307,8 +333,13 @@ __global__ void fused_qknorm_rope_warp(const QKNormRopeParamsT<kPackKV> __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<CacheDType, fp32_t>) {
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<kPackKV> __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<CacheDType, fp32_t>) {
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<kPackKV> __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<int>(lane_id * kElemsPerThread + i) / 2;
const float cos = cast<fp32_t>(load_cache_value(cos_ptr, half_idx));
const float sin = cast<fp32_t>(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<fp32_t>(load_cache_value(cos_ptr, cache_idx_0));
const float sin_0 = cast<fp32_t>(load_cache_value(sin_ptr, cache_idx_0));
const float cos_1 = cast<fp32_t>(load_cache_value(cos_ptr, cache_idx_1));
const float sin_1 = cast<fp32_t>(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;
}
}
}
@@ -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 |
@@ -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,
)
@@ -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)