[Feature] JIT Fused QK norm + qk norm clean up (#15835)

This commit is contained in:
DarkSharpness
2025-12-28 11:53:50 +08:00
committed by GitHub
parent 474a4699c5
commit 8e43980ebb
15 changed files with 827 additions and 127 deletions

View File

@@ -71,6 +71,7 @@ from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.utils import (
apply_qk_norm,
create_fused_set_kv_buffer_arg,
enable_fused_set_kv_buffer,
)
@@ -492,28 +493,6 @@ class LLaDA2MoeAttention(nn.Module):
self.alt_stream = alt_stream
def _apply_qk_norm(
self, q: torch.Tensor, k: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
# overlap qk norm
if self.alt_stream is not None and get_is_capture_mode():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.query_layernorm(q_by_head)
with torch.cuda.stream(self.alt_stream):
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.key_layernorm(k_by_head)
current_stream.wait_stream(self.alt_stream)
else:
q_by_head = q.reshape(-1, self.head_dim)
q_by_head = self.query_layernorm(q_by_head)
k_by_head = k.reshape(-1, self.head_dim)
k_by_head = self.key_layernorm(k_by_head)
q = q_by_head.view(q.shape)
k = k_by_head.view(k.shape)
return q, k
def forward(
self,
positions: torch.Tensor,
@@ -525,7 +504,14 @@ class LLaDA2MoeAttention(nn.Module):
qkv, _ = self.query_key_value(hidden_states)
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
if self.use_qk_norm:
q, k = self._apply_qk_norm(q, k)
q, k = apply_qk_norm(
q=q,
k=k,
q_norm=self.query_layernorm,
k_norm=self.key_layernorm,
head_dim=self.head_dim,
alt_stream=self.alt_stream,
)
q, k = self.rotary_emb(
positions,
q,