state_capturer: pin the exact host-cache size via mmap + cudaHostRegister (#37285)

This commit is contained in:
Yueming Yuan
2026-09-03 16:20:59 -07:00
committed by GitHub
parent 8b0501399e
commit 66d60433c1
4 changed files with 45 additions and 3 deletions
+4
View File
@@ -309,6 +309,8 @@ from sglang.srt.speculative.eagle_utils import (
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.uno_validation import validate_uno_request
from sglang.srt.state_capturer.indexer_topk import destroy_global_indexer_capturer
from sglang.srt.state_capturer.routed_experts import destroy_global_experts_capturer
from sglang.srt.utils import (
DynamicGradMode,
configure_gc_logger,
@@ -1798,6 +1800,8 @@ class Scheduler(
self.tree_cache.release_host_resources()
if self.decode_offload_manager is not None:
self.decode_offload_manager.release_host_resources()
destroy_global_experts_capturer()
destroy_global_indexer_capturer()
rank_consensus_checker.shutdown()
+29 -3
View File
@@ -5,10 +5,19 @@ from typing import Optional
import torch
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.pool_host.common import (
ALLOC_MEMORY_FUNCS,
HostTensorAllocator,
_cuda_host_unregister,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.utils import is_cuda, is_hip
logger = logging.getLogger(__name__)
_is_cuda = is_cuda()
_is_hip = is_hip()
_GB = 1024 * 1024 * 1024
_MB = 1024 * 1024
@@ -52,19 +61,31 @@ class BaseDeviceCache:
class BaseHostCache:
def __init__(self, num_tokens: int, num_layers: int, topk_size: int, name: str):
self.buffer = torch.zeros(
def __init__(
self, num_tokens: int, num_layers: int, topk_size: int, name: str, device: str
):
alloc = ALLOC_MEMORY_FUNCS[device]
self.buffer = alloc(
(num_tokens, num_layers, topk_size),
dtype=torch.int32,
device="cpu",
pin_memory=True,
allocator=HostTensorAllocator(),
)
self.buffer.zero_()
self.num_tokens = num_tokens
self.num_layers = num_layers
self.topk_size = topk_size
self.name = name
self._log_allocation()
def destroy(self):
if self.buffer is None:
return
if _is_cuda or _is_hip:
_cuda_host_unregister(self.buffer)
self.buffer = None
def get_buffer_size_bytes(self):
return get_tensor_size_bytes(self.buffer)
@@ -115,7 +136,9 @@ class BaseTopkCapturer:
self.num_layers = num_layers
self.topk_size = topk_size
self.host_cache = BaseHostCache(num_tokens, num_layers, topk_size, name=name)
self.host_cache = BaseHostCache(
num_tokens, num_layers, topk_size, name=name, device=device
)
self.device_cache = BaseDeviceCache(
max_batch_size,
num_layers,
@@ -127,6 +150,9 @@ class BaseTopkCapturer:
def capture(self, layer_id: int, topk_indices: torch.Tensor):
self.device_cache.capture(layer_id, topk_indices)
def destroy(self):
self.host_cache.destroy()
def _get_local_slice(
self,
forward_batch: ForwardBatch,
@@ -56,6 +56,12 @@ def set_global_indexer_capturer(capturer: Optional[IndexerTopkCapturer]):
get_resources().indexer_capturer = capturer
def destroy_global_indexer_capturer():
if (capturer := get_resources().indexer_capturer) is not None:
capturer.destroy()
get_resources().indexer_capturer = None
def maybe_capture_indexer_topk(
layer_id: int, topk_indices: Optional[torch.Tensor]
) -> Optional[torch.Tensor]:
@@ -144,6 +144,12 @@ def set_global_experts_capturer(capturer: Optional[RoutedExpertsCapturer]):
get_resources().experts_capturer = capturer
def destroy_global_experts_capturer():
if (capturer := get_resources().experts_capturer) is not None:
capturer.destroy()
get_resources().experts_capturer = None
def extract_routed_experts_from_meta_info(data):
# To solve the performance issue, we return the experts_ids in base64
# We left this function for user to change it back to normal int32