From 2a33724c9b22e403893b68399042eb149473f6c5 Mon Sep 17 00:00:00 2001 From: Zhonghua Deng Date: Mon, 15 Jun 2026 17:47:47 +0800 Subject: [PATCH] [perf] Reuse a pooled HTTP session for multimodal URL downloads (#28056) Signed-off-by: Abatom --- .../srt/multimodal/processors/mimo_audio.py | 26 ++++++++--------- python/sglang/srt/utils/common.py | 28 +++++++++++++++---- 2 files changed, 35 insertions(+), 19 deletions(-) diff --git a/python/sglang/srt/multimodal/processors/mimo_audio.py b/python/sglang/srt/multimodal/processors/mimo_audio.py index ae21abd09..47295e9a8 100644 --- a/python/sglang/srt/multimodal/processors/mimo_audio.py +++ b/python/sglang/srt/multimodal/processors/mimo_audio.py @@ -10,9 +10,9 @@ from typing import Optional import numpy as np import pybase64 -import requests import torch +from sglang.srt.utils import common from sglang.utils import logger try: @@ -138,8 +138,6 @@ class MiMoAudioPipeline: ) self._resamplers_max = max_resamplers - self.http_session = requests.Session() - @property def audio_token_per_second(self) -> float: return self.audio_input_id_per_second / self.audio_group_size @@ -188,18 +186,18 @@ class MiMoAudioPipeline: dl_start = time.perf_counter() timeout = int(os.getenv("REQUEST_TIMEOUT", "5")) try: - response = self.http_session.get( + with common.get_mm_http_session().get( audio, stream=True, timeout=timeout - ) - dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000 - if dl_elapsed_ms > 1000.0: - content_len = len(response.content) - logger.warning( - f"Slow audio download: {dl_elapsed_ms:.2f}ms, " - f"size={content_len / 1024:.1f}KB, url={audio}" - ) - file = io.BytesIO(response.content) - response.close() + ) as response: + response.raise_for_status() + dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000 + if dl_elapsed_ms > 1000.0: + content_len = len(response.content) + logger.warning( + f"Slow audio download: {dl_elapsed_ms:.2f}ms, " + f"size={content_len / 1024:.1f}KB, url={audio}" + ) + file = io.BytesIO(response.content) except Exception as e: dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000 logger.error( diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index ac44ba8ee..53f578252 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -782,6 +782,22 @@ def set_random_seed(seed: int) -> None: torch.xpu.manual_seed_all(seed) +_mm_http_session = threading.local() + + +def get_mm_http_session() -> requests.Session: + """Per-thread HTTP session for multimodal downloads, to pool/reuse TCP + connections. Pid-checked so a forked worker rebuilds its own, not the parent's. + """ + pid = os.getpid() + session = getattr(_mm_http_session, "session", None) + if session is None or getattr(_mm_http_session, "pid", None) != pid: + session = requests.Session() + _mm_http_session.session = session + _mm_http_session.pid = pid + return session + + def load_audio( audio_file: str, sr: Optional[int] = None, mono: bool = True ) -> np.ndarray: @@ -797,7 +813,7 @@ def load_audio( audio_file.startswith("http://") or audio_file.startswith("https://") ): timeout = int(os.getenv("REQUEST_TIMEOUT", "5")) - with requests.get(audio_file, timeout=timeout) as response: + with get_mm_http_session().get(audio_file, timeout=timeout) as response: response.raise_for_status() source = response.content elif isinstance(audio_file, str) and audio_file.startswith("file://"): @@ -958,7 +974,7 @@ def get_image_bytes(image_file: Union[str, bytes]) -> bytes: return image_file if image_file.startswith(("http://", "https://")): timeout = int(os.getenv("REQUEST_TIMEOUT", "3")) - response = requests.get(image_file, timeout=timeout) + response = get_mm_http_session().get(image_file, timeout=timeout) try: response.raise_for_status() result = response.content @@ -990,9 +1006,11 @@ def _normalize_video_input( elif isinstance(video_file, str): if video_file.startswith(("http://", "https://")): timeout = int(os.getenv("REQUEST_TIMEOUT", "10")) - response = requests.get(video_file, stream=True, timeout=timeout) - response.raise_for_status() - return response.content + with get_mm_http_session().get( + video_file, stream=True, timeout=timeout + ) as response: + response.raise_for_status() + return response.content elif video_file.startswith("data:"): _, encoded = video_file.split(",", 1) return pybase64.b64decode(encoded, validate=True)