[perf] Reuse a pooled HTTP session for multimodal URL downloads (#28056)

Signed-off-by: Abatom <abzhonghua@gmail.com>
This commit is contained in:
Zhonghua Deng
2026-06-15 17:47:47 +08:00
committed by GitHub
parent bf186cf8fc
commit 2a33724c9b
2 changed files with 35 additions and 19 deletions
@@ -10,9 +10,9 @@ from typing import Optional
import numpy as np import numpy as np
import pybase64 import pybase64
import requests
import torch import torch
from sglang.srt.utils import common
from sglang.utils import logger from sglang.utils import logger
try: try:
@@ -138,8 +138,6 @@ class MiMoAudioPipeline:
) )
self._resamplers_max = max_resamplers self._resamplers_max = max_resamplers
self.http_session = requests.Session()
@property @property
def audio_token_per_second(self) -> float: def audio_token_per_second(self) -> float:
return self.audio_input_id_per_second / self.audio_group_size return self.audio_input_id_per_second / self.audio_group_size
@@ -188,18 +186,18 @@ class MiMoAudioPipeline:
dl_start = time.perf_counter() dl_start = time.perf_counter()
timeout = int(os.getenv("REQUEST_TIMEOUT", "5")) timeout = int(os.getenv("REQUEST_TIMEOUT", "5"))
try: try:
response = self.http_session.get( with common.get_mm_http_session().get(
audio, stream=True, timeout=timeout audio, stream=True, timeout=timeout
) ) as response:
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000 response.raise_for_status()
if dl_elapsed_ms > 1000.0: dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
content_len = len(response.content) if dl_elapsed_ms > 1000.0:
logger.warning( content_len = len(response.content)
f"Slow audio download: {dl_elapsed_ms:.2f}ms, " logger.warning(
f"size={content_len / 1024:.1f}KB, url={audio}" 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() file = io.BytesIO(response.content)
except Exception as e: except Exception as e:
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000 dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
logger.error( logger.error(
+23 -5
View File
@@ -782,6 +782,22 @@ def set_random_seed(seed: int) -> None:
torch.xpu.manual_seed_all(seed) 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( def load_audio(
audio_file: str, sr: Optional[int] = None, mono: bool = True audio_file: str, sr: Optional[int] = None, mono: bool = True
) -> np.ndarray: ) -> np.ndarray:
@@ -797,7 +813,7 @@ def load_audio(
audio_file.startswith("http://") or audio_file.startswith("https://") audio_file.startswith("http://") or audio_file.startswith("https://")
): ):
timeout = int(os.getenv("REQUEST_TIMEOUT", "5")) 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() response.raise_for_status()
source = response.content source = response.content
elif isinstance(audio_file, str) and audio_file.startswith("file://"): 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 return image_file
if image_file.startswith(("http://", "https://")): if image_file.startswith(("http://", "https://")):
timeout = int(os.getenv("REQUEST_TIMEOUT", "3")) 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: try:
response.raise_for_status() response.raise_for_status()
result = response.content result = response.content
@@ -990,9 +1006,11 @@ def _normalize_video_input(
elif isinstance(video_file, str): elif isinstance(video_file, str):
if video_file.startswith(("http://", "https://")): if video_file.startswith(("http://", "https://")):
timeout = int(os.getenv("REQUEST_TIMEOUT", "10")) timeout = int(os.getenv("REQUEST_TIMEOUT", "10"))
response = requests.get(video_file, stream=True, timeout=timeout) with get_mm_http_session().get(
response.raise_for_status() video_file, stream=True, timeout=timeout
return response.content ) as response:
response.raise_for_status()
return response.content
elif video_file.startswith("data:"): elif video_file.startswith("data:"):
_, encoded = video_file.split(",", 1) _, encoded = video_file.split(",", 1)
return pybase64.b64decode(encoded, validate=True) return pybase64.b64decode(encoded, validate=True)