diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 2a110ca08..bb89bb894 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -705,6 +705,7 @@ class TboForwardBatchPreparer: for key in [ "req_pool_indices", + "req_pool_indices_cpu", "seq_lens", "seq_lens_cpu", "extend_seq_lens", diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py index 38805f7d1..c36f892c5 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer_kpool.py @@ -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 ( diff --git a/python/sglang/srt/layers/attention/dsa/kpool_plan.py b/python/sglang/srt/layers/attention/dsa/kpool_plan.py index 43d0f9c01..52497e881 100644 --- a/python/sglang/srt/layers/attention/dsa/kpool_plan.py +++ b/python/sglang/srt/layers/attention/dsa/kpool_plan.py @@ -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 diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 81ea87a14..bb0616c71 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -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, diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 40939377d..998a788cb 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -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: diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index bcee09dbe..3d0c2010a 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -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 diff --git a/test/registered/unit/layers/test_kpool_planner_cpu_mirror.py b/test/registered/unit/layers/test_kpool_planner_cpu_mirror.py new file mode 100644 index 000000000..2cf449750 --- /dev/null +++ b/test/registered/unit/layers/test_kpool_planner_cpu_mirror.py @@ -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() diff --git a/test/registered/unit/layers/test_kpool_stream_scheduling.py b/test/registered/unit/layers/test_kpool_stream_scheduling.py new file mode 100644 index 000000000..f43c5e6a4 --- /dev/null +++ b/test/registered/unit/layers/test_kpool_stream_scheduling.py @@ -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() diff --git a/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py b/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py index f1474f3dc..ed9ebdeda 100644 --- a/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py +++ b/test/registered/unit/model_executor/runner/test_prefill_cuda_graph_padding.py @@ -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(