[GLM-5.3-Flash] Reduce KPool planning synchronization and overlap indexer preparation (#39695)

Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
Yuxuan Zhang
2026-09-19 23:51:17 -07:00
committed by GitHub
co-authored by Xinyuan Tong
parent c8eb54c41d
commit 9f21fbc34b
9 changed files with 628 additions and 51 deletions
@@ -705,6 +705,7 @@ class TboForwardBatchPreparer:
for key in [
"req_pool_indices",
"req_pool_indices_cpu",
"seq_lens",
"seq_lens_cpu",
"extend_seq_lens",
@@ -14,6 +14,7 @@ from sglang.srt.layers.attention.dsa.dsa_indexer import (
rotate_activation,
)
from sglang.srt.layers.attention.dsa.dsa_topk_backend import TopkTransformMethod
from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp
from sglang.srt.layers.layernorm import LayerNorm
from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.utils import add_prefix, ceil_align, is_cuda, is_hip, is_npu
@@ -163,6 +164,36 @@ class IndexerKPool(MultiPlatformOp):
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
return weights
@torch.compile(dynamic=True)
def _project_and_scale_head_gates(self, x: torch.Tensor) -> torch.Tensor:
weights, _ = self.weights_proj(x.float())
return weights * self.n_heads**-0.5
@torch.compile(dynamic=True)
def _apply_q_scale_and_softmax_scale(
self, weights: torch.Tensor, q_scale: torch.Tensor
) -> torch.Tensor:
return weights.unsqueeze(-1) * q_scale * self.softmax_scale
def _resolve_head_gate_weights(self, x, q_scale, head_weights):
if head_weights is not None:
return self._apply_q_scale_and_softmax_scale(head_weights, q_scale)
return self._get_logits_head_gate(x, q_scale)
def _can_overlap_prefill(
self,
forward_batch: ForwardBatch,
return_indices: bool,
) -> bool:
return (
self.alt_stream is not None
and return_indices
and forward_batch.forward_mode.is_extend_without_speculative()
and not get_is_capture_mode()
and not is_in_breakable_cuda_graph()
and not dsa_use_prefill_cp(forward_batch)
)
@staticmethod
def _get_index_k_read_buffer(pool, layer_id: int) -> torch.Tensor:
if hasattr(pool, "get_broadcastable_index_k_with_scale_buffer"):
@@ -538,8 +569,11 @@ class IndexerKPool(MultiPlatformOp):
enable_dual_stream: bool,
forward_batch: ForwardBatch,
precompute_compress_gate: bool = False,
precompute_head_gate: bool = False,
):
gate_score = None
head_weights = None
apply_rope = not self.skip_rope and self.rope_head_dim > 0
if enable_dual_stream:
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
@@ -557,6 +591,10 @@ class IndexerKPool(MultiPlatformOp):
[self.rope_head_dim, self.head_dim - self.rope_head_dim],
dim=-1,
)
if precompute_head_gate:
head_weights = self._project_and_scale_head_gates(x)
if not apply_rope:
query = rotate_activation(query)
with torch.cuda.stream(self.alt_stream):
key, _ = self.wk(x)
key = self.k_norm(key)
@@ -584,15 +622,16 @@ class IndexerKPool(MultiPlatformOp):
key, [self.rope_head_dim, self.head_dim - self.rope_head_dim], dim=-1
)
if not self.skip_rope:
if apply_rope:
q_rope, k_rope = self.rotary_emb(positions, q_rope, k_rope)
query[..., : self.rope_head_dim] = q_rope
key[..., : self.rope_head_dim] = k_rope
query = rotate_activation(query)
if apply_rope or not enable_dual_stream:
query = rotate_activation(query)
return query, key, gate_score
return query, key, gate_score, head_weights
def _get_k_bf16(
self,
@@ -1293,7 +1332,7 @@ class IndexerKPool(MultiPlatformOp):
assert plan is not None, "DSA kpool target_verify requires kpool_write_plan"
num_draft_tokens = plan.num_draft_tokens
query, key, gate_score_maybe = self._get_q_k_bf16(
query, key, gate_score_maybe, head_weights = self._get_q_k_bf16(
q_lora,
x,
positions,
@@ -1302,6 +1341,7 @@ class IndexerKPool(MultiPlatformOp):
precompute_compress_gate=(
enable_dual_stream and self.compress_gate_stream is not None
),
precompute_head_gate=enable_dual_stream and return_indices,
)
pool = get_token_to_kv_pool()
@@ -1340,7 +1380,7 @@ class IndexerKPool(MultiPlatformOp):
self.alt_stream.wait_stream(self.compress_gate_stream)
if return_indices:
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
weights = self._get_logits_head_gate(x, q_scale)
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
with torch.cuda.stream(self.alt_stream):
_compress_write()
current_stream.wait_stream(self.alt_stream)
@@ -1348,7 +1388,7 @@ class IndexerKPool(MultiPlatformOp):
_compress_write()
if return_indices:
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
weights = self._get_logits_head_gate(x, q_scale)
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
if not return_indices:
return None
@@ -1476,44 +1516,24 @@ class IndexerKPool(MultiPlatformOp):
and forward_batch.forward_mode.is_decode_or_idle()
and self.compress_gate_stream is not None
)
query, key, gate_score = self._get_q_k_bf16(
query, key, gate_score, head_weights = self._get_q_k_bf16(
q_lora,
x,
positions,
enable_dual_stream,
forward_batch=forward_batch,
precompute_compress_gate=precompute_compress_gate,
precompute_head_gate=enable_dual_stream and return_indices,
)
weights = None
kpool_extend_cache = None
if enable_dual_stream and forward_batch.forward_mode.is_decode_or_idle():
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
if gate_score is not None:
self.alt_stream.wait_stream(self.compress_gate_stream)
with torch.cuda.stream(self.alt_stream):
self._compress_write(
x=x,
key=key,
positions=positions,
forward_batch=forward_batch,
layer_id=layer_id,
metadata=metadata,
gate_score=gate_score,
)
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
weights = self._get_logits_head_gate(x, q_scale)
current_stream.wait_stream(self.alt_stream)
else:
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
has_kpool_extend_plan = metadata.attn_metadata.kpool_extend_plan is not None
defer_kpool_cache_write = (
forward_batch.forward_mode.is_extend_without_speculative()
and return_indices
and not has_kpool_extend_plan
)
kpool_extend_cache = self._compress_write(
has_kpool_extend_plan = metadata.attn_metadata.kpool_extend_plan is not None
is_prefill = forward_batch.forward_mode.is_extend_without_speculative()
defer_kpool_cache_write = (
is_prefill and return_indices and not has_kpool_extend_plan
)
def compress_write():
return self._compress_write(
x=x,
key=key,
positions=positions,
@@ -1521,20 +1541,33 @@ class IndexerKPool(MultiPlatformOp):
layer_id=layer_id,
metadata=metadata,
gate_score=gate_score,
return_compressed=(
forward_batch.forward_mode.is_extend_without_speculative()
and return_indices
),
return_compressed=is_prefill and return_indices,
write_cache=not defer_kpool_cache_write,
)
if (
forward_batch.forward_mode.is_extend_without_speculative()
and not return_indices
):
return None
if weights is None:
weights = self._get_logits_head_gate(x, q_scale)
overlap_decode = (
enable_dual_stream and forward_batch.forward_mode.is_decode_or_idle()
)
if overlap_decode or self._can_overlap_prefill(forward_batch, return_indices):
current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream)
if gate_score is not None:
self.alt_stream.wait_stream(self.compress_gate_stream)
with torch.cuda.stream(self.alt_stream):
kpool_extend_cache = compress_write()
if return_indices:
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
current_stream.wait_stream(self.alt_stream)
else:
if return_indices:
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
kpool_extend_cache = compress_write()
if return_indices:
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
if not return_indices:
return None
if is_cuda():
if (
@@ -245,7 +245,10 @@ def _kpool_cpu_plan(
if isinstance(extend_seq_lens_cpu, torch.Tensor):
extend_seq_lens_cpu = extend_seq_lens_cpu.tolist()
seq_lens_cpu = forward_batch.seq_lens_cpu.tolist()
req_pool_indices_cpu = forward_batch.req_pool_indices.tolist()
req_pool_indices_cpu = getattr(forward_batch, "req_pool_indices_cpu", None)
if req_pool_indices_cpu is None:
req_pool_indices_cpu = forward_batch.req_pool_indices
req_pool_indices_cpu = req_pool_indices_cpu.tolist()
_append_compress_rows(
plan,
@@ -411,7 +414,9 @@ def _kpool_plan_to_gpu(
if need_paged:
req_to_token = get_req_to_token_pool().req_to_token
ragged_paged_page_table_row_index = torch.repeat_interleave(
local_req_pool_indices.to(torch.int32), ragged_q_len_t
local_req_pool_indices.to(torch.int32),
ragged_q_len_t,
output_size=sum(cpu.ragged_q_len),
)
ragged_paged_page_table = req_to_token
+1 -2
View File
@@ -2348,8 +2348,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# A full prefill has one result but can span several run_batch calls.
split_prefill_start: Optional[Tuple[int, float]] = None
# CPU mirror of req_pool_indices; schedule-path only (used in overlap_utils,
# not read by ForwardBatch), stale in spec draft window
req_pool_indices_cpu: torch.Tensor = None # shape: [b], int64
# Forward-pass metrics
@@ -3770,6 +3768,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
prefix_lens=self.prefix_lens,
req_to_token_pool=self.req_to_token_pool,
req_pool_indices=self.req_pool_indices,
req_pool_indices_cpu=self.req_pool_indices_cpu,
model_config=self.model_config,
forward_mode=self.forward_mode,
out_cache_loc=self.out_cache_loc,
@@ -547,6 +547,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# === Borrowed from ScheduleBatch: host metadata (CPU lists / mirrors) ===
# Optional seq_lens on cpu (CPU mirror of seq_lens)
seq_lens_cpu: Optional[torch.Tensor] = None
# Fresh only for non-speculative extend; speculative modes use device slots.
req_pool_indices_cpu: Optional[torch.Tensor] = None
# For logprob
top_logprobs_nums: Optional[List[int]] = None
@@ -908,6 +910,11 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
seq_lens_sum=batch.seq_lens_sum,
# Inputs aliased by reference from ScheduleBatch
seq_lens_cpu=seq_lens_cpu,
req_pool_indices_cpu=(
getattr(batch, "req_pool_indices_cpu", None)
if batch.forward_mode.is_extend_without_speculative()
else None
),
orig_seq_lens=batch.orig_seq_lens,
out_cache_loc_dsv4=batch.out_cache_loc_dsv4,
engram_history=batch.engram_history,
@@ -1671,6 +1678,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
# Keep token-aligned inputs consistent after padding.
self.input_embeds = self._pad_tensor_to_size(self.input_embeds, num_tokens)
self.req_pool_indices = self._pad_tensor_to_size(self.req_pool_indices, bs)
if self.req_pool_indices_cpu is not None:
self.req_pool_indices_cpu = self._pad_tensor_to_size(
self.req_pool_indices_cpu, bs
)
if self.lora_ids is not None:
self.lora_ids.extend((bs - len(self.lora_ids)) * [None])
@@ -1836,6 +1847,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
self.positions = self.positions[: self._original_num_tokens]
self.seq_lens = self.seq_lens[:bs]
self.req_pool_indices = self.req_pool_indices[:bs]
if self.req_pool_indices_cpu is not None:
self.req_pool_indices_cpu = self.req_pool_indices_cpu[:bs]
if self.seq_lens_cpu is not None:
self.seq_lens_cpu = self.seq_lens_cpu[:bs]
@@ -1845,6 +1858,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
self.positions = self.positions[:num_tokens]
self.seq_lens = self.seq_lens[:bs]
self.req_pool_indices = self.req_pool_indices[:bs]
if self.req_pool_indices_cpu is not None:
self.req_pool_indices_cpu = self.req_pool_indices_cpu[:bs]
if self.seq_lens_cpu is not None:
self.seq_lens_cpu = self.seq_lens_cpu[:bs]
if logits_output.next_token_logits is not None:
@@ -1176,6 +1176,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
padded_view.seq_lens = s["seq_lens"][:r]
padded_view.seq_lens_cpu = self._full_cg_seq_lens_cpu
padded_view.req_pool_indices = s["req_pool_indices"][:r]
if getattr(forward_batch, "req_pool_indices_cpu", None) is not None:
padded_view.req_pool_indices_cpu = padded_view._pad_tensor_to_size(
forward_batch.req_pool_indices_cpu, r
)
padded_view.extend_seq_lens = s["extend_seq_lens"][:r]
padded_view.extend_prefix_lens = s["extend_prefix_lens"][:r]
padded_view.max_seq_len_override = static_forward_batch.max_seq_len_override