Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 49 additions & 16 deletions python/sglang/srt/managers/cache_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -318,6 +318,11 @@ def __init__(
self.storage_host_pool = mem_pool_host
self.write_policy = write_policy
self.page_size = page_size
# Token granularity of one storage (L3) key. Defaults to the host pool
# page size; UnifiedRadixCache.init_hicache raises it to the radix-tree
# page size for compressed-DSA pools, where one tree page (the hash
# unit) spans multiple physical host pages ("span mode").
self.storage_page_size = page_size
self.io_backend = io_backend
self.enable_storage = False
self.storage_backend = None
Expand Down Expand Up @@ -744,6 +749,7 @@ def _generate_storage_config(
model_name=model_name,
tp_lcm_size=tp_lcm_size,
should_split_heads=should_split_heads,
storage_page_size=self.storage_page_size,
extra_config=storage_backend_extra_config,
)

Expand Down Expand Up @@ -1042,9 +1048,8 @@ def _page_get_zero_copy(
def _generic_page_get(
self, operation, hash_values, host_indices, extra_info=None
) -> int:
dummy_page_dst = [
self.storage_host_pool.get_dummy_flat_data_page() for _ in hash_values
]
pages_per_key = self.storage_page_size // self.page_size
dummy_page_dst = [self._span_dummy_page(pages_per_key) for _ in hash_values]
page_data = self.storage_backend.batch_get(hash_values, dummy_page_dst)
if page_data is None:
return 0
Expand All @@ -1057,13 +1062,36 @@ def _generic_page_get(
break
if operation.is_terminated():
break
self.storage_host_pool.set_from_flat_data_page(
host_indices[i * self.page_size],
self._restore_span_page(
host_indices[
i * self.storage_page_size : (i + 1) * self.storage_page_size
],
page_data[i],
pages_per_key,
)
count += 1
return count

def _span_dummy_page(self, pages_per_key: int) -> torch.Tensor:
if pages_per_key == 1:
return self.storage_host_pool.get_dummy_flat_data_page()
return torch.cat(
[
self.storage_host_pool.get_dummy_flat_data_page()
for _ in range(pages_per_key)
]
)

def _restore_span_page(
self, slots: torch.Tensor, data: torch.Tensor, pages_per_key: int
) -> None:
page_numel = data.numel() // pages_per_key
for j in range(pages_per_key):
self.storage_host_pool.set_from_flat_data_page(
int(slots[j * self.page_size]),
data[j * page_numel : (j + 1) * page_numel],
)

def _page_transfer(self, operation: PrefetchOperation) -> int:
# Transfer batch by batch
prefix_keys = operation.prefix_keys
Expand All @@ -1083,7 +1111,8 @@ def _page_transfer(self, operation: PrefetchOperation) -> int:
if all_success:
batch_hashes = operation.hash_value[i : i + STORAGE_BATCH_SIZE]
batch_host_indices = operation.host_indices[
i * self.page_size : (i + len(batch_hashes)) * self.page_size
i * self.storage_page_size : (i + len(batch_hashes))
* self.storage_page_size
]

# Get one batch token, and update the completed_tokens if succeed
Expand All @@ -1104,7 +1133,7 @@ def _page_transfer(self, operation: PrefetchOperation) -> int:
completed_pages += hit_pages
ack = PrefetchAck(
rid=operation.request_id,
completed_tokens=completed_pages * self.page_size,
completed_tokens=completed_pages * self.storage_page_size,
operation=operation,
)
self.prefetch_sync_queue.put(ack)
Expand Down Expand Up @@ -1198,7 +1227,7 @@ def _storage_hit_query(self, operation) -> tuple[list[str], int]:
storage_query_count = 0
hash_value = []
page_hashes = self.get_hash_str(
tokens_to_fetch, last_hash, page_size=self.page_size
tokens_to_fetch, last_hash, page_size=self.storage_page_size
)
operation.all_hash_values = page_hashes

Expand All @@ -1207,7 +1236,7 @@ def _storage_hit_query(self, operation) -> tuple[list[str], int]:
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])
storage_query_count += hit_page_num * self.page_size
storage_query_count += hit_page_num * self.storage_page_size
if hit_page_num < len(batch_hashes):
break
if prefix_keys and len(prefix_keys) > 0:
Expand Down Expand Up @@ -1241,7 +1270,7 @@ def prefetch_thread_func(self):
# Record the TP-synced hit count; the scheduler thread decides
# at drain time whether to revoke (below threshold) or allocate.
operation.hash_value = hash_value[
: (storage_hit_count // self.page_size)
: (storage_hit_count // self.storage_page_size)
]
operation.storage_hit_count = storage_hit_count
self.prefetch_hit_queue.put(operation)
Expand All @@ -1267,10 +1296,13 @@ def write_storage(

# todo: deprecate
def _generic_page_set(self, hash_values, host_indices, extra_info=None) -> bool:
data = [
self.storage_host_pool.get_data_page(host_indices[i * self.page_size])
for i in range(len(hash_values))
]
span = self.storage_page_size
pages_per_key = span // self.page_size
data = []
for i in range(len(hash_values)):
first_slots = host_indices[i * span : (i + 1) * span][:: self.page_size]
pages = [self.storage_host_pool.get_data_page(int(s)) for s in first_slots]
data.append(pages[0] if pages_per_key == 1 else torch.cat(pages))
return self.storage_backend.batch_set(hash_values, data)

def _page_set_zero_copy(self, hash_values, host_indices, extra_info=None) -> bool:
Expand All @@ -1285,7 +1317,8 @@ def _page_backup(self, operation):
for i in range(0, len(operation.hash_value), STORAGE_BATCH_SIZE):
batch_hashes = operation.hash_value[i : i + STORAGE_BATCH_SIZE]
batch_host_indices = operation.host_indices[
i * self.page_size : (i + len(batch_hashes)) * self.page_size
i * self.storage_page_size : (i + len(batch_hashes))
* self.storage_page_size
]
# Set one batch token, and record if success.
# todo: allow partial success
Expand All @@ -1299,7 +1332,7 @@ def _page_backup(self, operation):

if prefix_keys and len(prefix_keys) > 0:
prefix_keys += batch_hashes
operation.completed_tokens += self.page_size * len(batch_hashes)
operation.completed_tokens += self.storage_page_size * len(batch_hashes)

def backup_thread_func(self):
"""
Expand Down
83 changes: 65 additions & 18 deletions python/sglang/srt/mem_cache/hicache_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,10 @@ class HiCacheStorageConfig:
tp_lcm_size: Optional[int] = None
should_split_heads: bool = False
extra_config: Optional[dict] = None
# Token granularity of one storage key; 0 means "same as the host pool
# page size". Set to the radix-tree page size for compressed-DSA span
# mode, where one stored object holds multiple consecutive host pages.
storage_page_size: int = 0


@dataclass
Expand Down Expand Up @@ -165,6 +169,10 @@ class HiCacheStorage(ABC):
It abstracts the underlying storage mechanism, allowing different implementations to be used.
"""

# Whether the backend can store one key as multiple consecutive host
# pages (compressed-DSA span mode).
supports_page_spans: bool = False

# todo, the page size of storage backend does not have to be the same as the same as host memory pool
def register_mem_pool_host(self, mem_pool_host: HostKVCache):
self.mem_pool_host = mem_pool_host
Expand Down Expand Up @@ -371,6 +379,8 @@ def clear(self):


class HiCacheFile(HiCacheStorage):
supports_page_spans = True

def __init__(
self, storage_config: HiCacheStorageConfig, file_path: str = "/tmp/hicache"
):
Expand All @@ -386,6 +396,7 @@ def __init__(
)
attn_cp_rank = storage_config.attn_cp_rank
attn_cp_size = storage_config.attn_cp_size
self.storage_page_size = storage_config.storage_page_size or 0
model_name = "-".join(model_name.split("/")) if model_name else ""
enable_pp = pp_size > 1
self.config_suffix = f"_{model_name}"
Expand Down Expand Up @@ -663,45 +674,81 @@ def has_component(page_idx: int, name: str) -> bool:
def _log_key(self, pool_name: str, key: str) -> str:
return key if pool_name == PoolName.KV else f"{key}.{pool_name}"

def _read_page(self, pool_name: str, key: str, host_pool, page_offset: int) -> bool:
"""Read one page from storage into host_pool at page_offset."""
def _read_span(
self, pool_name: str, key: str, host_pool, slots: torch.Tensor, pool_page: int
) -> bool:
"""Read one storage object (one key) into host_pool pages."""
storage_key = self._log_key(pool_name, key)
data_page = self.get(storage_key, host_pool.get_dummy_flat_data_page())
first_slots = [int(s) for s in slots[::pool_page]]
dummy_page = host_pool.get_dummy_flat_data_page()
target = (
dummy_page
if len(first_slots) == 1
else torch.cat([host_pool.get_dummy_flat_data_page() for _ in first_slots])
)
data_page = self.get(storage_key, target)
if data_page is None:
return False
host_pool.set_from_flat_data_page(page_offset, data_page)
page_numel = dummy_page.numel()
for j, first_slot in enumerate(first_slots):
host_pool.set_from_flat_data_page(
first_slot, data_page[j * page_numel : (j + 1) * page_numel]
)
return True

def _write_page(
self, pool_name: str, key: str, host_pool, page_offset: int
def _write_span(
self, pool_name: str, key: str, host_pool, slots: torch.Tensor, pool_page: int
) -> bool:
"""Write one page from host_pool at page_offset to storage as raw bytes."""
"""Write one storage object (one key) from host_pool pages."""
storage_key = self._log_key(pool_name, key)
data_page = host_pool.get_data_page(page_offset, flat=True)
return self.set(storage_key, data_page)
first_slots = [int(s) for s in slots[::pool_page]]
pages = [host_pool.get_data_page(s, flat=True) for s in first_slots]
data = pages[0] if len(pages) == 1 else torch.cat(pages)
return self.set(storage_key, data)

def _batch_io_v2(self, transfers: List[PoolTransfer], op_fn):
results: dict[str, List[bool]] = {}
for transfer in transfers:
host_pool = self.registered_pools[transfer.name]
keys = transfer.keys or []
page_size = getattr(host_pool, "page_size", 1) or 1
expected = len(keys) * page_size
pool_page = getattr(host_pool, "page_size", 1) or 1
# Span mode: KV and KV-derived pools carry storage_page_size token
# slots per key (multiple consecutive host pages). Independent
# state pools (e.g. mamba) carry their own per-key entry count —
# one checkpoint slot per tree page — so the stride is derived
# from the transfer instead of assumed.
span = self.storage_page_size or pool_page
host_indices = transfer.host_indices

if host_indices is None or host_indices.numel() != expected:
kv_derived = (
transfer.name == PoolName.KV
or transfer.indices_from_pool == PoolName.KV
)
per_key = None
if host_indices is not None and keys:
per_key = host_indices.numel() // len(keys)
if per_key < 1 or host_indices.numel() != per_key * len(keys) or (
kv_derived and per_key != span
):
per_key = None
if per_key is None:
logger.error(
"%s indices length mismatch for %s: expected %s, got %s",
"%s indices length mismatch for %s: expected %s per key, got %s",
op_fn.__name__,
transfer.name,
expected,
span if kv_derived else "len(keys)-divisible",
host_indices.numel() if host_indices is not None else 0,
)
results[transfer.name] = [False] * len(keys)
continue

results[transfer.name] = [
op_fn(transfer.name, key, host_pool, host_indices[i * page_size].item())
op_fn(
transfer.name,
key,
host_pool,
host_indices[i * per_key : (i + 1) * per_key],
pool_page,
)
for i, key in enumerate(keys)
]
return results
Expand All @@ -711,14 +758,14 @@ def batch_get_v2(
transfers: List[PoolTransfer],
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> dict[str, List[bool]]:
return self._batch_io_v2(transfers, self._read_page)
return self._batch_io_v2(transfers, self._read_span)

def batch_set_v2(
self,
transfers: List[PoolTransfer],
extra_info: Optional[HiCacheStorageExtraInfo] = None,
) -> dict[str, List[bool]]:
return self._batch_io_v2(transfers, self._write_page)
return self._batch_io_v2(transfers, self._write_span)

def clear(self) -> bool:
try:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -577,7 +577,7 @@ def write_storage(

def _storage_hit_query(self, operation) -> tuple[list[str], int]:
hash_value = self.get_hash_str(
operation.token_ids, operation.last_hash, page_size=self.page_size
operation.token_ids, operation.last_hash, page_size=self.storage_page_size
)
operation.all_hash_values = hash_value

Expand All @@ -599,7 +599,7 @@ def _storage_hit_query(self, operation) -> tuple[list[str], int]:

return (
hash_value[:kv_hit_pages],
kv_hit_pages * self.page_size,
kv_hit_pages * self.storage_page_size,
)

def move_hybrid_indices(
Expand Down Expand Up @@ -711,7 +711,7 @@ def _page_backup(self, operation):
sidecar_ok = False
break
operation.completed_tokens = (
len(operation.hash_value) * self.page_size if sidecar_ok else 0
len(operation.hash_value) * self.storage_page_size if sidecar_ok else 0
)

def should_backup(self, transfer: PoolTransfer) -> bool:
Expand Down
15 changes: 15 additions & 0 deletions python/sglang/srt/mem_cache/storage/backend_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,21 @@ def _load_backend_class(
f"Class '{class_name}' not found in module '{module_path}': {e}"
) from e

@classmethod
def backend_supports_page_spans(cls, backend_name: str) -> bool:
entry = cls._registry.get(backend_name)
if entry is None:
return False
try:
backend_class = entry["loader"]()
except Exception:
logger.exception(
"Failed to load storage backend '%s' for span capability check",
backend_name,
)
return False
return bool(getattr(backend_class, "supports_page_spans", False))

@classmethod
def register_backend(cls, name: str, module_path: str, class_name: str) -> None:
"""Register a storage backend with lazy loading.
Expand Down
Loading
Loading