Support EPD error handling (#16670)

Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
siyu
2026-01-21 18:47:00 +08:00
committed by GitHub
co-authored by ZhengWG
parent e7224e9681
commit 7520b92927
4 changed files with 310 additions and 117 deletions
@@ -4,7 +4,8 @@ import pickle
import random import random
import threading import threading
import uuid import uuid
from typing import List, Optional from enum import IntEnum
from typing import TYPE_CHECKING, List, Optional
import aiohttp import aiohttp
import torch import torch
@@ -16,15 +17,28 @@ from sglang.srt.disaggregation.mooncake.transfer_engine import MooncakeTransferE
from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.distributed.parallel_state import GroupCoordinator
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.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import get_local_ip_auto, get_zmq_socket_on_host from sglang.srt.utils import get_local_ip_auto, get_zmq_socket_on_host
from sglang.srt.utils.hf_transformers_utils import get_processor from sglang.srt.utils.hf_transformers_utils import get_processor
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from sglang.srt.managers.scheduler import Scheduler
class EmbeddingData: class EmbeddingData:
def __init__(self, req_id, num_parts, part_idx, image_grid_dim, embedding=None): def __init__(
self,
req_id,
num_parts,
part_idx,
image_grid_dim,
embedding=None,
error_msg=None,
error_code=None,
):
self.req_id = req_id self.req_id = req_id
self.num_parts = num_parts self.num_parts = num_parts
self.part_idx = part_idx self.part_idx = part_idx
@@ -42,6 +56,8 @@ class EmbeddingData:
self.image_grid_dim if i == self.part_idx else None self.image_grid_dim if i == self.part_idx else None
for i in range(self.num_parts) for i in range(self.num_parts)
] ]
self.error_msg = error_msg
self.error_code = error_code
def add(self, embedding_data): def add(self, embedding_data):
assert self.req_id == embedding_data.req_id assert self.req_id == embedding_data.req_id
@@ -66,7 +82,7 @@ class EmbeddingData:
return sum(self.ready_list) == self.num_parts return sum(self.ready_list) == self.num_parts
def __repr__(self): def __repr__(self):
return f"EmbeddingData(req_id={self.req_id}, num_parts={self.num_parts}, part_idx={self.part_idx})" return f"EmbeddingData(req_id={self.req_id}, num_parts={self.num_parts}, part_idx={self.part_idx}) error_msg={self.error_msg}"
def copy_without_embedding(self): def copy_without_embedding(self):
new_data = EmbeddingData( new_data = EmbeddingData(
@@ -74,6 +90,8 @@ class EmbeddingData:
num_parts=self.num_parts, num_parts=self.num_parts,
part_idx=self.part_idx, part_idx=self.part_idx,
image_grid_dim=self.image_grid_dim, image_grid_dim=self.image_grid_dim,
error_msg=self.error_msg,
error_code=self.error_code,
) )
new_data.send_time = self.send_time new_data.send_time = self.send_time
new_data.dtype = self.dtype new_data.dtype = self.dtype
@@ -81,6 +99,12 @@ class EmbeddingData:
return new_data return new_data
class WaitingImageRequestStatus(IntEnum):
FAIL = -1
PENDING = 0
SUCCESS = 1
# For zmq_to_scheduler # For zmq_to_scheduler
class WaitingImageRequest: class WaitingImageRequest:
def __init__( def __init__(
@@ -107,7 +131,10 @@ class WaitingImageRequest:
) )
logger.info(f"Waiting for input {self.embedding_port = }") logger.info(f"Waiting for input {self.embedding_port = }")
self.recv_embedding_data = None self.recv_embedding_data = None
self.ready = False # ok=1 pending=0 fail=-1
self.status = WaitingImageRequestStatus.PENDING
self.error_msg = None
self.error_code = None
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):
@@ -163,7 +190,7 @@ class WaitingImageRequest:
) )
def _try_recv_mm_data(self): def _try_recv_mm_data(self):
if self.ready: if self.status != WaitingImageRequestStatus.PENDING:
return return
while self.recv_embedding_data is None or not self.recv_embedding_data.ready: while self.recv_embedding_data is None or not self.recv_embedding_data.ready:
try: try:
@@ -171,8 +198,17 @@ class WaitingImageRequest:
except zmq.Again: except zmq.Again:
# No data available yet, wait a bit and retry # No data available yet, wait a bit and retry
return return
recv_obj: EmbeddingData = pickle.loads(parts[0]) recv_obj: EmbeddingData = pickle.loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None:
logger.warning(
f"Received error signal from encoder for {self.rid}: {recv_obj.error_msg} {recv_obj.error_code = }"
)
self.error_msg = recv_obj.error_msg
self.error_code = recv_obj.error_code
self.status = WaitingImageRequestStatus.FAIL
self.recv_socket.close()
return
buffer = parts[1].buffer if hasattr(parts[1], "buffer") else parts[1] buffer = parts[1].buffer if hasattr(parts[1], "buffer") else parts[1]
recv_obj.embedding = torch.frombuffer(buffer, dtype=recv_obj.dtype).reshape( recv_obj.embedding = torch.frombuffer(buffer, dtype=recv_obj.dtype).reshape(
recv_obj.shape recv_obj.shape
@@ -191,7 +227,7 @@ class WaitingImageRequest:
) )
self.recv_req.mm_inputs = mm_inputs self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = mm_inputs["input_ids"] self.recv_req.input_ids = mm_inputs["input_ids"]
self.ready = True self.status = WaitingImageRequestStatus.SUCCESS
self.recv_socket.close() self.recv_socket.close()
@@ -215,6 +251,7 @@ class MMReceiver:
pp_rank: Optional[int] = None, pp_rank: Optional[int] = None,
tp_rank: Optional[int] = None, tp_rank: Optional[int] = None,
tp_group: Optional[GroupCoordinator] = None, tp_group: Optional[GroupCoordinator] = None,
scheduler: Optional["Scheduler"] = None,
): ):
self.context = zmq.asyncio.Context(20) self.context = zmq.asyncio.Context(20)
self.encoder_transfer_backend = server_args.encoder_transfer_backend self.encoder_transfer_backend = server_args.encoder_transfer_backend
@@ -237,6 +274,7 @@ class MMReceiver:
self.nnodes = server_args.nnodes self.nnodes = server_args.nnodes
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
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")
@@ -268,6 +306,41 @@ class MMReceiver:
hf_config, server_args, _processor, transport_mode hf_config, server_args, _processor, transport_mode
) )
def create_req(self, recv_req):
req = Req(
recv_req.rid,
recv_req.input_text,
recv_req.input_ids,
recv_req.sampling_params,
return_logprob=recv_req.return_logprob,
top_logprobs_num=recv_req.top_logprobs_num,
token_ids_logprob=recv_req.token_ids_logprob,
stream=recv_req.stream,
lora_id=recv_req.lora_id,
input_embeds=recv_req.input_embeds,
custom_logit_processor=recv_req.custom_logit_processor,
require_reasoning=recv_req.require_reasoning,
return_hidden_states=recv_req.return_hidden_states,
return_routed_experts=recv_req.return_routed_experts,
eos_token_ids=self.scheduler.model_config.hf_eos_token_id,
bootstrap_host=recv_req.bootstrap_host,
bootstrap_port=recv_req.bootstrap_port,
bootstrap_room=recv_req.bootstrap_room,
disagg_mode=self.scheduler.disaggregation_mode,
data_parallel_rank=recv_req.data_parallel_rank,
vocab_size=self.scheduler.model_config.vocab_size,
priority=recv_req.priority,
metrics_collector=(
self.scheduler.metrics_collector
if self.scheduler.enable_metrics
else None
),
http_worker_ipc=recv_req.http_worker_ipc,
dllm_config=self.scheduler.dllm_config,
)
req.tokenizer = self.scheduler.tokenizer
return req
# For zmq_to_scheduler # For zmq_to_scheduler
def process_waiting_requests(self, recv_reqs): def process_waiting_requests(self, recv_reqs):
new_recv_reqs = [] new_recv_reqs = []
@@ -290,12 +363,12 @@ class MMReceiver:
new_recv_reqs.append(recv_req) new_recv_reqs.append(recv_req)
if len(self.waiting_list) == 0: if len(self.waiting_list) == 0:
return new_recv_reqs return new_recv_reqs, []
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()
local_status.append(waiting_req.ready) 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)
@@ -306,14 +379,27 @@ class MMReceiver:
) )
new_waiting = [] new_waiting = []
abort_reqs = []
for i, waiting_req in enumerate(self.waiting_list): for i, waiting_req in enumerate(self.waiting_list):
if local_status[i].item(): status_value = local_status[i].item()
if status_value == WaitingImageRequestStatus.SUCCESS:
new_recv_reqs.append(waiting_req.recv_req) new_recv_reqs.append(waiting_req.recv_req)
else: elif status_value == WaitingImageRequestStatus.FAIL:
logger.error(
f"Waiting request {waiting_req.rid} failed: {waiting_req.error_msg} {waiting_req.error_code = }"
)
abort_reqs.append(
(
self.create_req(waiting_req.recv_req),
waiting_req.error_msg,
waiting_req.error_code,
)
)
else: # status_value == WaitingImageRequestStatus.PENDING
new_waiting.append(waiting_req) new_waiting.append(waiting_req)
self.waiting_list = new_waiting self.waiting_list = new_waiting
return new_recv_reqs return new_recv_reqs, abort_reqs
# For zmq_to_scheduler # For zmq_to_scheduler
def _run_encode_in_thread( def _run_encode_in_thread(
@@ -389,6 +475,16 @@ class MMReceiver:
] ]
responses = await asyncio.gather(*tasks) responses = await asyncio.gather(*tasks)
for response in responses:
if response.status != 200:
try:
err_data = await response.json()
msg = err_data.get("message", "Unknown encoder error")
except:
msg = await response.text()
logger.error(f"Encoder returned error {response.status}: {msg}")
return
response_json_list_unsort = [ response_json_list_unsort = [
await response.json() for response in responses await response.json() for response in responses
] ]
+190 -102
View File
@@ -7,6 +7,7 @@ import os
import pickle import pickle
import time import time
import traceback import traceback
from http import HTTPStatus
from typing import Dict, List, Optional, Set, Tuple from typing import Dict, List, Optional, Set, Tuple
import aiohttp import aiohttp
@@ -52,12 +53,30 @@ logger = logging.getLogger(__name__)
rid_lock = asyncio.Lock() 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()
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
) )
class MMError(Exception):
def __init__(self, message, code=HTTPStatus.INTERNAL_SERVER_ERROR):
self.message = message
self.code = code
super().__init__(self.message)
class BadRequestError(MMError):
def __init__(self, message):
super().__init__(message, code=HTTPStatus.BAD_REQUEST)
class InternalError(MMError):
def __init__(self, message):
super().__init__(message, code=HTTPStatus.INTERNAL_SERVER_ERROR)
class TensorWrapper: class TensorWrapper:
"""Wrapper to keep tensor alive while exposing buffer for zero-copy.""" """Wrapper to keep tensor alive while exposing buffer for zero-copy."""
@@ -196,6 +215,7 @@ class MMEncoder:
) )
self.embedding_to_send = dict() self.embedding_to_send = dict()
self.background_tasks: Set[asyncio.Task] = set()
logger.info(f"rank {rank} init finish ") logger.info(f"rank {rank} init finish ")
@@ -291,53 +311,56 @@ class MMEncoder:
return await asyncio.gather(*async_futures) return await asyncio.gather(*async_futures)
async def _encode(self, mm_items) -> torch.Tensor: async def _encode(self, mm_items) -> torch.Tensor:
images = await self._flatten_and_load_images(mm_items) try:
images = await self._flatten_and_load_images(mm_items)
except Exception as e:
raise BadRequestError(f"Failed to load images from input: {str(e)}")
kwargs = {"device": self.device} if self.use_image_processor_gpu else {} try:
images_input = self.image_processor(images=images, **kwargs) kwargs = {"device": self.device} if self.use_image_processor_gpu else {}
feature = images_input["pixel_values"] images_input = self.image_processor(images=images, **kwargs)
mm_item = MultimodalDataItem.from_dict( feature = images_input["pixel_values"]
{ mm_item = MultimodalDataItem.from_dict(
"modality": Modality.IMAGE, {
"feature": _convert(feature), "modality": Modality.IMAGE,
} "feature": _convert(feature),
) }
for k, v in images_input.items(): )
if k == "pixel_values": for k, v in images_input.items():
continue if k == "pixel_values":
mm_item.set(k, _convert(v)) continue
mm_item.set(k, _convert(v))
# support mm_cache # support mm_cache
mm_embedding = None mm_embedding = None
mm_hash = None mm_hash = None
start_time = time.perf_counter() if self.server_args.enable_prefix_mm_cache:
if self.server_args.enable_prefix_mm_cache: mm_item.set_pad_value()
mm_item.set_pad_value() mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash]) async with self.mm_cache_lock:
async with self.mm_cache_lock: mm_cache = self.mm_cache.get([mm_item.hash])
mm_cache = self.mm_cache.get([mm_item.hash]) if mm_cache is not None:
if mm_cache is not None: mm_embedding = mm_cache.embedding
mm_embedding = mm_cache.embedding
if mm_embedding is None: if mm_embedding is None:
with torch.inference_mode(): with torch.inference_mode():
mm_embedding: torch.Tensor = self.model.get_image_feature([mm_item]) mm_embedding: torch.Tensor = self.model.get_image_feature([mm_item])
mm_embedding = mm_embedding.cpu() mm_embedding = mm_embedding.cpu()
if len(mm_embedding.shape) != 2: if len(mm_embedding.shape) != 2:
mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1]) mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1])
if self.server_args.enable_prefix_mm_cache: if self.server_args.enable_prefix_mm_cache:
async with self.mm_cache_lock: async with self.mm_cache_lock:
self.mm_cache.set(mm_hash, EmbeddingResult(embedding=mm_embedding)) self.mm_cache.set(mm_hash, EmbeddingResult(embedding=mm_embedding))
end_time = time.perf_counter() if self.profiler is not None:
logger.info( self.profiler.step()
f"Vit time : {(end_time - start_time)*1000:.2f} ms {mm_embedding.shape = }"
)
if self.profiler is not None:
self.profiler.step()
return _get_image_grid_dim(images_input), mm_embedding return _get_image_grid_dim(images_input), mm_embedding
except BadRequestError as e:
raise BadRequestError(f"Bad request error: {str(e)}")
except Exception as e:
raise InternalError(f"Internal encoding error: {str(e)}")
async def _send( async def _send(
self, self,
@@ -377,26 +400,47 @@ class MMEncoder:
socket.send_multipart([pickle.dumps(mm_data)]) socket.send_multipart([pickle.dumps(mm_data)])
else: else:
new_mm_data = mm_data.copy_without_embedding() new_mm_data = mm_data.copy_without_embedding()
if new_mm_data.error_msg is not None:
socket.send_multipart([pickle.dumps(new_mm_data)])
return
embedding_tensor = TensorWrapper(mm_data.embedding) embedding_tensor = TensorWrapper(mm_data.embedding)
socket.send_multipart( socket.send_multipart(
[pickle.dumps(new_mm_data), embedding_tensor.__buffer__()] [pickle.dumps(new_mm_data), embedding_tensor.__buffer__()]
) )
async def encode(self, mm_items, req_id, num_parts, part_idx): async def encode(self, mm_items, req_id, num_parts, part_idx):
start_time = time.time() try:
image_grid_dim, mm_embedding = await self._encode(mm_items) image_grid_dim, mm_embedding = await self._encode(mm_items)
end_time = time.time()
logger.info(f"🕛 encode cost = {(end_time - start_time) * 1000:.2f}ms") if self.rank == 0:
if self.rank == 0: mm_data = EmbeddingData(
mm_data = EmbeddingData( req_id, num_parts, part_idx, image_grid_dim, mm_embedding
req_id, )
num_parts, self.embedding_to_send[req_id] = mm_data
part_idx, return (
image_grid_dim, mm_embedding.nbytes,
mm_embedding, mm_embedding.shape[0],
mm_embedding.shape[1],
None,
None,
) )
self.embedding_to_send[mm_data.req_id] = mm_data except Exception as e:
return mm_embedding.nbytes, mm_embedding.shape[0], mm_embedding.shape[1] error_code = getattr(e, "code", HTTPStatus.INTERNAL_SERVER_ERROR)
error_msg = str(e)
logger.error(f"Rank {self.rank} encode failed: {error_msg} {error_code = }")
if self.rank == 0:
mm_data = EmbeddingData(
req_id,
num_parts,
part_idx,
None,
error_msg=error_msg,
error_code=error_code,
)
self.embedding_to_send[req_id] = mm_data
logger.debug(f"Created error EmbeddingData: {mm_data}")
return 0, 0, 0, error_msg, error_code
# For zmq_to_tokenizer zmq_to_scheduler and mooncake # For zmq_to_tokenizer zmq_to_scheduler and mooncake
async def send( async def send(
@@ -624,55 +668,93 @@ def launch_server(server_args: ServerArgs):
@app.post("/encode") @app.post("/encode")
async def handle_encode_request(request: dict): async def handle_encode_request(request: dict):
# broadcast request req_id = request["req_id"]
request.update({"enter_time": time.time()}) try:
for socket in send_sockets:
socket.send_pyobj(request)
nbytes, embedding_len, embedding_dim = await encoder.encode( def start_background_send(req_id):
mm_items=request["mm_items"], task = asyncio.create_task(encoder.send_with_url(req_id=req_id))
req_id=request["req_id"], encoder.background_tasks.add(task)
num_parts=request["num_parts"], task.add_done_callback(encoder.background_tasks.discard)
part_idx=request["part_idx"],
) # broadcast request
if encoder.server_args.encoder_transfer_backend == "mooncake": request.update({"enter_time": time.time()})
del request["mm_items"] for socket in send_sockets:
request.update( socket.send_pyobj(request)
{
"embedding_size": nbytes, nbytes, embedding_len, embedding_dim, error_msg, error_code = (
"embedding_len": embedding_len, await encoder.encode(
"embedding_dim": embedding_dim, mm_items=request["mm_items"],
}
)
return ORJSONResponse(content=request)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
logger.info(f"{request['embedding_port'] = }")
if request["embedding_port"] is None:
await encoder.send_with_url(
req_id=request["req_id"], req_id=request["req_id"],
num_parts=request["num_parts"],
part_idx=request["part_idx"],
) )
else:
assert type(request["embedding_port"]) == list
tasks = []
for embedding_port in request["embedding_port"]:
tasks.append(
encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=embedding_port,
)
)
await asyncio.gather(*tasks)
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
await encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=request["embedding_port"],
) )
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None) if error_msg:
if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
if request["embedding_port"] is None:
start_background_send(req_id)
else:
for port in request["embedding_port"]:
await encoder.send(
req_id=req_id,
prefill_host=request["prefill_host"],
embedding_port=port,
)
return ORJSONResponse(
status_code=error_code,
content={"status": "error", "message": error_msg, "req_id": req_id},
)
if encoder.server_args.encoder_transfer_backend == "mooncake":
del request["mm_items"]
request.update(
{
"embedding_size": nbytes,
"embedding_len": embedding_len,
"embedding_dim": embedding_dim,
}
)
return ORJSONResponse(content=request)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
logger.info(f"{request['embedding_port'] = }")
if request["embedding_port"] is None:
await encoder.send_with_url(
req_id=request["req_id"],
)
else:
assert type(request["embedding_port"]) == list
tasks = []
for embedding_port in request["embedding_port"]:
tasks.append(
encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=embedding_port,
)
)
await asyncio.gather(*tasks)
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
await encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
embedding_port=request["embedding_port"],
)
encoder.embedding_to_send.pop(request["req_id"], None)
return ORJSONResponse(content=None)
except Exception as e:
error_msg = str(e)
logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}")
rid_to_err_msg[req_id] = error_msg
return ORJSONResponse(
status_code=HTTPStatus.INTERNAL_SERVER_ERROR,
content={
"status": "error",
"message": error_msg,
"req_id": req_id,
},
)
@app.post("/send") @app.post("/send")
@@ -746,7 +828,9 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None):
f"profile_id={encoder.profiler.profile_id}\n" f"profile_id={encoder.profiler.profile_id}\n"
) )
return Response(content=detail, status_code=200) return Response(content=detail, status_code=200)
return Response(content=(msg or "Start profiling failed.\n"), status_code=400) return Response(
content=(msg or "Start profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST
)
@app.api_route("/stop_profile", methods=["GET", "POST"]) @app.api_route("/stop_profile", methods=["GET", "POST"])
@@ -754,11 +838,15 @@ async def stop_profile_async():
if encoder is None: if encoder is None:
return Response(content="encoder not ready\n", status_code=503) return Response(content="encoder not ready\n", status_code=503)
if encoder.profiler is None: if encoder.profiler is None:
return Response(content="profiling not initialized\n", status_code=400) return Response(
content="profiling not initialized\n", status_code=HTTPStatus.BAD_REQUEST
)
req = ProfileReq(ProfileReqType.STOP_PROFILE) req = ProfileReq(ProfileReqType.STOP_PROFILE)
for socket in send_sockets: for socket in send_sockets:
socket.send_pyobj(req) socket.send_pyobj(req)
ok, msg = encoder.profiler.stop() ok, msg = encoder.profiler.stop()
if ok: if ok:
return Response(content="Stop profiling.\n", status_code=200) return Response(content="Stop profiling.\n", status_code=200)
return Response(content=(msg or "Stop profiling failed.\n"), status_code=400) return Response(
content=(msg or "Stop profiling failed.\n"), status_code=HTTPStatus.BAD_REQUEST
)
+12 -2
View File
@@ -959,9 +959,10 @@ class Scheduler(
self.mm_receiver = MMReceiver( self.mm_receiver = MMReceiver(
self.server_args, self.server_args,
hf_config=self.model_config.hf_config, hf_config=self.model_config.hf_config,
tp_rank=self.tp_rank,
pp_rank=self.pp_rank, pp_rank=self.pp_rank,
tp_rank=self.tp_rank,
tp_group=self.tp_group, tp_group=self.tp_group,
scheduler=self,
) )
def init_overlap(self): def init_overlap(self):
@@ -1260,7 +1261,16 @@ class Scheduler(
and self.server_args.language_only and self.server_args.language_only
and self.server_args.encoder_transfer_backend == "zmq_to_scheduler" and self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
): ):
recv_reqs = self.mm_receiver.process_waiting_requests(recv_reqs) recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
for req, error_msg, error_code in abort_reqs:
status_code = (
HTTPStatus.BAD_REQUEST
if error_code == 400
else HTTPStatus.INTERNAL_SERVER_ERROR
)
prepare_abort(req, error_msg, status_code=status_code)
self.stream_output([req], req.return_logprob)
if self.enable_trace: if self.enable_trace:
for req in recv_reqs: for req in recv_reqs:
@@ -1072,7 +1072,6 @@ class SchedulerOutputProcessorMixin:
if reqs or is_idle_batch: if reqs or is_idle_batch:
if self.model_config.is_multimodal_gen: if self.model_config.is_multimodal_gen:
return return
self.send_to_detokenizer.send_output( self.send_to_detokenizer.send_output(
BatchTokenIDOutput( BatchTokenIDOutput(
rids=rids, rids=rids,