[Feature] JIT Fused QK norm + qk norm clean up (#15835)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user