Files
sglang/python/sglang/srt/disaggregation/decode_kvcache_offload_manager.py

236 lines
8.6 KiB
Python

from __future__ import annotations
import logging
import threading
import time
from typing import TYPE_CHECKING
import torch
from sglang.srt.managers.cache_controller import HiCacheController
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.memory_pool import (
MHATokenToKVPool,
MLATokenToKVPool,
ReqToTokenPool,
)
from sglang.srt.mem_cache.memory_pool_host import (
MHATokenToKVPoolHost,
MLATokenToKVPoolHost,
)
from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
logger = logging.getLogger(__name__)
class DecodeKVCacheOffloadManager:
"""Manage decode-side KV cache offloading lifecycle and operations."""
def __init__(
self,
req_to_token_pool: ReqToTokenPool,
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
tp_group: torch.distributed.ProcessGroup,
tree_cache: BasePrefixCache,
server_args: ServerArgs,
) -> None:
self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
self.page_size = server_args.page_size
self.server_args = server_args
self.request_counter = 0
self.tree_cache = tree_cache
kv_cache = self.token_to_kv_pool_allocator.get_kvcache()
if isinstance(kv_cache, MHATokenToKVPool):
self.decode_host_mem_pool = MHATokenToKVPoolHost(
kv_cache,
server_args.hicache_ratio,
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
)
elif isinstance(kv_cache, MLATokenToKVPool):
self.decode_host_mem_pool = MLATokenToKVPoolHost(
kv_cache,
server_args.hicache_ratio,
server_args.hicache_size,
self.page_size,
server_args.hicache_mem_layout,
)
else:
raise ValueError("Unsupported KV cache type for decode offload")
self.tp_group = tp_group
self.tp_world_size = torch.distributed.get_world_size(group=self.tp_group)
self.cache_controller = HiCacheController(
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
mem_pool_host=self.decode_host_mem_pool,
page_size=self.page_size,
tp_group=tp_group,
io_backend=server_args.hicache_io_backend,
load_cache_event=threading.Event(),
storage_backend=server_args.hicache_storage_backend,
model_name=server_args.served_model_name,
storage_backend_extra_config=server_args.hicache_storage_backend_extra_config,
)
self.ongoing_offload = {}
self.ongoing_backup = {}
logger.info("Enable offload kv cache for decode side")
def offload_kv_cache(self, req) -> bool:
"""Offload incremental KV cache for decode side."""
if self.cache_controller is None or self.decode_host_mem_pool is None:
return False
if req.req_pool_idx == -1 or len(req.output_ids) == 0:
return False
token_indices = self.req_to_token_pool.req_to_token[req.req_pool_idx]
if token_indices.dim() == 0 or token_indices.numel() == 0:
return False
# Prefill side offloads page-aligned origin_input_ids, decode side offloads the incremental part
all_tokens = req.origin_input_ids + req.output_ids[:-1]
prefill_offloaded_len = (
len(req.origin_input_ids) // self.page_size * self.page_size
)
incremental_len = len(all_tokens) - prefill_offloaded_len
incremental_aligned_len = incremental_len // self.page_size * self.page_size
if incremental_aligned_len == 0:
return False
# Extract incremental tokens and indices
start, end = (
prefill_offloaded_len,
prefill_offloaded_len + incremental_aligned_len,
)
incremental_tokens = all_tokens[start:end]
incremental_indices = token_indices[start:end]
# Early free prefill-offloaded GPU memory
if prefill_offloaded_len > 0:
self.token_to_kv_pool_allocator.free(token_indices[:prefill_offloaded_len])
# Asynchronously offload incremental KV cache from device to host
self.request_counter += 1
ack_id = self.request_counter
host_indices = self.cache_controller.write(
device_indices=incremental_indices.long(),
node_id=ack_id,
)
if host_indices is None:
logger.error(f"Not enough host memory for request {req.rid}")
return False
self.ongoing_offload[ack_id] = (
req,
host_indices,
incremental_tokens,
time.time(),
prefill_offloaded_len,
)
return True
def check_offload_progress(self):
"""Check the progress of offload from device to host and backup from host to storage."""
cc = self.cache_controller
qsizes = torch.tensor(
[
len(cc.ack_write_queue),
cc.ack_backup_queue.qsize(),
],
dtype=torch.int,
)
if self.tp_world_size > 1:
torch.distributed.all_reduce(
qsizes, op=torch.distributed.ReduceOp.MIN, group=self.tp_group
)
n_write, n_backup = map(int, qsizes.tolist())
self._check_offload_progress(n_write)
self._check_backup_progress(n_backup)
def _check_offload_progress(self, finish_count):
"""Check the progress of offload from device to host."""
while finish_count > 0:
_, finish_event, ack_list = self.cache_controller.ack_write_queue.pop(0)
finish_event.synchronize()
for ack_id in ack_list:
(
req,
host_indices,
incremental_tokens,
start_time,
prefill_offloaded_len,
) = self.ongoing_offload.pop(ack_id)
self._release_finished_req(req, prefill_offloaded_len)
self._trigger_backup(
req,
host_indices,
incremental_tokens,
start_time,
prefill_offloaded_len,
)
finish_count -= 1
def _release_finished_req(self, req: Req, prefill_offloaded_len: int):
kv_committed_len = req.pop_committed_kv_cache()
kv_indices = self.req_to_token_pool.req_to_token[
req.req_pool_idx, prefill_offloaded_len:kv_committed_len
]
# Free the incremental part of the request
self.token_to_kv_pool_allocator.free(kv_indices)
self.req_to_token_pool.free(req.req_pool_idx)
self.tree_cache.protected_size_ -= len(req.prefix_indices)
def _check_backup_progress(self, finish_count):
"""Check the progress of backup from host to storage."""
for _ in range(finish_count):
storage_operation = self.cache_controller.ack_backup_queue.get()
ack_id = storage_operation.id
req_id, host_indices, start_time = self.ongoing_backup.pop(ack_id)
# Release host memory
self.decode_host_mem_pool.free(host_indices)
logger.info(
f"Finished backup request {req_id}, free host memory, len:{len(host_indices)}, cost time:{time.time() - start_time:.2f} seconds."
)
def _trigger_backup(
self, req, host_indices, incremental_tokens, start_time, prefill_offloaded_len
):
"""Trigger async backup from host to storage."""
prefill_hashes = self._compute_prefix_hash(
req.origin_input_ids[:prefill_offloaded_len]
)
last_prefill_hash = prefill_hashes[-1] if prefill_offloaded_len > 0 else ""
page_hashes = self._compute_prefix_hash(incremental_tokens, last_prefill_hash)
ack_id = self.cache_controller.write_storage(
host_indices,
incremental_tokens,
hash_value=page_hashes,
)
self.ongoing_backup[ack_id] = (req.rid, host_indices, start_time)
def _compute_prefix_hash(self, tokens, prior_hash=""):
page_hashes = []
last_hash = prior_hash
for offset in range(0, len(tokens), self.page_size):
page_tokens = tokens[offset : offset + self.page_size]
last_hash = self.cache_controller.get_hash_str(page_tokens, last_hash)
page_hashes.append(last_hash)
return page_hashes