[VLM] perf: optimize CUDA IPC for multimodal transfer by caching IPC pool handles (#21418)
This commit is contained in:
@@ -451,6 +451,7 @@ class Envs:
|
|||||||
|
|
||||||
# VLM Item CUDA IPC Transport
|
# VLM Item CUDA IPC Transport
|
||||||
SGLANG_USE_CUDA_IPC_TRANSPORT = EnvBool(False)
|
SGLANG_USE_CUDA_IPC_TRANSPORT = EnvBool(False)
|
||||||
|
SGLANG_USE_IPC_POOL_HANDLE_CACHE = EnvBool(False)
|
||||||
SGLANG_MM_FEATURE_CACHE_MB = EnvInt(4 * 1024)
|
SGLANG_MM_FEATURE_CACHE_MB = EnvInt(4 * 1024)
|
||||||
SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC = EnvFloat(0.05)
|
SGLANG_MM_ITEM_MEM_POOL_RECYCLE_INTERVAL_SEC = EnvFloat(0.05)
|
||||||
|
|
||||||
|
|||||||
@@ -40,6 +40,7 @@ _is_npu = is_npu()
|
|||||||
_is_xpu = is_xpu()
|
_is_xpu = is_xpu()
|
||||||
|
|
||||||
SGL_USE_CUDA_IPC = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()
|
SGL_USE_CUDA_IPC = envs.SGLANG_USE_CUDA_IPC_TRANSPORT.get()
|
||||||
|
_IPC_POOL_HANDLE_CACHE = envs.SGLANG_USE_IPC_POOL_HANDLE_CACHE.get()
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
@@ -1150,7 +1151,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
# post-process
|
# post-process
|
||||||
for item in all_collected_items:
|
for item in all_collected_items:
|
||||||
if isinstance(item.feature, torch.Tensor) and item.feature.is_cuda:
|
if isinstance(item.feature, torch.Tensor) and item.feature.is_cuda:
|
||||||
sync_flag, available_slice = (
|
sync_flag, available_slice, byte_offset = (
|
||||||
self.cudaipc_mmfeature_pool.return_a_slice_tensor_with_flag(
|
self.cudaipc_mmfeature_pool.return_a_slice_tensor_with_flag(
|
||||||
item.feature
|
item.feature
|
||||||
)
|
)
|
||||||
@@ -1163,6 +1164,13 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
data=available_slice,
|
data=available_slice,
|
||||||
info_data=item.feature,
|
info_data=item.feature,
|
||||||
sync_buffer_meta=sync_flag,
|
sync_buffer_meta=sync_flag,
|
||||||
|
pool_ipc_handle=(
|
||||||
|
self.cudaipc_mmfeature_pool._pool_ipc_handle
|
||||||
|
if _IPC_POOL_HANDLE_CACHE
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
pool_byte_offset=byte_offset,
|
||||||
|
pool_device_index=self.cudaipc_mmfeature_pool._pool_device_index,
|
||||||
)
|
)
|
||||||
elif not self.server_args.keep_mm_feature_on_device:
|
elif not self.server_args.keep_mm_feature_on_device:
|
||||||
item.feature = item.feature.cpu()
|
item.feature = item.feature.cpu()
|
||||||
@@ -1171,7 +1179,7 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
and item.precomputed_embeddings.is_cuda
|
and item.precomputed_embeddings.is_cuda
|
||||||
):
|
):
|
||||||
|
|
||||||
sync_flag, available_slice = (
|
sync_flag, available_slice, byte_offset = (
|
||||||
self.cudaipc_mmfeature_pool.return_a_slice_tensor_with_flag(
|
self.cudaipc_mmfeature_pool.return_a_slice_tensor_with_flag(
|
||||||
item.precomputed_embeddings
|
item.precomputed_embeddings
|
||||||
)
|
)
|
||||||
@@ -1185,6 +1193,13 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
data=available_slice,
|
data=available_slice,
|
||||||
info_data=item.precomputed_embeddings,
|
info_data=item.precomputed_embeddings,
|
||||||
sync_buffer_meta=sync_flag,
|
sync_buffer_meta=sync_flag,
|
||||||
|
pool_ipc_handle=(
|
||||||
|
self.cudaipc_mmfeature_pool._pool_ipc_handle
|
||||||
|
if _IPC_POOL_HANDLE_CACHE
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
pool_byte_offset=byte_offset,
|
||||||
|
pool_device_index=self.cudaipc_mmfeature_pool._pool_device_index,
|
||||||
)
|
)
|
||||||
elif not self.server_args.keep_mm_feature_on_device:
|
elif not self.server_args.keep_mm_feature_on_device:
|
||||||
item.precomputed_embeddings = item.precomputed_embeddings.cpu()
|
item.precomputed_embeddings = item.precomputed_embeddings.cpu()
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ import logging
|
|||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from multiprocessing import shared_memory
|
from multiprocessing import shared_memory
|
||||||
from typing import Tuple
|
from typing import Any, Tuple
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
@@ -22,6 +22,49 @@ MM_ITEM_MEMORY_POOL_RECYCLE_INTERVAL = (
|
|||||||
SHM_LOCK_FILE = "/tmp/shm_wr_lock.lock"
|
SHM_LOCK_FILE = "/tmp/shm_wr_lock.lock"
|
||||||
|
|
||||||
|
|
||||||
|
# Cache for pool-level IPC handles on the consumer side.
|
||||||
|
# Key: the pool CUDA IPC handle tuple. Value: opened UntypedStorage.
|
||||||
|
_pool_storage_cache: dict = {}
|
||||||
|
_pool_cache_lock = threading.Lock()
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_pool_cache_key(pool_handle, pool_device_index: int) -> tuple[Any, ...]:
|
||||||
|
normalized_handle = (
|
||||||
|
pool_handle if isinstance(pool_handle, tuple) else tuple(pool_handle)
|
||||||
|
)
|
||||||
|
return (pool_device_index, normalized_handle)
|
||||||
|
|
||||||
|
|
||||||
|
def _open_pooled_storage_uncached(pool_handle):
|
||||||
|
return torch.UntypedStorage._new_shared_cuda(*pool_handle)
|
||||||
|
|
||||||
|
|
||||||
|
def _pool_handle_cache_get_or_open(cache_key, pool_handle):
|
||||||
|
storage = _pool_storage_cache.get(cache_key)
|
||||||
|
if storage is None:
|
||||||
|
with _pool_cache_lock:
|
||||||
|
storage = _pool_storage_cache.get(cache_key)
|
||||||
|
if storage is None:
|
||||||
|
storage = _open_pooled_storage_uncached(pool_handle)
|
||||||
|
_pool_storage_cache[cache_key] = storage
|
||||||
|
return storage
|
||||||
|
|
||||||
|
|
||||||
|
def _pool_handle_cache_set(cache_key, storage):
|
||||||
|
with _pool_cache_lock:
|
||||||
|
_pool_storage_cache[cache_key] = storage
|
||||||
|
|
||||||
|
|
||||||
|
def _pool_handle_cache_invalidate(cache_key):
|
||||||
|
with _pool_cache_lock:
|
||||||
|
_pool_storage_cache.pop(cache_key, None)
|
||||||
|
|
||||||
|
|
||||||
|
def _pool_handle_cache_clear():
|
||||||
|
with _pool_cache_lock:
|
||||||
|
_pool_storage_cache.clear()
|
||||||
|
|
||||||
|
|
||||||
class ShmSyncBuffer:
|
class ShmSyncBuffer:
|
||||||
def __init__(self, byte_size: int = 4):
|
def __init__(self, byte_size: int = 4):
|
||||||
self.buffer = shared_memory.SharedMemory(create=True, size=byte_size)
|
self.buffer = shared_memory.SharedMemory(create=True, size=byte_size)
|
||||||
@@ -80,6 +123,9 @@ class MmItemMemoryPool:
|
|||||||
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()
|
||||||
|
storage = self.memory_pool.untyped_storage()
|
||||||
|
self._pool_ipc_handle = storage._share_cuda_()
|
||||||
|
self._pool_device_index = self.memory_pool.device.index
|
||||||
|
|
||||||
self.sync_flag_list = []
|
self.sync_flag_list = []
|
||||||
|
|
||||||
@@ -181,8 +227,9 @@ class MmItemMemoryPool:
|
|||||||
return (
|
return (
|
||||||
available_chunk.sync_flag.meta_data,
|
available_chunk.sync_flag.meta_data,
|
||||||
self.memory_pool[available_chunk.start : available_chunk.end],
|
self.memory_pool[available_chunk.start : available_chunk.end],
|
||||||
|
available_chunk.start,
|
||||||
)
|
)
|
||||||
return None, None
|
return None, None, None
|
||||||
|
|
||||||
def recycle_chunks(self):
|
def recycle_chunks(self):
|
||||||
|
|
||||||
@@ -229,6 +276,9 @@ class CudaIpcTensorTransportProxy:
|
|||||||
data: torch.Tensor,
|
data: torch.Tensor,
|
||||||
info_data: torch.Tensor,
|
info_data: torch.Tensor,
|
||||||
sync_buffer_meta,
|
sync_buffer_meta,
|
||||||
|
pool_ipc_handle=None,
|
||||||
|
pool_byte_offset: int = 0,
|
||||||
|
pool_device_index: int = 0,
|
||||||
):
|
):
|
||||||
|
|
||||||
if (not isinstance(data, torch.Tensor)) or (
|
if (not isinstance(data, torch.Tensor)) or (
|
||||||
@@ -238,6 +288,23 @@ class CudaIpcTensorTransportProxy:
|
|||||||
f"Input 'data' must be a torch.Tensor, but got {type(data)}"
|
f"Input 'data' must be a torch.Tensor, but got {type(data)}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if pool_ipc_handle is not None:
|
||||||
|
self.proxy_state = {
|
||||||
|
"ipc_extra": {
|
||||||
|
"pool_handle": pool_ipc_handle,
|
||||||
|
"pool_byte_offset": pool_byte_offset,
|
||||||
|
"pool_device_index": pool_device_index,
|
||||||
|
"shape": data.shape,
|
||||||
|
"dtype": data.dtype,
|
||||||
|
"stride": data.stride(),
|
||||||
|
"storage_offset": 0,
|
||||||
|
"nbytes": data.numel() * data.element_size(),
|
||||||
|
"recons_shape": info_data.shape,
|
||||||
|
"recons_dtype": info_data.dtype,
|
||||||
|
},
|
||||||
|
"tensor_data": None,
|
||||||
|
}
|
||||||
|
else:
|
||||||
self.proxy_state = self.get_proxy_state(data, info_data)
|
self.proxy_state = self.get_proxy_state(data, info_data)
|
||||||
self.reconstruct_tensor = None
|
self.reconstruct_tensor = None
|
||||||
self.sync_data_meta = sync_buffer_meta
|
self.sync_data_meta = sync_buffer_meta
|
||||||
@@ -283,44 +350,45 @@ class CudaIpcTensorTransportProxy:
|
|||||||
|
|
||||||
return state
|
return state
|
||||||
|
|
||||||
def reconstruct_on_target_device(self, rebuild_device_idx):
|
def _reconstruct_from_ipc_extra(self, ipc_extra, *, use_cache: bool):
|
||||||
rebuild_device = torch.device(f"cuda:{rebuild_device_idx}")
|
shape = ipc_extra["shape"]
|
||||||
if (
|
dtype = ipc_extra["dtype"]
|
||||||
isinstance(self.reconstruct_tensor, torch.Tensor)
|
stride = ipc_extra["stride"]
|
||||||
and self.reconstruct_tensor.device == rebuild_device
|
target_device = torch.device(f"cuda:{ipc_extra['pool_device_index']}")
|
||||||
):
|
cache_key = _normalize_pool_cache_key(
|
||||||
return self.reconstruct_tensor
|
ipc_extra["pool_handle"], ipc_extra["pool_device_index"]
|
||||||
|
|
||||||
if self.proxy_state["ipc_extra"]:
|
|
||||||
ipc_extra = self.proxy_state["ipc_extra"]
|
|
||||||
(
|
|
||||||
handle,
|
|
||||||
shape,
|
|
||||||
dtype,
|
|
||||||
stride,
|
|
||||||
source_device_index,
|
|
||||||
s_offset,
|
|
||||||
recons_shape,
|
|
||||||
recons_dtype,
|
|
||||||
) = (
|
|
||||||
ipc_extra["handle"],
|
|
||||||
ipc_extra["shape"],
|
|
||||||
ipc_extra["dtype"],
|
|
||||||
ipc_extra["stride"],
|
|
||||||
ipc_extra["device_index"],
|
|
||||||
ipc_extra["storage_offset"],
|
|
||||||
ipc_extra["recons_shape"],
|
|
||||||
ipc_extra["recons_dtype"],
|
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
|
||||||
target_device = torch.device(f"cuda:{source_device_index}")
|
|
||||||
with torch.cuda.device(target_device):
|
with torch.cuda.device(target_device):
|
||||||
storage = torch.UntypedStorage._new_shared_cuda(*handle)
|
if use_cache:
|
||||||
slice_tensor = torch.empty(
|
storage = _pool_handle_cache_get_or_open(
|
||||||
0, dtype=dtype, device=target_device
|
cache_key, ipc_extra["pool_handle"]
|
||||||
).set_(storage, storage_offset=s_offset, size=shape, stride=stride)
|
)
|
||||||
|
storage_to_cache = None
|
||||||
|
else:
|
||||||
|
storage = _open_pooled_storage_uncached(ipc_extra["pool_handle"])
|
||||||
|
storage_to_cache = storage
|
||||||
|
slice_storage = storage[
|
||||||
|
ipc_extra["pool_byte_offset"] : ipc_extra["pool_byte_offset"]
|
||||||
|
+ ipc_extra["nbytes"]
|
||||||
|
]
|
||||||
|
slice_tensor = torch.empty(0, dtype=dtype, device=target_device).set_(
|
||||||
|
slice_storage,
|
||||||
|
storage_offset=ipc_extra["storage_offset"],
|
||||||
|
size=shape,
|
||||||
|
stride=stride,
|
||||||
|
)
|
||||||
|
|
||||||
|
return slice_tensor, target_device, cache_key, storage_to_cache
|
||||||
|
|
||||||
|
def _copy_slice_tensor_to_target(
|
||||||
|
self,
|
||||||
|
slice_tensor: torch.Tensor,
|
||||||
|
rebuild_device: torch.device,
|
||||||
|
recons_shape,
|
||||||
|
recons_dtype,
|
||||||
|
):
|
||||||
|
with torch.cuda.device(rebuild_device):
|
||||||
reconstructed_tensor = torch.empty(
|
reconstructed_tensor = torch.empty(
|
||||||
recons_shape, dtype=recons_dtype, device=rebuild_device
|
recons_shape, dtype=recons_dtype, device=rebuild_device
|
||||||
).contiguous()
|
).contiguous()
|
||||||
@@ -336,9 +404,70 @@ class CudaIpcTensorTransportProxy:
|
|||||||
|
|
||||||
self.close_shm()
|
self.close_shm()
|
||||||
|
|
||||||
|
return reconstructed_tensor
|
||||||
|
|
||||||
|
def reconstruct_on_target_device(self, rebuild_device_idx):
|
||||||
|
rebuild_device = torch.device(f"cuda:{rebuild_device_idx}")
|
||||||
|
if (
|
||||||
|
isinstance(self.reconstruct_tensor, torch.Tensor)
|
||||||
|
and self.reconstruct_tensor.device == rebuild_device
|
||||||
|
):
|
||||||
|
return self.reconstruct_tensor
|
||||||
|
|
||||||
|
if self.proxy_state["ipc_extra"]:
|
||||||
|
ipc_extra = self.proxy_state["ipc_extra"]
|
||||||
|
recons_shape = ipc_extra["recons_shape"]
|
||||||
|
recons_dtype = ipc_extra["recons_dtype"]
|
||||||
|
|
||||||
|
if "pool_handle" in ipc_extra:
|
||||||
|
try:
|
||||||
|
(
|
||||||
|
slice_tensor,
|
||||||
|
_target_device,
|
||||||
|
cache_key,
|
||||||
|
storage_to_cache,
|
||||||
|
) = self._reconstruct_from_ipc_extra(ipc_extra, use_cache=True)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.info(f"Error: Failed to deserialize from CUDA IPC handle ({e}).")
|
cache_key = _normalize_pool_cache_key(
|
||||||
raise e
|
ipc_extra["pool_handle"], ipc_extra["pool_device_index"]
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
"Failed to deserialize from cached pooled CUDA IPC handle (%s). "
|
||||||
|
"Invalidating cache entry and retrying uncached.",
|
||||||
|
e,
|
||||||
|
)
|
||||||
|
_pool_handle_cache_invalidate(cache_key)
|
||||||
|
(
|
||||||
|
slice_tensor,
|
||||||
|
_target_device,
|
||||||
|
_cache_key,
|
||||||
|
storage_to_cache,
|
||||||
|
) = self._reconstruct_from_ipc_extra(ipc_extra, use_cache=False)
|
||||||
|
if storage_to_cache is not None:
|
||||||
|
_pool_handle_cache_set(cache_key, storage_to_cache)
|
||||||
|
else:
|
||||||
|
# Non-pooled path: open handle directly (original behavior)
|
||||||
|
try:
|
||||||
|
storage = torch.UntypedStorage._new_shared_cuda(
|
||||||
|
*ipc_extra["handle"]
|
||||||
|
)
|
||||||
|
target_device = torch.device(f"cuda:{ipc_extra['device_index']}")
|
||||||
|
with torch.cuda.device(target_device):
|
||||||
|
slice_tensor = torch.empty(
|
||||||
|
0, dtype=ipc_extra["dtype"], device=target_device
|
||||||
|
).set_(
|
||||||
|
storage,
|
||||||
|
storage_offset=ipc_extra["storage_offset"],
|
||||||
|
size=ipc_extra["shape"],
|
||||||
|
stride=ipc_extra["stride"],
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.info("Failed to deserialize from CUDA IPC handle (%s).", e)
|
||||||
|
raise
|
||||||
|
|
||||||
|
reconstructed_tensor = self._copy_slice_tensor_to_target(
|
||||||
|
slice_tensor, rebuild_device, recons_shape, recons_dtype
|
||||||
|
)
|
||||||
elif isinstance(self.proxy_state["tensor_data"], torch.Tensor):
|
elif isinstance(self.proxy_state["tensor_data"], torch.Tensor):
|
||||||
reconstructed_tensor = self.proxy_state["tensor_data"].to(
|
reconstructed_tensor = self.proxy_state["tensor_data"].to(
|
||||||
rebuild_device, non_blocking=True
|
rebuild_device, non_blocking=True
|
||||||
|
|||||||
Reference in New Issue
Block a user