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.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.speculative.uno_validation import validate_uno_request
|
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 (
|
from sglang.srt.utils import (
|
||||||
DynamicGradMode,
|
DynamicGradMode,
|
||||||
configure_gc_logger,
|
configure_gc_logger,
|
||||||
@@ -1798,6 +1800,8 @@ class Scheduler(
|
|||||||
self.tree_cache.release_host_resources()
|
self.tree_cache.release_host_resources()
|
||||||
if self.decode_offload_manager is not None:
|
if self.decode_offload_manager is not None:
|
||||||
self.decode_offload_manager.release_host_resources()
|
self.decode_offload_manager.release_host_resources()
|
||||||
|
destroy_global_experts_capturer()
|
||||||
|
destroy_global_indexer_capturer()
|
||||||
|
|
||||||
rank_consensus_checker.shutdown()
|
rank_consensus_checker.shutdown()
|
||||||
|
|
||||||
|
|||||||
@@ -5,10 +5,19 @@ from typing import Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
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.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.utils import is_cuda, is_hip
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_is_cuda = is_cuda()
|
||||||
|
_is_hip = is_hip()
|
||||||
|
|
||||||
_GB = 1024 * 1024 * 1024
|
_GB = 1024 * 1024 * 1024
|
||||||
_MB = 1024 * 1024
|
_MB = 1024 * 1024
|
||||||
|
|
||||||
@@ -52,19 +61,31 @@ class BaseDeviceCache:
|
|||||||
|
|
||||||
|
|
||||||
class BaseHostCache:
|
class BaseHostCache:
|
||||||
def __init__(self, num_tokens: int, num_layers: int, topk_size: int, name: str):
|
def __init__(
|
||||||
self.buffer = torch.zeros(
|
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),
|
(num_tokens, num_layers, topk_size),
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device="cpu",
|
device="cpu",
|
||||||
pin_memory=True,
|
pin_memory=True,
|
||||||
|
allocator=HostTensorAllocator(),
|
||||||
)
|
)
|
||||||
|
self.buffer.zero_()
|
||||||
self.num_tokens = num_tokens
|
self.num_tokens = num_tokens
|
||||||
self.num_layers = num_layers
|
self.num_layers = num_layers
|
||||||
self.topk_size = topk_size
|
self.topk_size = topk_size
|
||||||
self.name = name
|
self.name = name
|
||||||
self._log_allocation()
|
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):
|
def get_buffer_size_bytes(self):
|
||||||
return get_tensor_size_bytes(self.buffer)
|
return get_tensor_size_bytes(self.buffer)
|
||||||
|
|
||||||
@@ -115,7 +136,9 @@ class BaseTopkCapturer:
|
|||||||
self.num_layers = num_layers
|
self.num_layers = num_layers
|
||||||
self.topk_size = topk_size
|
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(
|
self.device_cache = BaseDeviceCache(
|
||||||
max_batch_size,
|
max_batch_size,
|
||||||
num_layers,
|
num_layers,
|
||||||
@@ -127,6 +150,9 @@ class BaseTopkCapturer:
|
|||||||
def capture(self, layer_id: int, topk_indices: torch.Tensor):
|
def capture(self, layer_id: int, topk_indices: torch.Tensor):
|
||||||
self.device_cache.capture(layer_id, topk_indices)
|
self.device_cache.capture(layer_id, topk_indices)
|
||||||
|
|
||||||
|
def destroy(self):
|
||||||
|
self.host_cache.destroy()
|
||||||
|
|
||||||
def _get_local_slice(
|
def _get_local_slice(
|
||||||
self,
|
self,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
|
|||||||
@@ -56,6 +56,12 @@ def set_global_indexer_capturer(capturer: Optional[IndexerTopkCapturer]):
|
|||||||
get_resources().indexer_capturer = capturer
|
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(
|
def maybe_capture_indexer_topk(
|
||||||
layer_id: int, topk_indices: Optional[torch.Tensor]
|
layer_id: int, topk_indices: Optional[torch.Tensor]
|
||||||
) -> Optional[torch.Tensor]:
|
) -> Optional[torch.Tensor]:
|
||||||
|
|||||||
@@ -144,6 +144,12 @@ def set_global_experts_capturer(capturer: Optional[RoutedExpertsCapturer]):
|
|||||||
get_resources().experts_capturer = capturer
|
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):
|
def extract_routed_experts_from_meta_info(data):
|
||||||
# To solve the performance issue, we return the experts_ids in base64
|
# 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
|
# We left this function for user to change it back to normal int32
|
||||||
|
|||||||
Reference in New Issue
Block a user