Support EPD error handling (#16670)
Co-authored-by: ZhengWG <zwg0606@gmail.com>
This commit is contained in:
@@ -4,7 +4,8 @@ import pickle
|
||||
import random
|
||||
import threading
|
||||
import uuid
|
||||
from typing import List, Optional
|
||||
from enum import IntEnum
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
import aiohttp
|
||||
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.managers.io_struct import TokenizedGenerateReqInput
|
||||
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.utils import get_local_ip_auto, get_zmq_socket_on_host
|
||||
from sglang.srt.utils.hf_transformers_utils import get_processor
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
|
||||
|
||||
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.num_parts = num_parts
|
||||
self.part_idx = part_idx
|
||||
@@ -42,6 +56,8 @@ class EmbeddingData:
|
||||
self.image_grid_dim if i == self.part_idx else None
|
||||
for i in range(self.num_parts)
|
||||
]
|
||||
self.error_msg = error_msg
|
||||
self.error_code = error_code
|
||||
|
||||
def add(self, embedding_data):
|
||||
assert self.req_id == embedding_data.req_id
|
||||
@@ -66,7 +82,7 @@ class EmbeddingData:
|
||||
return sum(self.ready_list) == self.num_parts
|
||||
|
||||
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):
|
||||
new_data = EmbeddingData(
|
||||
@@ -74,6 +90,8 @@ class EmbeddingData:
|
||||
num_parts=self.num_parts,
|
||||
part_idx=self.part_idx,
|
||||
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.dtype = self.dtype
|
||||
@@ -81,6 +99,12 @@ class EmbeddingData:
|
||||
return new_data
|
||||
|
||||
|
||||
class WaitingImageRequestStatus(IntEnum):
|
||||
FAIL = -1
|
||||
PENDING = 0
|
||||
SUCCESS = 1
|
||||
|
||||
|
||||
# For zmq_to_scheduler
|
||||
class WaitingImageRequest:
|
||||
def __init__(
|
||||
@@ -107,7 +131,10 @@ class WaitingImageRequest:
|
||||
)
|
||||
logger.info(f"Waiting for input {self.embedding_port = }")
|
||||
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):
|
||||
async def _send_single_request(session, url, payload):
|
||||
@@ -163,7 +190,7 @@ class WaitingImageRequest:
|
||||
)
|
||||
|
||||
def _try_recv_mm_data(self):
|
||||
if self.ready:
|
||||
if self.status != WaitingImageRequestStatus.PENDING:
|
||||
return
|
||||
while self.recv_embedding_data is None or not self.recv_embedding_data.ready:
|
||||
try:
|
||||
@@ -171,8 +198,17 @@ class WaitingImageRequest:
|
||||
except zmq.Again:
|
||||
# No data available yet, wait a bit and retry
|
||||
return
|
||||
|
||||
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]
|
||||
recv_obj.embedding = torch.frombuffer(buffer, dtype=recv_obj.dtype).reshape(
|
||||
recv_obj.shape
|
||||
@@ -191,7 +227,7 @@ class WaitingImageRequest:
|
||||
)
|
||||
self.recv_req.mm_inputs = mm_inputs
|
||||
self.recv_req.input_ids = mm_inputs["input_ids"]
|
||||
self.ready = True
|
||||
self.status = WaitingImageRequestStatus.SUCCESS
|
||||
self.recv_socket.close()
|
||||
|
||||
|
||||
@@ -215,6 +251,7 @@ class MMReceiver:
|
||||
pp_rank: Optional[int] = None,
|
||||
tp_rank: Optional[int] = None,
|
||||
tp_group: Optional[GroupCoordinator] = None,
|
||||
scheduler: Optional["Scheduler"] = None,
|
||||
):
|
||||
self.context = zmq.asyncio.Context(20)
|
||||
self.encoder_transfer_backend = server_args.encoder_transfer_backend
|
||||
@@ -237,6 +274,7 @@ class MMReceiver:
|
||||
self.nnodes = server_args.nnodes
|
||||
self.hostname = get_local_ip_auto()
|
||||
self.waiting_list: List[WaitingImageRequest] = []
|
||||
self.scheduler = scheduler
|
||||
if hf_config is not None:
|
||||
transport_mode = _determine_tensor_transport_mode(server_args)
|
||||
import_processors("sglang.srt.multimodal.processors")
|
||||
@@ -268,6 +306,41 @@ class MMReceiver:
|
||||
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
|
||||
def process_waiting_requests(self, recv_reqs):
|
||||
new_recv_reqs = []
|
||||
@@ -290,12 +363,12 @@ class MMReceiver:
|
||||
new_recv_reqs.append(recv_req)
|
||||
|
||||
if len(self.waiting_list) == 0:
|
||||
return new_recv_reqs
|
||||
return new_recv_reqs, []
|
||||
|
||||
local_status = []
|
||||
for waiting_req in self.waiting_list:
|
||||
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)
|
||||
|
||||
@@ -306,14 +379,27 @@ class MMReceiver:
|
||||
)
|
||||
|
||||
new_waiting = []
|
||||
abort_reqs = []
|
||||
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)
|
||||
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)
|
||||
|
||||
self.waiting_list = new_waiting
|
||||
return new_recv_reqs
|
||||
return new_recv_reqs, abort_reqs
|
||||
|
||||
# For zmq_to_scheduler
|
||||
def _run_encode_in_thread(
|
||||
@@ -389,6 +475,16 @@ class MMReceiver:
|
||||
]
|
||||
|
||||
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 = [
|
||||
await response.json() for response in responses
|
||||
]
|
||||
|
||||
@@ -7,6 +7,7 @@ import os
|
||||
import pickle
|
||||
import time
|
||||
import traceback
|
||||
from http import HTTPStatus
|
||||
from typing import Dict, List, Optional, Set, Tuple
|
||||
|
||||
import aiohttp
|
||||
@@ -52,12 +53,30 @@ logger = logging.getLogger(__name__)
|
||||
rid_lock = asyncio.Lock()
|
||||
rid_to_receive_endpoint: Dict[str, List[str]] = dict()
|
||||
rid_to_receive_count: Dict[str, int] = dict()
|
||||
rid_to_err_msg: Dict[str, str] = dict()
|
||||
|
||||
use_image_processor_gpu = (
|
||||
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:
|
||||
"""Wrapper to keep tensor alive while exposing buffer for zero-copy."""
|
||||
|
||||
@@ -196,6 +215,7 @@ class MMEncoder:
|
||||
)
|
||||
|
||||
self.embedding_to_send = dict()
|
||||
self.background_tasks: Set[asyncio.Task] = set()
|
||||
|
||||
logger.info(f"rank {rank} init finish ")
|
||||
|
||||
@@ -291,53 +311,56 @@ class MMEncoder:
|
||||
return await asyncio.gather(*async_futures)
|
||||
|
||||
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 {}
|
||||
images_input = self.image_processor(images=images, **kwargs)
|
||||
feature = images_input["pixel_values"]
|
||||
mm_item = MultimodalDataItem.from_dict(
|
||||
{
|
||||
"modality": Modality.IMAGE,
|
||||
"feature": _convert(feature),
|
||||
}
|
||||
)
|
||||
for k, v in images_input.items():
|
||||
if k == "pixel_values":
|
||||
continue
|
||||
mm_item.set(k, _convert(v))
|
||||
try:
|
||||
kwargs = {"device": self.device} if self.use_image_processor_gpu else {}
|
||||
images_input = self.image_processor(images=images, **kwargs)
|
||||
feature = images_input["pixel_values"]
|
||||
mm_item = MultimodalDataItem.from_dict(
|
||||
{
|
||||
"modality": Modality.IMAGE,
|
||||
"feature": _convert(feature),
|
||||
}
|
||||
)
|
||||
for k, v in images_input.items():
|
||||
if k == "pixel_values":
|
||||
continue
|
||||
mm_item.set(k, _convert(v))
|
||||
|
||||
# support mm_cache
|
||||
mm_embedding = None
|
||||
mm_hash = None
|
||||
# support mm_cache
|
||||
mm_embedding = None
|
||||
mm_hash = None
|
||||
|
||||
start_time = time.perf_counter()
|
||||
if self.server_args.enable_prefix_mm_cache:
|
||||
mm_item.set_pad_value()
|
||||
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
|
||||
async with self.mm_cache_lock:
|
||||
mm_cache = self.mm_cache.get([mm_item.hash])
|
||||
if mm_cache is not None:
|
||||
mm_embedding = mm_cache.embedding
|
||||
if self.server_args.enable_prefix_mm_cache:
|
||||
mm_item.set_pad_value()
|
||||
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
|
||||
async with self.mm_cache_lock:
|
||||
mm_cache = self.mm_cache.get([mm_item.hash])
|
||||
if mm_cache is not None:
|
||||
mm_embedding = mm_cache.embedding
|
||||
|
||||
if mm_embedding is None:
|
||||
with torch.inference_mode():
|
||||
mm_embedding: torch.Tensor = self.model.get_image_feature([mm_item])
|
||||
mm_embedding = mm_embedding.cpu()
|
||||
if len(mm_embedding.shape) != 2:
|
||||
mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1])
|
||||
if mm_embedding is None:
|
||||
with torch.inference_mode():
|
||||
mm_embedding: torch.Tensor = self.model.get_image_feature([mm_item])
|
||||
mm_embedding = mm_embedding.cpu()
|
||||
if len(mm_embedding.shape) != 2:
|
||||
mm_embedding = mm_embedding.reshape(-1, mm_embedding.shape[-1])
|
||||
|
||||
if self.server_args.enable_prefix_mm_cache:
|
||||
async with self.mm_cache_lock:
|
||||
self.mm_cache.set(mm_hash, EmbeddingResult(embedding=mm_embedding))
|
||||
end_time = time.perf_counter()
|
||||
logger.info(
|
||||
f"Vit time : {(end_time - start_time)*1000:.2f} ms {mm_embedding.shape = }"
|
||||
)
|
||||
if self.profiler is not None:
|
||||
self.profiler.step()
|
||||
if self.server_args.enable_prefix_mm_cache:
|
||||
async with self.mm_cache_lock:
|
||||
self.mm_cache.set(mm_hash, EmbeddingResult(embedding=mm_embedding))
|
||||
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(
|
||||
self,
|
||||
@@ -377,26 +400,47 @@ class MMEncoder:
|
||||
socket.send_multipart([pickle.dumps(mm_data)])
|
||||
else:
|
||||
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)
|
||||
socket.send_multipart(
|
||||
[pickle.dumps(new_mm_data), embedding_tensor.__buffer__()]
|
||||
)
|
||||
|
||||
async def encode(self, mm_items, req_id, num_parts, part_idx):
|
||||
start_time = time.time()
|
||||
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:
|
||||
mm_data = EmbeddingData(
|
||||
req_id,
|
||||
num_parts,
|
||||
part_idx,
|
||||
image_grid_dim,
|
||||
mm_embedding,
|
||||
try:
|
||||
image_grid_dim, mm_embedding = await self._encode(mm_items)
|
||||
|
||||
if self.rank == 0:
|
||||
mm_data = EmbeddingData(
|
||||
req_id, num_parts, part_idx, image_grid_dim, mm_embedding
|
||||
)
|
||||
self.embedding_to_send[req_id] = mm_data
|
||||
return (
|
||||
mm_embedding.nbytes,
|
||||
mm_embedding.shape[0],
|
||||
mm_embedding.shape[1],
|
||||
None,
|
||||
None,
|
||||
)
|
||||
self.embedding_to_send[mm_data.req_id] = mm_data
|
||||
return mm_embedding.nbytes, mm_embedding.shape[0], mm_embedding.shape[1]
|
||||
except Exception as e:
|
||||
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
|
||||
async def send(
|
||||
@@ -624,55 +668,93 @@ def launch_server(server_args: ServerArgs):
|
||||
|
||||
@app.post("/encode")
|
||||
async def handle_encode_request(request: dict):
|
||||
# broadcast request
|
||||
request.update({"enter_time": time.time()})
|
||||
for socket in send_sockets:
|
||||
socket.send_pyobj(request)
|
||||
req_id = request["req_id"]
|
||||
try:
|
||||
|
||||
nbytes, embedding_len, embedding_dim = await encoder.encode(
|
||||
mm_items=request["mm_items"],
|
||||
req_id=request["req_id"],
|
||||
num_parts=request["num_parts"],
|
||||
part_idx=request["part_idx"],
|
||||
)
|
||||
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(
|
||||
def start_background_send(req_id):
|
||||
task = asyncio.create_task(encoder.send_with_url(req_id=req_id))
|
||||
encoder.background_tasks.add(task)
|
||||
task.add_done_callback(encoder.background_tasks.discard)
|
||||
|
||||
# broadcast request
|
||||
request.update({"enter_time": time.time()})
|
||||
for socket in send_sockets:
|
||||
socket.send_pyobj(request)
|
||||
|
||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
||||
await encoder.encode(
|
||||
mm_items=request["mm_items"],
|
||||
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")
|
||||
@@ -746,7 +828,9 @@ async def start_profile_async(obj: Optional[ProfileReqInput] = None):
|
||||
f"profile_id={encoder.profiler.profile_id}\n"
|
||||
)
|
||||
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"])
|
||||
@@ -754,11 +838,15 @@ async def stop_profile_async():
|
||||
if encoder is None:
|
||||
return Response(content="encoder not ready\n", status_code=503)
|
||||
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)
|
||||
for socket in send_sockets:
|
||||
socket.send_pyobj(req)
|
||||
ok, msg = encoder.profiler.stop()
|
||||
if ok:
|
||||
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
|
||||
)
|
||||
|
||||
@@ -959,9 +959,10 @@ class Scheduler(
|
||||
self.mm_receiver = MMReceiver(
|
||||
self.server_args,
|
||||
hf_config=self.model_config.hf_config,
|
||||
tp_rank=self.tp_rank,
|
||||
pp_rank=self.pp_rank,
|
||||
tp_rank=self.tp_rank,
|
||||
tp_group=self.tp_group,
|
||||
scheduler=self,
|
||||
)
|
||||
|
||||
def init_overlap(self):
|
||||
@@ -1260,7 +1261,16 @@ class Scheduler(
|
||||
and self.server_args.language_only
|
||||
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:
|
||||
for req in recv_reqs:
|
||||
|
||||
@@ -1072,7 +1072,6 @@ class SchedulerOutputProcessorMixin:
|
||||
if reqs or is_idle_batch:
|
||||
if self.model_config.is_multimodal_gen:
|
||||
return
|
||||
|
||||
self.send_to_detokenizer.send_output(
|
||||
BatchTokenIDOutput(
|
||||
rids=rids,
|
||||
|
||||
Reference in New Issue
Block a user