[PP][DeepSeek V4] Overlap communication and optimize SM120 prefill (#38792)

Co-authored-by: Yangmin Li <yangminl@nvidia.com>
Co-authored-by: YAMY <74099316+YAMY1234@users.noreply.github.com>
This commit is contained in:
jmydurant
2026-09-18 21:18:04 -07:00
committed by GitHub
co-authored by Yangmin Li YAMY
parent 929230a6f0
commit 5e9342d16f
17 changed files with 418 additions and 110 deletions
@@ -927,6 +927,7 @@ def _make_dsv4_draft(*, unified, mapping=None):
pool._unified_kv = unified
pool.compression_ratios = [0]
pool.page_size = 256
pool.swa_page_size = 256
pool.sliding_window = 128
pool.full_to_swa_index_mapping = mapping
pool.unified_swa_window = 128
@@ -941,7 +942,7 @@ def _make_dsv4_draft(*, unified, mapping=None):
)
else:
pool.swa_kv_pool = SimpleNamespace(
kv_buffer=[torch.empty((2, 16), dtype=torch.uint8)]
page_size=256, kv_buffer=[torch.empty((2, 16), dtype=torch.uint8)]
)
return pool
@@ -0,0 +1,102 @@
import unittest
from collections import defaultdict, deque
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import Mock, call
import torch
from sglang.srt.managers.scheduler_pp_mixin import SchedulerPPMixin
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class FakeStream:
def __init__(self, stream_id):
self.cuda_stream = stream_id
class FakeEvent:
def __init__(self):
self.recorded_stream = None
def record(self, stream):
self.recorded_stream = stream
def _make_scheduler(**attrs):
scheduler = object.__new__(SchedulerPPMixin)
scheduler.__dict__.update(attrs)
return scheduler
class TestPPCommOverlap(CustomTestCase):
def test_graph_proxy_send_records_forward_reuse_fence(self):
comm_stream = FakeStream(4)
work = Mock()
works = [SimpleNamespace(work=work)]
scheduler = _make_scheduler(
pp_comm_stream=comm_stream,
pp_comm_stream_ctx=nullcontext(),
pp_send_done_event=None,
device_module=SimpleNamespace(Event=FakeEvent),
)
scheduler._pp_commit_comm_work(works, fence_next_forward=True)
work.wait.assert_called_once_with()
self.assertEqual(works, [])
self.assertIs(scheduler.pp_send_done_event.recorded_stream, comm_stream)
def test_no_fence_event_without_comm_stream(self):
scheduler = _make_scheduler(
pp_comm_stream=None,
pp_comm_stream_ctx=nullcontext(),
pp_send_done_event=None,
)
scheduler._pp_commit_comm_work([SimpleNamespace(work=Mock())], True)
self.assertIsNone(scheduler.pp_send_done_event)
def test_forward_waits_for_graph_send_and_proxy_receive(self):
schedule_stream = FakeStream(1)
send_done_event = object()
recv_event = object()
forward_stream = Mock()
scheduler = _make_scheduler(
schedule_stream=schedule_stream,
forward_stream=forward_stream,
pp_send_done_event=send_done_event,
pp_proxy_recv_event=recv_event,
)
scheduler._pp_wait_forward_dependencies()
forward_stream.wait_stream.assert_called_once_with(schedule_stream)
self.assertEqual(
forward_stream.wait_event.call_args_list,
[call(send_done_event), call(recv_event)],
)
self.assertIsNone(scheduler.pp_send_done_event)
self.assertIsNone(scheduler.pp_proxy_recv_event)
def test_inbox_returns_original_receive_event(self):
recv_event = object()
tensor_dict = {"__msg_type__": "output", "value": torch.arange(2)}
scheduler = _make_scheduler(
_pp_tensor_dict_inbox=defaultdict(
deque, {"output": deque([(tensor_dict, recv_event)])}
),
)
received, event = scheduler._pp_recv_typed_dict("output")
self.assertIs(received, tensor_dict)
self.assertIs(event, recv_event)
if __name__ == "__main__":
unittest.main()
@@ -14,6 +14,7 @@ from sglang.srt.mem_cache.deepseek_v4_memory_pool import (
DeepSeekV4SingleKVPool,
DeepSeekV4TokenToKVPool,
_CompressedPoolConfig,
_num_dsv4_physical_kv_pages,
)
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
@@ -23,6 +24,40 @@ register_cpu_ci(est_time=5, suite="base-a-test-cpu")
class TestDSV4CompressedPools(CustomTestCase):
def test_physical_kv_pages_cover_reserved_logical_page(self):
size = 8192
self.assertEqual(_num_dsv4_physical_kv_pages(size, 256, 256), 33)
self.assertEqual(_num_dsv4_physical_kv_pages(size, 64, 256), 132)
self.assertGreaterEqual(
_num_dsv4_physical_kv_pages(size, 64, 256) * 64,
size + 256,
)
def test_state_buf_item_covers_one_logical_swa_page(self):
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
pool._unified_kv = False
pool.swa_page_size = 256
pool.compress_state_pools = []
pool.indexer_compress_state_pools = []
for physical_page_size in (256, 64):
with self.subTest(physical_page_size=physical_page_size):
row_bytes = physical_page_size * 4
buf = torch.empty((8, row_bytes), dtype=torch.uint8)
pool.swa_kv_pool = SimpleNamespace(
page_size=physical_page_size, kv_buffer=[buf]
)
data_ptrs, data_lens, item_lens = pool.get_state_buf_infos()
self.assertEqual(data_ptrs, [buf.data_ptr()])
self.assertEqual(data_lens, [buf.nbytes])
self.assertEqual(item_lens, [256 * 4])
def test_state_buf_infos_without_paged_swa(self):
pool = DeepSeekV4TokenToKVPool.__new__(DeepSeekV4TokenToKVPool)
pool.swa_kv_pool = None
pool.compress_state_pools = []
pool.indexer_compress_state_pools = []
self.assertEqual(pool.get_state_buf_infos(), ([], [], []))
def test_pp_mapping_and_pd_buffer_order(self):
for unified, stage_ratios in product(
(False, True), ([4, 0, 128, 4], [128], [0])
@@ -11,6 +11,7 @@ from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
_evict_swa_for_device_alloc,
_MambaStrategy,
_MambaSwaStrategy,
_require_single_row_dsv4_swa_pages,
_split_hicache_size,
_SwaStrategy,
build_full_draft_pools,
@@ -22,6 +23,23 @@ from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
class TestDeepSeekV4SWAPageLayout(CustomTestCase):
def test_split_physical_rows_are_rejected_for_hicache_consumers(self):
with self.assertRaisesRegex(ValueError, "direct SWA KV layout"):
_require_single_row_dsv4_swa_pages(
logical_page_size=256,
physical_page_size=64,
consumer="test consumer",
)
def test_matching_page_geometry_is_supported(self):
_require_single_row_dsv4_swa_pages(
logical_page_size=256,
physical_page_size=256,
consumer="test consumer",
)
class _Pool:
def __init__(self, kv_bytes):
self._kv_bytes = kv_bytes
@@ -179,7 +179,8 @@ class TestHybridDevicePoolAssembler(CustomTestCase):
kvcache.end_layer = 4
kvcache.swa_page_size = 2
kvcache.swa_kv_pool = SimpleNamespace(
kv_buffer=[torch.zeros((8, 3), dtype=torch.uint8) for _ in range(3)]
page_size=2,
kv_buffer=[torch.zeros((8, 3), dtype=torch.uint8) for _ in range(3)],
)
kvcache.c4_kv_pool = SimpleNamespace(
kv_buffer=[torch.zeros((8, 5), dtype=torch.uint8) for _ in range(2)],