[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:
co-authored by
Xinyuan Tong
parent
c8eb54c41d
commit
9f21fbc34b
@@ -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
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user