Support unified memory page-envelope transfers in PD (#39477)

Co-authored-by: yhzhuang <yhzhuang@fb.com>
Co-authored-by: Lianmin Zheng <lianminzheng@gmail.com>
Co-authored-by: Yonghao Zhuang <yhzhuang@users.noreply.github.com>
Co-authored-by: Cheng Wan <cheng.wan@radixark.ai>
This commit is contained in:
Yonghao Zhuang
2026-09-18 17:39:50 -07:00
committed by GitHub
co-authored by yhzhuang Lianmin Zheng Yonghao Zhuang Cheng Wan
parent d0730a0e8b
commit 5931fd60ee
27 changed files with 1012 additions and 247 deletions
@@ -17,6 +17,9 @@ from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.runtime_context import get_context, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.separate_buffer_allocator_double import (
bind_separate_buffer_capacity,
)
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
@@ -71,7 +74,8 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue.retracted_queue = reqs.copy()
queue.num_reserved_decode_tokens = 0
queue.req_to_token_pool = SimpleNamespace(available_size=lambda: len(reqs))
queue.token_to_kv_pool_allocator = SimpleNamespace(page_size=page_size)
queue.token_to_kv_pool_allocator = MagicMock(page_size=page_size)
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
queue.tree_cache = MagicMock()
queue.scheduler = SimpleNamespace(
sliding_window_size=2047,
@@ -81,6 +85,9 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue._swa_aware_allocatable_token_budgets = MagicMock(
return_value=(physical_available, physical_available)
)
queue._allocatable_token_budgets = MagicMock(
side_effect=lambda **_: physical_available
)
queue._swa_tail_allocatable_token_budget = MagicMock(
side_effect=lambda **_: physical_available
)
@@ -120,6 +127,10 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue.retracted_queue = []
queue._resolve_pending_reqs = MagicMock()
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
# `_uses_swa_reservation` consults the allocator once tail prealloc is
# off, so this abort path needs one even though it never allocates.
queue.token_to_kv_pool_allocator = MagicMock()
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
queue._allocatable_token_budgets = MagicMock(return_value=0)
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
@@ -175,6 +186,10 @@ class TestDecodeQueueCleanup(CustomTestCase):
queue._resolve_pending_reqs = MagicMock()
queue._update_handshake_waiters = MagicMock()
queue._uses_swa_tail_prealloc = MagicMock(return_value=False)
# `_uses_swa_reservation` consults the allocator once tail prealloc is
# off, so this abort path needs one even though it never allocates.
queue.token_to_kv_pool_allocator = MagicMock()
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
queue._allocatable_token_budgets = MagicMock(return_value=0)
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
@@ -234,6 +249,9 @@ class TestDecodeQueueCleanup(CustomTestCase):
)
queue._hicache_pending_restore_tokens = MagicMock(return_value=0)
queue._pre_alloc = MagicMock()
queue.token_to_kv_pool_allocator = MagicMock()
bind_separate_buffer_capacity(queue.token_to_kv_pool_allocator)
queue.tree_cache = MagicMock()
queue.req_to_token_pool = MagicMock()
queue.req_to_token_pool.available_size.return_value = 1
# Non-hybrid pools have no mamba allocator; MagicMock would otherwise
@@ -7,6 +7,7 @@ from types import SimpleNamespace
import numpy as np
from sglang.srt.disaggregation.ascend.conn import AscendKVManager
from sglang.srt.disaggregation.base.conn import StateType
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
from sglang.srt.disaggregation.prefill import _transfer_start_layer
@@ -15,6 +16,8 @@ from sglang.srt.disaggregation.utils import (
build_transfer_entry_pairs,
)
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool
from sglang.srt.runtime_context import get_memory, publish, reset_context
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -71,6 +74,7 @@ class TestTransferStartLayer(CustomTestCase):
class _RecordingKVManager:
get_mha_kv_ptrs_with_pp = CommonKVManager.get_mha_kv_ptrs_with_pp
get_mla_kv_ptrs_with_pp = CommonKVManager.get_mla_kv_ptrs_with_pp
def __init__(self, *, prefill_start_layer: int, pp_size: int):
self.is_mla_backend = False
@@ -143,6 +147,50 @@ class TestHybridSendUsesLayerIdPairing(CustomTestCase):
self._run_case(model_full_ids=ids, stage_full_ids=ids[:5], start_offset=0)
class TestSingleRegionSWATransfer(CustomTestCase):
def test_one_region_full_generates_transfer_block(self):
publish(ServerArgs(model_path="dummy"), role="tokenizer")
self.addCleanup(reset_context)
manager = _RecordingKVManager(prefill_start_layer=0, pp_size=1)
manager.kv_args.kv_data_ptrs = [1000]
manager.kv_args.kv_item_lens = [64]
manager.kv_args.kv_layer_ids = []
manager._validate_envelope_kv_layout = (
MooncakeKVManager._validate_envelope_kv_layout.__get__(manager)
)
manager._send_kvcache_generic = MooncakeKVManager._send_kvcache_generic.__get__(
manager
)
with get_memory().override(enable_unified_memory=True):
rc = MooncakeKVManager.send_kvcache(
manager,
mooncake_session_id="session",
prefill_kv_indices=np.array([3, 4], dtype=np.int32),
dst_kv_ptrs=[2000],
dst_kv_indices=np.array([7, 8], dtype=np.int32),
dst_kv_item_len=64,
executor=None,
)
self.assertEqual(rc, 0)
self.assertEqual(manager.blocks, [(1192, 2448, 128)])
def test_one_region_swa_generates_transfer_block(self):
manager = _RecordingKVManager(prefill_start_layer=0, pp_size=1)
rc = MooncakeKVManager._send_kvcache_generic(
manager,
mooncake_session_id="session",
src_data_ptrs=[1000],
dst_data_ptrs=[2000],
item_lens=[64],
prefill_data_indices=np.array([3, 4], dtype=np.int32),
dst_data_indices=np.array([7, 8], dtype=np.int32),
executor=None,
state_type=StateType.SWA,
)
self.assertEqual(rc, 0)
self.assertEqual(manager.blocks, [(1000 + 3 * 64, 2000 + 7 * 64, 2 * 64)])
class _RecordingAscendManager:
def __init__(self):
self.is_hybrid_mla_backend = True