[HiCache] Support DCP with DSpark (#35221)
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -1013,9 +1013,11 @@ def build_full_draft_pools(
|
||||
controller = tree_cache.cache_controller
|
||||
host_pool_group = controller.mem_pool_host
|
||||
|
||||
# Note(kpham-sgl): DCP x DSpark draft KV is replicated and spans the virtual
|
||||
# loc space, so match the target host's logical_size instead of physical size.
|
||||
draft_host_pool = _build_mha_mla_host_pool(
|
||||
pool=pool,
|
||||
host_to_device_ratio=host_pool_group.size / pool.size,
|
||||
host_to_device_ratio=host_pool_group.logical_size / pool.size,
|
||||
page_size=controller.page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
allocator_type=_get_allocator_type(server_args),
|
||||
|
||||
@@ -117,7 +117,7 @@ def _register_legacy_hicache_draft(
|
||||
# so that host indices stay 1-to-1 between target and draft KV caches.
|
||||
primary_host_pool = tree_cache.cache_controller.mem_pool_host
|
||||
host_pool_kwargs = dict(
|
||||
host_to_device_ratio=primary_host_pool.size / pool.size,
|
||||
host_to_device_ratio=primary_host_pool.logical_size / pool.size,
|
||||
host_size=0,
|
||||
page_size=page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
|
||||
@@ -7502,11 +7502,11 @@ class ServerArgs:
|
||||
"backup and the storage keys must become dcp_rank-aware "
|
||||
"first. Run HiCache+DCP with L1/L2 only."
|
||||
)
|
||||
if self.speculative_algorithm is not None:
|
||||
if self.speculative_algorithm not in (None, "DSPARK"):
|
||||
raise NotImplementedError(
|
||||
"HiCache with --dcp-size > 1 does not support speculative "
|
||||
"decoding yet (the draft-model host pool has no DCP index "
|
||||
"translation)."
|
||||
"HiCache with --dcp-size > 1 only supports DSPARK speculative "
|
||||
"decoding; other draft-model host pools have no DCP index "
|
||||
"translation."
|
||||
)
|
||||
if self.enable_lmcache:
|
||||
raise NotImplementedError(
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
||||
_split_hicache_size,
|
||||
@@ -54,6 +55,39 @@ class TestDraftSidecarPoolDispatch(CustomTestCase):
|
||||
self.assertEqual(specs, [])
|
||||
self.assertEqual(entries, [])
|
||||
|
||||
def test_full_builder_sizes_sidecar_for_anchor_logical_space(self):
|
||||
draft_kv_pool = SimpleNamespace(layer_num=1, size=800)
|
||||
draft_host_pool = SimpleNamespace(layer_num=1)
|
||||
tree_cache = SimpleNamespace(
|
||||
cache_controller=SimpleNamespace(
|
||||
mem_pool_host=SimpleNamespace(size=100, logical_size=800),
|
||||
page_size=512,
|
||||
)
|
||||
)
|
||||
server_args = SimpleNamespace(hicache_mem_layout="page_first")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler."
|
||||
"_build_mha_mla_host_pool",
|
||||
return_value=draft_host_pool,
|
||||
) as build_host_pool,
|
||||
patch(
|
||||
"sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler."
|
||||
"_get_allocator_type",
|
||||
return_value="default",
|
||||
),
|
||||
):
|
||||
specs, entries = build_full_draft_pools(
|
||||
draft_kv_pool=draft_kv_pool,
|
||||
tree_cache=tree_cache,
|
||||
server_args=server_args,
|
||||
)
|
||||
|
||||
self.assertEqual(build_host_pool.call_args.kwargs["host_to_device_ratio"], 1.0)
|
||||
self.assertEqual(len(specs), 1)
|
||||
self.assertIs(entries[0].host_pool, draft_host_pool)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user