[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
|
controller = tree_cache.cache_controller
|
||||||
host_pool_group = controller.mem_pool_host
|
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(
|
draft_host_pool = _build_mha_mla_host_pool(
|
||||||
pool=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,
|
page_size=controller.page_size,
|
||||||
layout=server_args.hicache_mem_layout,
|
layout=server_args.hicache_mem_layout,
|
||||||
allocator_type=_get_allocator_type(server_args),
|
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.
|
# so that host indices stay 1-to-1 between target and draft KV caches.
|
||||||
primary_host_pool = tree_cache.cache_controller.mem_pool_host
|
primary_host_pool = tree_cache.cache_controller.mem_pool_host
|
||||||
host_pool_kwargs = dict(
|
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,
|
host_size=0,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
layout=server_args.hicache_mem_layout,
|
layout=server_args.hicache_mem_layout,
|
||||||
|
|||||||
@@ -7502,11 +7502,11 @@ class ServerArgs:
|
|||||||
"backup and the storage keys must become dcp_rank-aware "
|
"backup and the storage keys must become dcp_rank-aware "
|
||||||
"first. Run HiCache+DCP with L1/L2 only."
|
"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(
|
raise NotImplementedError(
|
||||||
"HiCache with --dcp-size > 1 does not support speculative "
|
"HiCache with --dcp-size > 1 only supports DSPARK speculative "
|
||||||
"decoding yet (the draft-model host pool has no DCP index "
|
"decoding; other draft-model host pools have no DCP index "
|
||||||
"translation)."
|
"translation."
|
||||||
)
|
)
|
||||||
if self.enable_lmcache:
|
if self.enable_lmcache:
|
||||||
raise NotImplementedError(
|
raise NotImplementedError(
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
from sglang.srt.mem_cache.hybrid_cache.hybrid_pool_assembler import (
|
||||||
_split_hicache_size,
|
_split_hicache_size,
|
||||||
@@ -54,6 +55,39 @@ class TestDraftSidecarPoolDispatch(CustomTestCase):
|
|||||||
self.assertEqual(specs, [])
|
self.assertEqual(specs, [])
|
||||||
self.assertEqual(entries, [])
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user