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:
co-authored by
yhzhuang
Lianmin Zheng
Yonghao Zhuang
Cheng Wan
parent
d0730a0e8b
commit
5931fd60ee
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user