diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 4c0189a21..e6dfaf99f 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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() diff --git a/python/sglang/srt/state_capturer/base.py b/python/sglang/srt/state_capturer/base.py index 0fa8bcdad..6f6b87578 100644 --- a/python/sglang/srt/state_capturer/base.py +++ b/python/sglang/srt/state_capturer/base.py @@ -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, diff --git a/python/sglang/srt/state_capturer/indexer_topk.py b/python/sglang/srt/state_capturer/indexer_topk.py index b20b5e450..e31951c1f 100644 --- a/python/sglang/srt/state_capturer/indexer_topk.py +++ b/python/sglang/srt/state_capturer/indexer_topk.py @@ -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]: diff --git a/python/sglang/srt/state_capturer/routed_experts.py b/python/sglang/srt/state_capturer/routed_experts.py index 57f512d8b..fab5c4d92 100644 --- a/python/sglang/srt/state_capturer/routed_experts.py +++ b/python/sglang/srt/state_capturer/routed_experts.py @@ -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