[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
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,65 +1516,58 @@ 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():
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,
forward_batch=forward_batch,
layer_id=layer_id,
metadata=metadata,
gate_score=gate_score,
return_compressed=is_prefill and return_indices,
write_cache=not defer_kpool_cache_write,
)
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):
self._compress_write(
x=x,
key=key,
positions=positions,
forward_batch=forward_batch,
layer_id=layer_id,
metadata=metadata,
gate_score=gate_score,
)
kpool_extend_cache = 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)
current_stream.wait_stream(self.alt_stream)
else:
if return_indices:
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(
x=x,
key=key,
positions=positions,
forward_batch=forward_batch,
layer_id=layer_id,
metadata=metadata,
gate_score=gate_score,
return_compressed=(
forward_batch.forward_mode.is_extend_without_speculative()
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
kpool_extend_cache = compress_write()
if return_indices:
weights = self._resolve_head_gate_weights(x, q_scale, head_weights)
if weights is None:
weights = self._get_logits_head_gate(x, q_scale)
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
@@ -0,0 +1,147 @@
"""CPU coverage for KPool request-slot selection and paged query rows."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.environ import envs
from sglang.srt.layers.attention.dsa import kpool_plan
from sglang.srt.layers.attention.dsa.dsa_topk_backend import TopkTransformMethod
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class _NoDeviceReadTensor(torch.Tensor):
def tolist(self):
raise AssertionError("Reading device request slots would synchronize")
def _batch(extend_lens=(3, 5), seq_lens=(5, 11), slots=(9, 2)):
indices = torch.tensor(slots, dtype=torch.int64)
return SimpleNamespace(
batch_size=len(slots),
extend_seq_lens_cpu=list(extend_lens),
seq_lens_cpu=torch.tensor(seq_lens, dtype=torch.int64),
seq_lens=torch.tensor(seq_lens, dtype=torch.int64),
req_pool_indices=indices.as_subclass(_NoDeviceReadTensor),
req_pool_indices_cpu=indices,
)
def _expected_plan():
return kpool_plan._KPoolCpuPlan(
pool_batch_idx=[0, 1],
pool_req=[9, 2],
pool_pool_id=[0, 1],
pool_n_from_tail=[2, 2],
pool_chunk_src=[0, 3],
pool_tail_logical_base=[0, 4],
tail_req=[9, 2],
tail_dst_logical_start=[4, 8],
tail_chunk_src=[2, 5],
tail_n_write=[1, 3],
ragged_q_len=[3, 5],
ragged_pool_pages=[1, 1],
cu_pages_excl=[0, 1],
cu_q_len_excl=[0, 3],
total_pool_pages=2,
)
class TestKPoolPlannerCpuMirror(unittest.TestCase):
def test_mirror_avoids_device_read_and_preserves_compression_rows(self):
for tensor_lengths in (False, True):
with self.subTest(tensor_lengths=tensor_lengths):
batch = _batch()
if tensor_lengths:
batch.extend_seq_lens_cpu = torch.tensor(batch.extend_seq_lens_cpu)
plan = kpool_plan._kpool_cpu_plan(batch, 4, 64)
self.assertEqual(plan, _expected_plan())
def test_absent_or_none_mirror_uses_device_slots(self):
for absent in (False, True):
with self.subTest(absent=absent):
batch = _batch()
batch.req_pool_indices = batch.req_pool_indices_cpu
if absent:
del batch.req_pool_indices_cpu
else:
batch.req_pool_indices_cpu = None
plan = kpool_plan._kpool_cpu_plan(batch, 4, 64)
self.assertEqual(plan, _expected_plan())
def test_empty_mirror_needs_no_device_read(self):
plan = kpool_plan._kpool_cpu_plan(_batch((), (), ()), 4, 64)
self.assertEqual(plan, kpool_plan._KPoolCpuPlan())
def test_paged_query_rows_use_local_lengths_and_explicit_output_size(self):
batch = _batch()
cpu_plan = kpool_plan._kpool_cpu_plan(
batch,
4,
64,
local_extend_seq_lens_cpu=[1, 2],
local_seq_lens_cpu=[3, 8],
)
expected = _expected_plan()
expected.ragged_q_len = [1, 2]
expected.ragged_pool_pages = [0, 1]
expected.cu_pages_excl = [0, 0]
expected.cu_q_len_excl = [0, 1]
expected.total_pool_pages = 1
self.assertEqual(cpu_plan, expected)
original_tensor = torch.tensor
def cpu_tensor(*args, **kwargs):
kwargs.pop("pin_memory", None)
return original_tensor(*args, **kwargs)
page_table = torch.zeros(2, 4, dtype=torch.int32)
req_to_token = torch.zeros(10, 256, dtype=torch.int32)
seq_lens = torch.tensor([3, 7, 8], dtype=torch.int32)
local_slots = torch.tensor([9, 2], dtype=torch.int64)
with (
envs.SGLANG_DSA_FUSE_TOPK.override(True),
patch.object(torch, "tensor", side_effect=cpu_tensor),
patch.object(
torch, "repeat_interleave", wraps=torch.repeat_interleave
) as repeat,
patch.object(
kpool_plan,
"kpool_build_ragged_layout",
return_value=(torch.empty(0), torch.empty(0), torch.empty(0)),
),
patch.object(kpool_plan, "dsa_use_prefill_cp", return_value=False),
patch.object(kpool_plan, "_RAGGED_SCRATCH_K_U8", None),
patch.object(kpool_plan, "_RAGGED_SCRATCH_K_SCALE", None),
patch.object(
kpool_plan,
"get_req_to_token_pool",
return_value=SimpleNamespace(req_to_token=req_to_token),
),
):
plan = kpool_plan._kpool_plan_to_gpu(
cpu_plan,
batch,
page_table,
page_table,
seq_lens,
local_slots,
4,
64,
TopkTransformMethod.PAGED,
)
repeat.assert_called_once()
self.assertEqual(repeat.call_args.kwargs["output_size"], 3)
self.assertEqual(plan.ragged_paged_page_table_row_index.tolist(), [9, 2, 2])
self.assertEqual(plan.ragged_paged_page_table_row_index.dtype, torch.int32)
self.assertIs(plan.ragged_paged_page_table, req_to_token)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,316 @@
"""CPU scheduling checks: stream dependencies, cache contracts and gate math."""
import unittest
from contextlib import contextmanager, nullcontext
from types import MethodType, SimpleNamespace
from unittest.mock import Mock, patch
import torch
import torch.nn.functional as F
from sglang.kernels.ops.attention.dsa import triton_kernel
from sglang.srt.layers.attention.dsa import dsa_indexer_kpool as indexer_module
from sglang.srt.layers.attention.dsa.dsa_indexer_kpool import IndexerKPool
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
def _eager(method):
return getattr(method, "_torchdynamo_orig_callable", method)
class _Stream:
def __init__(self, name, trace):
self.name = name
self.trace = trace
def wait_stream(self, other):
self.trace.append(("wait", self.name, other.name))
class _Streams:
def __init__(self):
self.trace = []
self.current = _Stream("main", self.trace)
self.alt = _Stream("alt", self.trace)
self.gate = _Stream("gate", self.trace)
@contextmanager
def use(self, stream):
previous = self.current
self.current = stream
try:
yield
finally:
self.current = previous
def record(self, operation):
self.trace.append((operation, self.current.name))
class TestKPoolStreamScheduling(unittest.TestCase):
def test_prefill_overlap_excludes_cp_and_graph_capture(self):
for mode in (ForwardMode.EXTEND, ForwardMode.DECODE, ForwardMode.TARGET_VERIFY):
for has_stream in (False, True):
for capture, breakable, cp in (
(False, False, False),
(True, False, False),
(False, True, False),
(False, False, True),
):
with (
self.subTest(
mode=mode,
stream=has_stream,
capture=capture,
breakable=breakable,
cp=cp,
),
patch.object(
indexer_module, "get_is_capture_mode", return_value=capture
),
patch.object(
indexer_module,
"is_in_breakable_cuda_graph",
return_value=breakable,
),
patch.object(
indexer_module, "dsa_use_prefill_cp", return_value=cp
),
):
indexer = SimpleNamespace(
alt_stream=object() if has_stream else None
)
actual = IndexerKPool._can_overlap_prefill(
indexer,
SimpleNamespace(forward_mode=mode),
return_indices=True,
)
self.assertEqual(
actual,
mode == ForwardMode.EXTEND
and has_stream
and not (capture or breakable or cp),
)
def test_precomputed_head_gate_matches_original_math(self):
torch.manual_seed(19)
for dtype in (torch.bfloat16, torch.float16, torch.float32):
x = torch.randn(11, 32).to(dtype)
matrix = torch.randn(8, 32)
q_scale = torch.rand(11, 8, 1)
indexer = SimpleNamespace(
weights_proj=lambda value: (F.linear(value, matrix), None),
n_heads=8,
softmax_scale=128**-0.5,
)
expected = _eager(IndexerKPool._get_logits_head_gate)(indexer, x, q_scale)
projected = _eager(IndexerKPool._project_and_scale_head_gates)(indexer, x)
actual = _eager(IndexerKPool._apply_q_scale_and_softmax_scale)(
indexer, projected, q_scale
)
torch.testing.assert_close(actual, expected, rtol=0, atol=0)
def test_projection_reordering_preserves_rope_and_third_stream(self):
for skip_rope in (False, True):
streams = _Streams()
x = torch.arange(32, dtype=torch.float32).reshape(4, 8)
indexer = SimpleNamespace(
alt_stream=streams.alt,
compress_gate_stream=streams.gate,
half_device_sm_count=8,
head_dim=4,
rope_head_dim=2,
skip_rope=skip_rope,
index_kpool_compress_gate=torch.ones(4, 8),
)
def project_q(value):
streams.record("project_q")
return value.clone(), None
def project_k(value):
streams.record("project_k")
return value[:, :4].clone(), None
def head_gate(value):
streams.record("head_gate")
return value[:, :2].clone()
def rotate(value):
streams.record("rotate")
return value.flip(-1)
def rope(positions, q, k):
streams.record("rope")
return q + 1, k + 2
indexer.wq_b = project_q
indexer.wk = project_k
indexer.k_norm = lambda value: value
indexer.rotary_emb = rope
indexer._project_and_scale_head_gates = head_gate
with (
patch.object(
torch.cuda, "current_stream", side_effect=lambda: streams.current
),
patch.object(torch.cuda, "stream", side_effect=streams.use),
patch.object(
indexer_module.deep_gemm_wrapper,
"configure_deep_gemm_num_sms",
return_value=nullcontext(),
),
patch.object(indexer_module, "rotate_activation", side_effect=rotate),
):
actual = IndexerKPool._get_q_k_bf16(
indexer,
x,
x,
torch.arange(4),
True,
None,
precompute_compress_gate=True,
precompute_head_gate=True,
)
trace = streams.trace[:]
expected = IndexerKPool._get_q_k_bf16(
indexer, x, x, torch.arange(4), False, None
)
torch.testing.assert_close(actual[0], expected[0])
torch.testing.assert_close(actual[1], expected[1])
torch.testing.assert_close(
actual[2], F.linear(x, indexer.index_kpool_compress_gate)
)
self.assertIsNotNone(actual[3])
self.assertIn(("wait", "gate", "main"), trace)
join = trace.index(("wait", "main", "alt"))
self.assertLess(trace.index(("head_gate", "main")), join)
if skip_rope:
self.assertLess(trace.index(("rotate", "main")), join)
self.assertNotIn(("rope", "main"), trace)
else:
self.assertGreater(trace.index(("rope", "main")), join)
self.assertGreater(trace.index(("rotate", "main")), join)
def test_prefill_waits_before_topk_and_preserves_deferred_cache(self):
for has_plan in (False, True):
for return_indices in (False, True):
for cp in (False, True):
for num_tokens in (0, 128, 8192):
with self.subTest(
plan=has_plan,
indices=return_indices,
cp=cp,
tokens=num_tokens,
):
self._run_prefill(has_plan, return_indices, cp, num_tokens)
def _run_prefill(self, has_plan, return_indices, cp, num_tokens):
streams = _Streams()
x = torch.ones(num_tokens, 8)
compressed = object()
calls = []
metadata = SimpleNamespace(
attn_metadata=SimpleNamespace(
kpool_extend_plan=object() if has_plan else None
)
)
batch = SimpleNamespace(
forward_mode=ForwardMode.EXTEND,
seq_lens_cpu=torch.tensor([8192 + num_tokens]),
)
prepare_qk = Mock(return_value=(x, x, None, None))
indexer = SimpleNamespace(
alt_stream=streams.alt,
compress_gate_stream=streams.gate,
index_topk=16,
index_kpool=4,
index_kpool_compress=True,
block_size=128,
scale_fmt=None,
_get_q_k_bf16=prepare_qk,
)
indexer._can_overlap_prefill = MethodType(
IndexerKPool._can_overlap_prefill, indexer
)
def compress(**kwargs):
streams.record("compress")
calls.append(kwargs)
return compressed
def quant(*args):
streams.record("quant")
return x, torch.ones(num_tokens, 1, 1)
def head(*args):
streams.record("head")
return x
def topk(*args, **kwargs):
streams.record("topk")
if not has_plan:
self.assertIs(kwargs["kpool_extend_cache"], compressed)
return x
indexer._compress_write = compress
indexer._resolve_head_gate_weights = head
indexer._get_topk_ragged = topk
indexer._get_topk_ragged_kpool_plan = topk
with (
patch.object(indexer_module, "is_cuda", return_value=True),
patch.object(indexer_module, "is_hip", return_value=False),
patch.object(indexer_module, "is_npu", return_value=False),
patch.object(indexer_module, "get_is_capture_mode", return_value=False),
patch.object(
indexer_module, "is_in_breakable_cuda_graph", return_value=False
),
patch.object(indexer_module, "dsa_use_prefill_cp", return_value=cp),
patch.object(
indexer_module,
"get_attn_backend",
return_value=SimpleNamespace(
get_indexer_metadata=lambda *args: metadata
),
),
patch.object(
torch.cuda, "current_stream", side_effect=lambda: streams.current
),
patch.object(torch.cuda, "stream", side_effect=streams.use),
patch.object(triton_kernel, "act_quant", side_effect=quant),
):
actual = IndexerKPool._forward_cuda_impl(
indexer, x, x, torch.arange(num_tokens), batch, 0, return_indices
)
self.assertEqual(calls[0]["return_compressed"], return_indices)
self.assertEqual(calls[0]["write_cache"], has_plan or not return_indices)
self.assertFalse(prepare_qk.call_args.args[3])
self.assertFalse(prepare_qk.call_args.kwargs["precompute_head_gate"])
overlap = not cp and return_indices
self.assertIn(("compress", "alt" if overlap else "main"), streams.trace)
if overlap:
self.assertLess(
streams.trace.index(("wait", "alt", "main")),
streams.trace.index(("compress", "alt")),
)
join = streams.trace.index(("wait", "main", "alt"))
self.assertGreater(join, streams.trace.index(("compress", "alt")))
self.assertGreater(join, streams.trace.index(("quant", "main")))
self.assertGreater(join, streams.trace.index(("head", "main")))
self.assertLess(join, streams.trace.index(("topk", "main")))
else:
self.assertFalse(any(event[0] == "wait" for event in streams.trace))
if return_indices:
self.assertIs(actual, x)
self.assertIn(("quant", "main"), streams.trace)
else:
self.assertIsNone(actual)
self.assertNotIn(("quant", "main"), streams.trace)
self.assertNotIn(("head", "main"), streams.trace)
self.assertNotIn(("topk", "main"), streams.trace)
if __name__ == "__main__":
unittest.main()
@@ -10,6 +10,7 @@ from sglang.srt.model_executor import forward_batch_info
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
prefill_graph_tolerates_sum_len,
)
@@ -93,6 +94,62 @@ class TestPrefillCudaGraphPadding(CustomTestCase):
forward_batch, num_qo_tokens=16
)
def test_full_replay_pads_request_slot_cpu_mirror(self):
for batch_size, has_mirror in ((2, True), (4, True), (2, False)):
with self.subTest(batch_size=batch_size, has_mirror=has_mirror):
runner = self._make_runner()
runner._is_full_backend = True
runner._capture_req_slots = 4
runner._full_cg_seq_lens_cpu = torch.full((4,), -1)
slots = torch.tensor([7, 2, 9, 5])[:batch_size]
seq_lens = torch.arange(1, batch_size + 1)
static_slots = torch.zeros(4, dtype=slots.dtype)
static_slots[:batch_size].copy_(slots)
static_lens = torch.zeros(4, dtype=seq_lens.dtype)
static_lens[:batch_size].copy_(seq_lens)
runner._prefill_static_buffers = {
"req_pool_indices": static_slots,
"seq_lens": static_lens,
"extend_seq_lens": static_lens.clone(),
"extend_prefix_lens": torch.zeros(4, dtype=torch.int64),
}
attn_backend = mock.Mock()
runner.model_runner = SimpleNamespace(attn_backend=attn_backend)
batch = ForwardBatch(
forward_mode=ForwardMode.EXTEND,
batch_size=batch_size,
input_ids=torch.arange(batch_size),
req_pool_indices=slots,
req_pool_indices_cpu=slots if has_mirror else None,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens,
out_cache_loc=torch.arange(batch_size),
seq_lens_sum=int(seq_lens.sum()),
)
runner._prepare_forward_metadata_for_replay(
batch, batch, shape_key=ShapeKey(size=16)
)
attn_backend.init_forward_metadata_out_graph.assert_called_once()
padded = attn_backend.init_forward_metadata_out_graph.call_args.args[0]
self.assertIsInstance(padded, ForwardBatch)
self.assertIsNot(padded, batch)
self.assertEqual(padded.batch_size, 4)
torch.testing.assert_close(padded.seq_lens_cpu, static_lens)
torch.testing.assert_close(padded.req_pool_indices, static_slots)
if has_mirror:
torch.testing.assert_close(
padded.req_pool_indices_cpu, static_slots
)
self.assertEqual(padded.req_pool_indices_cpu.device.type, "cpu")
self.assertIs(batch.req_pool_indices_cpu, slots)
else:
self.assertIsNone(padded.req_pool_indices_cpu)
self.assertEqual(batch.batch_size, batch_size)
torch.testing.assert_close(batch.seq_lens_cpu, seq_lens)
attn_backend.init_forward_metadata.assert_not_called()
def _megamoe_no_prefill_cp(self, graph_has_dp_gather=False):
return (
mock.patch.object(