[perf] Reuse a pooled HTTP session for multimodal URL downloads (#28056)
Signed-off-by: Abatom <abzhonghua@gmail.com>
This commit is contained in:
@@ -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,9 +186,10 @@ 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:
|
||||||
|
response.raise_for_status()
|
||||||
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
|
dl_elapsed_ms = (time.perf_counter() - dl_start) * 1000
|
||||||
if dl_elapsed_ms > 1000.0:
|
if dl_elapsed_ms > 1000.0:
|
||||||
content_len = len(response.content)
|
content_len = len(response.content)
|
||||||
@@ -199,7 +198,6 @@ class MiMoAudioPipeline:
|
|||||||
f"size={content_len / 1024:.1f}KB, url={audio}"
|
f"size={content_len / 1024:.1f}KB, url={audio}"
|
||||||
)
|
)
|
||||||
file = io.BytesIO(response.content)
|
file = io.BytesIO(response.content)
|
||||||
response.close()
|
|
||||||
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(
|
||||||
|
|||||||
@@ -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,7 +1006,9 @@ 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(
|
||||||
|
video_file, stream=True, timeout=timeout
|
||||||
|
) as response:
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
return response.content
|
return response.content
|
||||||
elif video_file.startswith("data:"):
|
elif video_file.startswith("data:"):
|
||||||
|
|||||||
Reference in New Issue
Block a user