[EPD] Add notification mechanism to fix server hang and add timeout env var (#18229)
This commit is contained in:
@@ -3,9 +3,11 @@ import logging
|
|||||||
import pickle
|
import pickle
|
||||||
import random
|
import random
|
||||||
import threading
|
import threading
|
||||||
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from enum import IntEnum
|
from enum import IntEnum
|
||||||
|
from http import HTTPStatus
|
||||||
from typing import TYPE_CHECKING, List, Optional
|
from typing import TYPE_CHECKING, List, Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
@@ -18,6 +20,7 @@ from sglang.srt.distributed.parallel_state import (
|
|||||||
GroupCoordinator,
|
GroupCoordinator,
|
||||||
get_mooncake_transfer_engine,
|
get_mooncake_transfer_engine,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput
|
from sglang.srt.managers.io_struct import TokenizedGenerateReqInput
|
||||||
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
|
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
@@ -106,6 +109,7 @@ class WaitingImageRequestStatus(IntEnum):
|
|||||||
FAIL = -1
|
FAIL = -1
|
||||||
PENDING = 0
|
PENDING = 0
|
||||||
SUCCESS = 1
|
SUCCESS = 1
|
||||||
|
TIMEOUT = -2
|
||||||
|
|
||||||
|
|
||||||
# For zmq_to_scheduler
|
# For zmq_to_scheduler
|
||||||
@@ -138,6 +142,7 @@ class WaitingImageRequest:
|
|||||||
self.status = WaitingImageRequestStatus.PENDING
|
self.status = WaitingImageRequestStatus.PENDING
|
||||||
self.error_msg = None
|
self.error_msg = None
|
||||||
self.error_code = None
|
self.error_code = None
|
||||||
|
self.start_time = time.time()
|
||||||
|
|
||||||
def send_encode_request(self):
|
def send_encode_request(self):
|
||||||
async def _send_single_request(session, url, payload):
|
async def _send_single_request(session, url, payload):
|
||||||
@@ -300,6 +305,7 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
self.hostname = get_local_ip_auto()
|
self.hostname = get_local_ip_auto()
|
||||||
self.waiting_list: List[WaitingImageRequest] = []
|
self.waiting_list: List[WaitingImageRequest] = []
|
||||||
self.scheduler = scheduler
|
self.scheduler = scheduler
|
||||||
|
self.wait_timeout = envs.SGLANG_ENCODER_RECV_TIMEOUT.get()
|
||||||
if hf_config is not None:
|
if hf_config is not None:
|
||||||
transport_mode = _determine_tensor_transport_mode(server_args)
|
transport_mode = _determine_tensor_transport_mode(server_args)
|
||||||
import_processors("sglang.srt.multimodal.processors")
|
import_processors("sglang.srt.multimodal.processors")
|
||||||
@@ -394,9 +400,12 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
if len(self.waiting_list) == 0:
|
if len(self.waiting_list) == 0:
|
||||||
return new_recv_reqs, []
|
return new_recv_reqs, []
|
||||||
|
|
||||||
|
current_time = time.time()
|
||||||
local_status = []
|
local_status = []
|
||||||
for waiting_req in self.waiting_list:
|
for waiting_req in self.waiting_list:
|
||||||
waiting_req._try_recv_mm_data()
|
waiting_req._try_recv_mm_data()
|
||||||
|
if current_time - waiting_req.start_time > self.wait_timeout:
|
||||||
|
waiting_req.status = WaitingImageRequestStatus.TIMEOUT
|
||||||
local_status.append(waiting_req.status)
|
local_status.append(waiting_req.status)
|
||||||
|
|
||||||
local_status = torch.tensor(local_status, device="cpu", dtype=torch.int32)
|
local_status = torch.tensor(local_status, device="cpu", dtype=torch.int32)
|
||||||
@@ -424,6 +433,17 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
waiting_req.error_code,
|
waiting_req.error_code,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
elif status_value == WaitingImageRequestStatus.TIMEOUT:
|
||||||
|
logger.error(
|
||||||
|
f"Timed out waiting for image embeddings for request {waiting_req.rid}"
|
||||||
|
)
|
||||||
|
abort_reqs.append(
|
||||||
|
(
|
||||||
|
self.create_req(waiting_req.recv_req),
|
||||||
|
f"Timeout waiting for image embedding after {self.wait_timeout}s",
|
||||||
|
HTTPStatus.REQUEST_TIMEOUT,
|
||||||
|
)
|
||||||
|
)
|
||||||
else: # status_value == WaitingImageRequestStatus.PENDING
|
else: # status_value == WaitingImageRequestStatus.PENDING
|
||||||
new_waiting.append(waiting_req)
|
new_waiting.append(waiting_req)
|
||||||
|
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ from sglang.srt.distributed.parallel_state import (
|
|||||||
init_distributed_environment,
|
init_distributed_environment,
|
||||||
initialize_model_parallel,
|
initialize_model_parallel,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import initialize_dp_attention
|
from sglang.srt.layers.dp_attention import initialize_dp_attention
|
||||||
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqInput, ProfileReqType
|
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqInput, ProfileReqType
|
||||||
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
||||||
@@ -55,6 +56,8 @@ rid_lock = asyncio.Lock()
|
|||||||
rid_to_receive_endpoint: Dict[str, List[str]] = dict()
|
rid_to_receive_endpoint: Dict[str, List[str]] = dict()
|
||||||
rid_to_receive_count: Dict[str, int] = dict()
|
rid_to_receive_count: Dict[str, int] = dict()
|
||||||
rid_to_err_msg: Dict[str, str] = dict()
|
rid_to_err_msg: Dict[str, str] = dict()
|
||||||
|
cond_dict_lock = asyncio.Lock()
|
||||||
|
rid_to_cond: Dict[str, asyncio.Condition] = {}
|
||||||
|
|
||||||
use_image_processor_gpu = (
|
use_image_processor_gpu = (
|
||||||
int(os.getenv("SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU", "0")) == 1
|
int(os.getenv("SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU", "0")) == 1
|
||||||
@@ -197,6 +200,7 @@ class MMEncoder:
|
|||||||
self.io_executor = concurrent.futures.ThreadPoolExecutor(
|
self.io_executor = concurrent.futures.ThreadPoolExecutor(
|
||||||
max_workers=int(os.environ.get("SGLANG_ENCODER_MM_LOAD_WORKERS", 4))
|
max_workers=int(os.environ.get("SGLANG_ENCODER_MM_LOAD_WORKERS", 4))
|
||||||
)
|
)
|
||||||
|
self.send_timeout = envs.SGLANG_ENCODER_SEND_TIMEOUT.get()
|
||||||
|
|
||||||
if schedule_path is not None:
|
if schedule_path is not None:
|
||||||
self.schedule_socket = get_zmq_socket(
|
self.schedule_socket = get_zmq_socket(
|
||||||
@@ -475,7 +479,8 @@ class MMEncoder:
|
|||||||
sent_urls: Set[str] = set()
|
sent_urls: Set[str] = set()
|
||||||
all_tasks: List[Tuple[asyncio.Task, str]] = []
|
all_tasks: List[Tuple[asyncio.Task, str]] = []
|
||||||
start_time = asyncio.get_running_loop().time()
|
start_time = asyncio.get_running_loop().time()
|
||||||
timeout = 60.0
|
timeout = self.send_timeout
|
||||||
|
cond = await get_condition(req_id)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
while True:
|
while True:
|
||||||
@@ -504,14 +509,18 @@ class MMEncoder:
|
|||||||
f"All {expected_count} endpoints initiated for {req_id}. Breaking loop."
|
f"All {expected_count} endpoints initiated for {req_id}. Breaking loop."
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
|
remaining = timeout - (asyncio.get_running_loop().time() - start_time)
|
||||||
if asyncio.get_running_loop().time() - start_time > timeout:
|
if remaining <= 0:
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Timeout waiting for all endpoints for {req_id}. Initiated {len(sent_urls)}/{expected_count}"
|
f"[{req_id}] Timeout! Sent {len(sent_urls)}/{expected_count}"
|
||||||
)
|
)
|
||||||
break
|
break
|
||||||
|
|
||||||
await asyncio.sleep(0.001)
|
async with cond:
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(cond.wait(), timeout=remaining)
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
continue
|
||||||
|
|
||||||
if all_tasks:
|
if all_tasks:
|
||||||
logger.info(
|
logger.info(
|
||||||
@@ -535,6 +544,8 @@ class MMEncoder:
|
|||||||
async with rid_lock:
|
async with rid_lock:
|
||||||
rid_to_receive_endpoint.pop(req_id, None)
|
rid_to_receive_endpoint.pop(req_id, None)
|
||||||
rid_to_receive_count.pop(req_id, None)
|
rid_to_receive_count.pop(req_id, None)
|
||||||
|
async with cond_dict_lock:
|
||||||
|
rid_to_cond.pop(req_id, None)
|
||||||
self.embedding_to_send.pop(req_id, None)
|
self.embedding_to_send.pop(req_id, None)
|
||||||
|
|
||||||
async def get_embedding_port(self, prefill_url):
|
async def get_embedding_port(self, prefill_url):
|
||||||
@@ -674,6 +685,13 @@ def launch_server(server_args: ServerArgs):
|
|||||||
uvicorn.run(app, host=server_args.host, port=server_args.port)
|
uvicorn.run(app, host=server_args.host, port=server_args.port)
|
||||||
|
|
||||||
|
|
||||||
|
async def get_condition(rid):
|
||||||
|
async with cond_dict_lock:
|
||||||
|
if rid not in rid_to_cond:
|
||||||
|
rid_to_cond[rid] = asyncio.Condition()
|
||||||
|
return rid_to_cond[rid]
|
||||||
|
|
||||||
|
|
||||||
@app.post("/encode")
|
@app.post("/encode")
|
||||||
async def handle_encode_request(request: dict):
|
async def handle_encode_request(request: dict):
|
||||||
req_id = request["req_id"]
|
req_id = request["req_id"]
|
||||||
@@ -789,6 +807,9 @@ async def handle_scheduler_receive_url_request(request: dict):
|
|||||||
rid_to_receive_count[rid] = request["receive_count"]
|
rid_to_receive_count[rid] = request["receive_count"]
|
||||||
assert rid_to_receive_count[rid] == request["receive_count"]
|
assert rid_to_receive_count[rid] == request["receive_count"]
|
||||||
rid_to_receive_endpoint[rid].add(request["receive_url"])
|
rid_to_receive_endpoint[rid].add(request["receive_url"])
|
||||||
|
cond = await get_condition(rid)
|
||||||
|
async with cond:
|
||||||
|
cond.notify_all()
|
||||||
|
|
||||||
|
|
||||||
@app.get("/health")
|
@app.get("/health")
|
||||||
|
|||||||
@@ -466,6 +466,10 @@ class Envs:
|
|||||||
SGLANG_USE_AITER_FP8_PER_TOKEN = EnvBool(False)
|
SGLANG_USE_AITER_FP8_PER_TOKEN = EnvBool(False)
|
||||||
# fmt: on
|
# fmt: on
|
||||||
|
|
||||||
|
# EPD
|
||||||
|
SGLANG_ENCODER_RECV_TIMEOUT = EnvFloat(180.0)
|
||||||
|
SGLANG_ENCODER_SEND_TIMEOUT = EnvFloat(180.0)
|
||||||
|
|
||||||
|
|
||||||
envs = Envs()
|
envs = Envs()
|
||||||
EnvField._allow_set_name = False
|
EnvField._allow_set_name = False
|
||||||
|
|||||||
@@ -224,6 +224,8 @@ class SchedulerMetricsMixin:
|
|||||||
else:
|
else:
|
||||||
msg += f"input throughput (token/s): {self.last_input_throughput:.2f}, "
|
msg += f"input throughput (token/s): {self.last_input_throughput:.2f}, "
|
||||||
|
|
||||||
|
if self.server_args.language_only:
|
||||||
|
msg += f"waiting-image-req: {len(self.mm_receiver.waiting_list)}, "
|
||||||
graph_backend = defaultdict(
|
graph_backend = defaultdict(
|
||||||
lambda: "cuda graph",
|
lambda: "cuda graph",
|
||||||
{
|
{
|
||||||
@@ -398,6 +400,9 @@ class SchedulerMetricsMixin:
|
|||||||
msg += f"#transfer-req: {len(self.disagg_decode_transfer_queue.queue)}, "
|
msg += f"#transfer-req: {len(self.disagg_decode_transfer_queue.queue)}, "
|
||||||
msg += f"#retracted-req: {len(self.disagg_decode_prealloc_queue.retracted_queue)}, "
|
msg += f"#retracted-req: {len(self.disagg_decode_prealloc_queue.retracted_queue)}, "
|
||||||
|
|
||||||
|
if self.server_args.language_only:
|
||||||
|
msg += f"waiting-image-req: {len(self.mm_receiver.waiting_list)}, "
|
||||||
|
|
||||||
graph_backend = defaultdict(
|
graph_backend = defaultdict(
|
||||||
lambda: "cuda graph",
|
lambda: "cuda graph",
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user