state_capturer: pin the exact host-cache size via mmap + cudaHostRegister (#37285)
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user