[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:
co-authored by
Yangmin Li
YAMY
parent
929230a6f0
commit
5e9342d16f
@@ -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)],
|
||||
|
||||
Reference in New Issue
Block a user