test: cover CP HiCache host integration
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user