test: cover CP HiCache host integration

This commit is contained in:
2026-05-08 02:59:27 +08:00
parent 39f40c02b4
commit 3e9152194b
4 changed files with 41 additions and 15 deletions

View File

@@ -22,13 +22,12 @@ from typing import TYPE_CHECKING, List, NamedTuple, Optional
import torch
from sglang.srt.mem_cache.hicache_storage import (
HiCacheStorageConfig,
HiCacheStorageExtraInfo,
)
if TYPE_CHECKING:
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.hicache_storage import (
HiCacheStorageConfig,
HiCacheStorageExtraInfo,
)
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
from sglang.srt.distributed import (
@@ -926,9 +925,7 @@ class HiCacheController:
for page in pages:
self.host_mem_release_queue.put(page)
def _page_get_zero_copy(
self, operation, hash_values, host_indices, extra_info=None
):
def _page_get_zero_copy(self, operation, hash_values, host_indices, extra_info=None):
results = self.storage_backend.batch_get_v1(
hash_values, host_indices, extra_info
)
@@ -975,6 +972,8 @@ class HiCacheController:
]
prev_completed_tokens = operation.completed_tokens
# Get one batch token, and update the completed_tokens if succeed
from sglang.srt.mem_cache.hicache_storage import HiCacheStorageExtraInfo
extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys)
self.page_get_func(operation, batch_hashes, batch_host_indices, extra_info)
# Check termination
@@ -1036,6 +1035,8 @@ class HiCacheController:
batch_tokens[i : i + self.page_size], last_hash
)
batch_hashes.append(last_hash)
from sglang.srt.mem_cache.hicache_storage import HiCacheStorageExtraInfo
extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys)
hit_page_num = self.storage_backend.batch_exists(batch_hashes, extra_info)
hash_value.extend(batch_hashes[:hit_page_num])
@@ -1137,6 +1138,8 @@ class HiCacheController:
]
# Set one batch token, and record if success.
# todo: allow partial success
from sglang.srt.mem_cache.hicache_storage import HiCacheStorageExtraInfo
extra_info = HiCacheStorageExtraInfo(prefix_keys=prefix_keys)
success = self.page_set_func(batch_hashes, batch_host_indices, extra_info)
if not success:

View File

@@ -1,14 +1,18 @@
from __future__ import annotations
import hashlib
import logging
import os
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import Any, List, Optional
from typing import TYPE_CHECKING, Any, List, Optional
import torch
from sglang.srt.environ import envs
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
if TYPE_CHECKING:
from sglang.srt.mem_cache.memory_pool_host import HostKVCache
logger = logging.getLogger(__name__)

View File

@@ -34,11 +34,6 @@ from sglang.srt.mem_cache.memory_pool import (
MLATokenToKVPool,
NSATokenToKVPool,
)
from sglang.srt.mem_cache.memory_pool_host import (
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
NSATokenToKVPoolHost,
)
from sglang.srt.mem_cache.radix_cache import (
RadixCache,
RadixKey,
@@ -136,6 +131,12 @@ class HiRadixCache(RadixCache):
)
self.kv_cache = params.token_to_kv_pool_allocator.get_kvcache()
from sglang.srt.mem_cache.memory_pool_host import (
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
NSATokenToKVPoolHost,
)
if isinstance(self.kv_cache, MHATokenToKVPool):
self.token_to_kv_pool_host = MHATokenToKVPoolHost(
self.kv_cache,