[BugFix][EPD] Harden zmq_to_scheduler receiver failures; sync error info across TP (#31592)

Co-authored-by: siyu <liusy58@linux.alibaba.com>
This commit is contained in:
Zheng Wengang
2026-07-24 16:37:52 +08:00
committed by GitHub
co-authored by siyu
parent 58f417049d
commit 364b5f23e6
2 changed files with 105 additions and 45 deletions
@@ -431,7 +431,9 @@ class EmbeddingData:
self.shape = list(embedding.shape) if embedding is not None else None
self.cached_embedding = None
self.error_msg = error_msg
self.error_code = error_code
# Coerce to plain int: this object crosses process boundaries via
# safe_pickle_loads, whose allowlist blocks http.HTTPStatus.
self.error_code = int(error_code) if error_code is not None else None
# Store additional metadata (e.g., video_timestamps for qwen3_vl)
for key, value in kwargs.items():
setattr(self, key, value)
@@ -839,51 +841,78 @@ class WaitingImageRequest:
except zmq.Again:
# No data available yet, wait a bit and retry
return
recv_obj: EmbeddingData = safe_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 = }"
try:
recv_obj: EmbeddingData = safe_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
# Extract original req_id from part_req_id and drop stale payloads
# that may arrive on a reused ZMQ port after a prior request aborted.
original_req_id = extract_original_req_id(recv_obj.req_id)
if original_req_id != self.recv_req.rid:
logger.warning(
f"Dropping stale embedding data: expected rid={self.recv_req.rid}, "
f"got rid={recv_obj.req_id} (likely from ZMQ port reuse)"
)
continue
recv_obj.req_id = original_req_id
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)
.clone()
)
self.error_msg = recv_obj.error_msg
self.error_code = recv_obj.error_code
if self.recv_embedding_data is None:
self.recv_embedding_data = (
MultiModalEmbeddingData.from_embedding_data(
recv_obj, model_type=self.model_type
)
)
else:
self.recv_embedding_data.add(recv_obj)
except Exception as e:
# A message the scheduler cannot decode (blocked unpickle,
# bad shape/dtype, ...) must fail this request, not crash the
# scheduler event loop; FAIL still reaches the TP-wide status
# all-reduce in _process_waiting_requests.
logger.exception(
"Failed to decode embedding message for rid=%s", self.rid
)
self.error_msg = f"Failed to decode embedding message: {e}"
self.status = WaitingImageRequestStatus.FAIL
self._cleanup_gpu_buffer()
self.recv_socket.close()
return
# Extract original req_id from part_req_id and drop stale payloads
# that may arrive on a reused ZMQ port after a prior request aborted.
original_req_id = extract_original_req_id(recv_obj.req_id)
if original_req_id != self.recv_req.rid:
logger.warning(
f"Dropping stale embedding data: expected rid={self.recv_req.rid}, "
f"got rid={recv_obj.req_id} (likely from ZMQ port reuse)"
)
continue
recv_obj.req_id = original_req_id
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)
.clone()
# Assemble mm_inputs. Wrapped so an assembly failure still reaches the
# TP-wide status all-reduce in _process_waiting_requests instead of
# raising past it.
try:
recv_embedding = self.recv_embedding_data.get_embedding(is_concat=True)
mm_inputs = self.mm_processor.get_mm_data(
self.recv_req.input_text,
recv_embedding,
**self.recv_embedding_data.get_mm_extra_meta(),
)
if self.recv_embedding_data is None:
self.recv_embedding_data = MultiModalEmbeddingData.from_embedding_data(
recv_obj, model_type=self.model_type
)
else:
self.recv_embedding_data.add(recv_obj)
recv_embedding = self.recv_embedding_data.get_embedding(is_concat=True)
mm_inputs = self.mm_processor.get_mm_data(
self.recv_req.input_text,
recv_embedding,
**self.recv_embedding_data.get_mm_extra_meta(),
)
self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = array("q", mm_inputs.input_ids)
self.status = WaitingImageRequestStatus.SUCCESS
self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = array("q", mm_inputs.input_ids)
self.status = WaitingImageRequestStatus.SUCCESS
except Exception as e:
logger.exception(
"Failed to assemble multimodal inputs for rid=%s", self.rid
)
self.status = WaitingImageRequestStatus.FAIL
self.error_msg = f"Failed to assemble multimodal inputs: {e}"
self._cleanup_gpu_buffer()
self.recv_socket.close()
def _cleanup_gpu_buffer(self):
@@ -1177,6 +1206,10 @@ class WaitingImageRDMARequest(WaitingImageRequest):
parts = self.recv_socket.recv_multipart(flags=zmq.NOBLOCK, copy=False)
except zmq.Again:
return
except zmq.ZMQError:
# The RDMA pipeline thread closed the socket after an encoder
# error (e.g. OOM). It already set status=FAIL; just bail.
return
recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None:
@@ -1803,6 +1836,31 @@ class MMReceiverBase(ABC):
)
obj.need_wait_for_mm_inputs = False
def _sync_fail_info_across_tp(self, waiting_req: WaitingImageRequest) -> None:
"""Share encoder error fields across TP ranks before abort.
The encoder sends ZMQ error signals to each TP rank's receive socket,
but they can arrive at different times. ``all_reduce`` on status makes
every rank enter FAIL together while only some ranks have populated
``error_msg`` / ``error_code``. attn_tp_rank 0 streams the abort to the
client, so merge the best-known payload from all ranks first.
"""
if self.tp_size <= 1 or self.tp_group is None:
return
gathered = self.tp_group.all_gather_object(
(waiting_req.error_msg, waiting_req.error_code)
)
best_msg = waiting_req.error_msg
best_code = waiting_req.error_code
for msg, code in gathered:
if msg is not None:
best_msg = msg
if code is not None:
best_code = code
waiting_req.error_msg = best_msg
waiting_req.error_code = best_code
# For zmq_to_scheduler
def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs):
new_recv_reqs = []
@@ -1861,6 +1919,7 @@ class MMReceiverBase(ABC):
if status_value == WaitingImageRequestStatus.SUCCESS:
new_recv_reqs.append(waiting_req.recv_req)
elif status_value == WaitingImageRequestStatus.FAIL:
self._sync_fail_info_across_tp(waiting_req)
logger.error(
f"Waiting request {waiting_req.rid} failed: {waiting_req.error_msg} {waiting_req.error_code = }"
)
@@ -226,11 +226,12 @@ class SchedulerRequestReceiver:
):
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
)
if error_code is None:
status_code = HTTPStatus.INTERNAL_SERVER_ERROR
elif isinstance(error_code, HTTPStatus):
status_code = error_code
else:
status_code = HTTPStatus(int(error_code))
prepare_abort(req, error_msg, status_code=status_code)
self.stream_output([req], req.return_logprob)
return recv_reqs