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
Original file line number Diff line number Diff line change
Expand Up @@ -231,6 +231,8 @@ def _apply_qwen3_qk_norm_rope(
head_dim: int,
cos_sin_cache: torch.Tensor,
rope_cache_positions: torch.Tensor,
*,
round_norm_before_rope: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
return apply_qk_norm_rope(
q=q,
Expand All @@ -242,6 +244,7 @@ def _apply_qwen3_qk_norm_rope(
is_neox=True,
positions=rope_cache_positions,
allow_strided_qk=True,
round_norm_before_rope=round_norm_before_rope,
)


Expand All @@ -256,6 +259,8 @@ def _apply_qwen3_qk_norm_rope_pack_kv(
head_dim: int,
cos_sin_cache: torch.Tensor,
rope_cache_positions: torch.Tensor,
*,
round_norm_before_rope: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
batch_size, suffix_tokens, _, _ = q.shape
prefix_tokens = k_prefix.shape[1]
Expand Down Expand Up @@ -283,6 +288,7 @@ def _apply_qwen3_qk_norm_rope_pack_kv(
eps=q_norm.variance_epsilon,
head_dim=head_dim,
rope_dim=cos_sin_cache.shape[-1],
round_norm_before_rope=round_norm_before_rope,
)
return q, packed_kv[0], packed_kv[1]

Expand Down Expand Up @@ -746,6 +752,7 @@ def forward(
cos_sin_cache: torch.Tensor,
rope_cache_positions: torch.Tensor,
use_fused_qk_norm_rope: bool,
round_norm_before_rope: bool = False,
) -> torch.Tensor:
"""Cross-attention from GEN to cached UND K/V.

Expand Down Expand Up @@ -795,6 +802,7 @@ def forward(
True,
q.dtype,
cos_sin_cache.dtype,
round_norm_before_rope=round_norm_before_rope,
pack_kv=True,
)
)
Expand All @@ -810,6 +818,7 @@ def forward(
self.head_dim,
cos_sin_cache,
rope_cache_positions,
round_norm_before_rope=round_norm_before_rope,
)
out = self.attn.forward(q, packed_k, packed_v)
elif use_fused_qk_norm_rope:
Expand All @@ -821,6 +830,7 @@ def forward(
self.head_dim,
cos_sin_cache,
rope_cache_positions,
round_norm_before_rope=round_norm_before_rope,
)
else:
q, k = _apply_qwen3_qk_norm_rope_split(
Expand Down Expand Up @@ -960,6 +970,7 @@ def forward(
cos_sin_cache: torch.Tensor,
rope_cache_positions: torch.Tensor,
use_fused_qk_norm_rope: bool,
round_norm_before_rope: bool = False,
residual: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
# Fused add+rmsnorm: each `(hidden_states, residual) = norm(...)`
Expand All @@ -979,6 +990,7 @@ def forward(
cos_sin_cache,
rope_cache_positions,
use_fused_qk_norm_rope,
round_norm_before_rope,
)

hidden_states, residual = self.post_attention_layernorm(hidden_states, residual)
Expand Down Expand Up @@ -1126,6 +1138,7 @@ def __init__(
self.rms_norm_eps = arch.rms_norm_eps
self.hidden_act = arch.hidden_act
self.rope_theta = arch.rope_theta
self._gen_layers_torch_compiled = False

# The checkpoint may override the activation (and thus the MLP weight
# layout), so bind the arch-derived mappings on the instance.
Expand Down Expand Up @@ -1618,11 +1631,21 @@ def forward(
vis_pos_ids = vis_pos_ids.view(
3, batch_size, self.sp_size, local_seq_len
)[:, :, self.sp_rank, :]
self.cached_gen_rope_inputs[cache_key] = (
cos_sin_gen, gen_rope_cache_positions = (
self.language_model.rotary_emb.build_rope_cache_inputs(
vis_pos_ids, cache_dtype=hidden_gen.dtype
)
)
if T == 1 and not self._gen_layers_torch_compiled:
# build_rope_cache_inputs already rounds through cache_dtype
# before returning FP32 storage. Keep that rounded cache in the
# activation dtype so the exact fused QKNorm+RoPE kernel can
# consume it without repeating the cast in every GEN layer.
cos_sin_gen = cos_sin_gen.to(hidden_gen.dtype)
self.cached_gen_rope_inputs[cache_key] = (
cos_sin_gen,
gen_rope_cache_positions,
)

cos_sin_gen, gen_rope_cache_positions = self.cached_gen_rope_inputs[cache_key]

Expand All @@ -1631,7 +1654,22 @@ def forward(
# fused add+rmsnorm path instead of separate add + norm kernels.
cached_kv_for_key = self.cached_kv[cache_key]
residual: torch.Tensor | None = None
use_fused_qk_norm_rope = T > 1
round_norm_before_rope = T == 1
use_fused_qk_norm_rope = T > 1 or (
hidden_gen.device.type == "cuda"
and not torch.compiler.is_compiling()
and not self._gen_layers_torch_compiled
and get_sp_world_size() == 1
and can_use_fused_inplace_qknorm_rope(
self.head_dim,
cos_sin_gen.shape[-1],
True,
hidden_gen.dtype,
cos_sin_gen.dtype,
round_norm_before_rope=True,
pack_kv=True,
)
)
for i, layer in enumerate(self.gen_layers):
k_und, v_und = cached_kv_for_key[i]
hidden_gen, residual = layer(
Expand All @@ -1641,6 +1679,7 @@ def forward(
cos_sin_gen,
gen_rope_cache_positions,
use_fused_qk_norm_rope,
round_norm_before_rope,
residual=residual,
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -870,6 +870,7 @@ def _maybe_enable_torch_compile(
len(gen_layers),
compile_kwargs,
)
transformer._gen_layers_torch_compiled = True
for i, layer in enumerate(gen_layers):
gen_layers[i] = torch.compile(layer, **compile_kwargs)
else:
Expand Down
93 changes: 93 additions & 0 deletions test/registered/kernels/ops/diffusion/test_qknorm_rope.py
Original file line number Diff line number Diff line change
Expand Up @@ -413,6 +413,99 @@ def test_qknorm_rope_pack_kv_matches_separate_ops() -> None:
assert torch.equal(packed_v_ref, packed_kv[1])


def test_qknorm_rope_pack_kv_preserves_split_bf16_rounding() -> None:
from sgl_kernel import rotary_embedding

from sglang.kernels.ops.diffusion.qknorm_rope import (
fused_qknorm_rope_pack_kv,
)
from sglang.kernels.ops.layernorm.norm import fused_inplace_qknorm

batch_size = 1
prefix_tokens, suffix_tokens = 17, 1024
num_q_heads, num_kv_heads, head_dim = 32, 8, 64
num_heads = num_q_heads + 2 * num_kv_heads
qkv = torch.randn(
batch_size,
suffix_tokens,
num_heads,
head_dim,
device=DEVICE,
dtype=DTYPE,
)
prefix_qkv = torch.randn(
batch_size,
prefix_tokens,
num_heads,
head_dim,
device=DEVICE,
dtype=DTYPE,
)
k_prefix = prefix_qkv[:, :, num_q_heads : num_q_heads + num_kv_heads]
v_prefix = prefix_qkv[:, :, num_q_heads + num_kv_heads :]
q_weight = torch.randn(head_dim, device=DEVICE, dtype=DTYPE)
k_weight = torch.randn(head_dim, device=DEVICE, dtype=DTYPE)
positions = torch.arange(
batch_size * suffix_tokens, device=DEVICE, dtype=torch.int64
)
cos_sin_cache = create_cos_sin_cache(head_dim, batch_size * suffix_tokens).to(DTYPE)

qkv_ref = qkv.clone()
q_ref = qkv_ref[:, :, :num_q_heads]
k_ref = qkv_ref[:, :, num_q_heads : num_q_heads + num_kv_heads]
v_ref = qkv_ref[:, :, num_q_heads + num_kv_heads :]
fused_inplace_qknorm(
q_ref.view(-1, num_q_heads, head_dim),
k_ref.view(-1, num_kv_heads, head_dim),
q_weight,
k_weight,
eps=1e-6,
)
rotary_embedding(
positions,
q_ref.view(-1, num_q_heads * head_dim),
k_ref.view(-1, num_kv_heads * head_dim),
head_dim,
cos_sin_cache,
True,
)
packed_k_ref = torch.cat([k_prefix, k_ref], dim=1)
packed_v_ref = torch.cat([v_prefix, v_ref], dim=1)

qkv_fused = qkv.clone()
q_fused = qkv_fused[:, :, :num_q_heads]
k_fused = qkv_fused[:, :, num_q_heads : num_q_heads + num_kv_heads]
v_fused = qkv_fused[:, :, num_q_heads + num_kv_heads :]
packed_kv = torch.empty(
2,
batch_size,
prefix_tokens + suffix_tokens,
num_kv_heads,
head_dim,
device=DEVICE,
dtype=DTYPE,
)
fused_qknorm_rope_pack_kv(
q_fused,
k_fused,
v_fused,
k_prefix,
v_prefix,
packed_kv,
q_weight,
k_weight,
cos_sin_cache,
positions,
is_neox=True,
rope_dim=head_dim,
round_norm_before_rope=True,
)

assert torch.equal(q_ref, q_fused)
assert torch.equal(packed_k_ref, packed_kv[0])
assert torch.equal(packed_v_ref, packed_kv[1])


def test_qknorm_rope_accepts_empty_token_dimension() -> None:
from sglang.kernels.ops.diffusion.qknorm_rope import fused_inplace_qknorm_rope

Expand Down
Loading