[VLM] Boost Memory Pool based CUDA IPC (#14123)
Co-authored-by: luoyuan.luo <luoyuan.luo@antgroup.com>
This commit is contained in:
@@ -23,6 +23,7 @@ from sglang.srt.utils import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
||||||
MM_FEATURE_CACHE_SIZE,
|
MM_FEATURE_CACHE_SIZE,
|
||||||
|
MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
|
||||||
CudaIpcTensorTransportProxy,
|
CudaIpcTensorTransportProxy,
|
||||||
MmItemMemoryPool,
|
MmItemMemoryPool,
|
||||||
)
|
)
|
||||||
@@ -225,7 +226,10 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
]
|
]
|
||||||
|
|
||||||
if SGL_USE_CUDA_IPC:
|
if SGL_USE_CUDA_IPC:
|
||||||
self.cudaipc_mmfeature_pool = MmItemMemoryPool(MM_FEATURE_CACHE_SIZE)
|
self.cudaipc_mmfeature_pool = MmItemMemoryPool(
|
||||||
|
MM_FEATURE_CACHE_SIZE,
|
||||||
|
MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL,
|
||||||
|
)
|
||||||
|
|
||||||
def process_mm_data(
|
def process_mm_data(
|
||||||
self, input_text, images=None, videos=None, audios=None, **kwargs
|
self, input_text, images=None, videos=None, audios=None, **kwargs
|
||||||
|
|||||||
@@ -284,6 +284,17 @@ def get_int_env_var(name: str, default: int = 0) -> int:
|
|||||||
return default
|
return default
|
||||||
|
|
||||||
|
|
||||||
|
def get_float_env_var(name: str, default: float = 0.0) -> float:
|
||||||
|
# FIXME: move your environment variable to sglang.srt.environ
|
||||||
|
value = os.getenv(name)
|
||||||
|
if value is None or not value.strip():
|
||||||
|
return default
|
||||||
|
try:
|
||||||
|
return float(value)
|
||||||
|
except ValueError:
|
||||||
|
return default
|
||||||
|
|
||||||
|
|
||||||
def support_triton(backend: str) -> bool:
|
def support_triton(backend: str) -> bool:
|
||||||
return backend not in ["torch_native", "intel_amx"]
|
return backend not in ["torch_native", "intel_amx"]
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
import fcntl
|
import fcntl
|
||||||
import logging
|
import logging
|
||||||
|
import threading
|
||||||
|
import time
|
||||||
from multiprocessing import shared_memory
|
from multiprocessing import shared_memory
|
||||||
from typing import Tuple
|
from typing import Tuple
|
||||||
|
|
||||||
@@ -7,16 +9,22 @@ import numpy as np
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import get_int_env_var
|
from sglang.srt.utils import get_float_env_var, get_int_env_var
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
MM_FEATURE_CACHE_SIZE = (
|
MM_FEATURE_CACHE_SIZE = (
|
||||||
2 * 1024 * 1024 * 1024
|
4 * 1024 * 1024 * 1024
|
||||||
if not get_int_env_var("SGLANG_MM_FEATURE_CACHE_MB")
|
if not get_int_env_var("SGLANG_MM_FEATURE_CACHE_MB")
|
||||||
else get_int_env_var("SGLANG_MM_FEATURE_CACHE_MB") * 1024 * 1024
|
else get_int_env_var("SGLANG_MM_FEATURE_CACHE_MB") * 1024 * 1024
|
||||||
)
|
)
|
||||||
|
|
||||||
|
MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL = (
|
||||||
|
0.05
|
||||||
|
if not get_float_env_var("SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC")
|
||||||
|
else get_float_env_var("SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC")
|
||||||
|
)
|
||||||
|
|
||||||
SHM_LOCK_FILE = "/tmp/shm_wr_lock.lock"
|
SHM_LOCK_FILE = "/tmp/shm_wr_lock.lock"
|
||||||
|
|
||||||
|
|
||||||
@@ -57,20 +65,24 @@ class MmItemMemoryChunk:
|
|||||||
def try_to_recycle(self) -> bool:
|
def try_to_recycle(self) -> bool:
|
||||||
try:
|
try:
|
||||||
tp_num = get_global_server_args().tp_size
|
tp_num = get_global_server_args().tp_size
|
||||||
except:
|
except Exception:
|
||||||
logger.info(
|
logger.info(
|
||||||
"get_global_server_args has not been inited , skip this turn 's recycle"
|
"get_global_server_args has not been inited , skip this turn 's recycle"
|
||||||
)
|
)
|
||||||
tp_num = -1
|
return False
|
||||||
if self.sync_flag.buffer_wrapper.item() == float(tp_num):
|
|
||||||
self.sync_flag.buffer_wrapper *= 0
|
val = float(self.sync_flag.buffer_wrapper.item())
|
||||||
|
logger.debug(f"[try_to_recycle] area={self.area}, flag={val}, tp_size={tp_num}")
|
||||||
|
|
||||||
|
if val == float(tp_num):
|
||||||
|
self.sync_flag.buffer_wrapper *= 0.0
|
||||||
return True
|
return True
|
||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
class MmItemMemoryPool:
|
class MmItemMemoryPool:
|
||||||
def __init__(self, memory_size):
|
def __init__(self, memory_size, recycle_interval):
|
||||||
self.memory_pool = torch.empty(
|
self.memory_pool = torch.empty(
|
||||||
memory_size, dtype=torch.int8, device="cuda"
|
memory_size, dtype=torch.int8, device="cuda"
|
||||||
).contiguous()
|
).contiguous()
|
||||||
@@ -81,6 +93,38 @@ class MmItemMemoryPool:
|
|||||||
self.available_chunks = [init_chunk]
|
self.available_chunks = [init_chunk]
|
||||||
self.occupied_chunks = []
|
self.occupied_chunks = []
|
||||||
|
|
||||||
|
self._lock = threading.Lock()
|
||||||
|
|
||||||
|
self._recycle_interval = recycle_interval
|
||||||
|
self._stop_recycler = False
|
||||||
|
self._recycle_thread = threading.Thread(
|
||||||
|
target=self._recycle_loop, name="MmItemMemoryPoolRecycler", daemon=True
|
||||||
|
)
|
||||||
|
self._recycle_thread.start()
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
f"[MmItemMemoryPool] init: memory_size={memory_size}, "
|
||||||
|
f"recycle_interval={recycle_interval}s"
|
||||||
|
)
|
||||||
|
|
||||||
|
def shutdown(self):
|
||||||
|
self._stop_recycler = True
|
||||||
|
if self._recycle_thread.is_alive():
|
||||||
|
self._recycle_thread.join(timeout=1.0)
|
||||||
|
|
||||||
|
def _recycle_loop(self):
|
||||||
|
while not self._stop_recycler:
|
||||||
|
try:
|
||||||
|
with self._lock:
|
||||||
|
self.recycle_chunks()
|
||||||
|
self.merge_chunks()
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
f"[MmItemMemoryPool] recycle loop error: {e}", exc_info=True
|
||||||
|
)
|
||||||
|
|
||||||
|
time.sleep(self._recycle_interval)
|
||||||
|
|
||||||
def clear_sync_flag_list(self):
|
def clear_sync_flag_list(self):
|
||||||
# call each chunk's __del__
|
# call each chunk's __del__
|
||||||
self.sync_flag_list.clear()
|
self.sync_flag_list.clear()
|
||||||
@@ -137,9 +181,7 @@ class MmItemMemoryPool:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
def return_a_slice_tensor_with_flag(self, src_tensor: torch.Tensor):
|
def return_a_slice_tensor_with_flag(self, src_tensor: torch.Tensor):
|
||||||
self.recycle_chunks()
|
with self._lock:
|
||||||
self.merge_chunks()
|
|
||||||
|
|
||||||
available_chunk = self.get_available_chunk(src_tensor)
|
available_chunk = self.get_available_chunk(src_tensor)
|
||||||
if available_chunk is not None:
|
if available_chunk is not None:
|
||||||
return (
|
return (
|
||||||
|
|||||||
Reference in New Issue
Block a user