[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.shape = list(embedding.shape) if embedding is not None else None
self.cached_embedding = None self.cached_embedding = None
self.error_msg = error_msg 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) # Store additional metadata (e.g., video_timestamps for qwen3_vl)
for key, value in kwargs.items(): for key, value in kwargs.items():
setattr(self, key, value) setattr(self, key, value)
@@ -839,6 +841,7 @@ 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
try:
recv_obj: EmbeddingData = safe_pickle_loads(parts[0]) recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None: if getattr(recv_obj, "error_msg", None) is not None:
logger.warning( logger.warning(
@@ -869,12 +872,31 @@ class WaitingImageRequest:
) )
if self.recv_embedding_data is None: if self.recv_embedding_data is None:
self.recv_embedding_data = MultiModalEmbeddingData.from_embedding_data( self.recv_embedding_data = (
MultiModalEmbeddingData.from_embedding_data(
recv_obj, model_type=self.model_type recv_obj, model_type=self.model_type
) )
)
else: else:
self.recv_embedding_data.add(recv_obj) 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
# 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) recv_embedding = self.recv_embedding_data.get_embedding(is_concat=True)
mm_inputs = self.mm_processor.get_mm_data( mm_inputs = self.mm_processor.get_mm_data(
self.recv_req.input_text, self.recv_req.input_text,
@@ -884,6 +906,13 @@ class WaitingImageRequest:
self.recv_req.mm_inputs = mm_inputs self.recv_req.mm_inputs = mm_inputs
self.recv_req.input_ids = array("q", mm_inputs.input_ids) self.recv_req.input_ids = array("q", mm_inputs.input_ids)
self.status = WaitingImageRequestStatus.SUCCESS 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() self.recv_socket.close()
def _cleanup_gpu_buffer(self): def _cleanup_gpu_buffer(self):
@@ -1177,6 +1206,10 @@ class WaitingImageRDMARequest(WaitingImageRequest):
parts = self.recv_socket.recv_multipart(flags=zmq.NOBLOCK, copy=False) parts = self.recv_socket.recv_multipart(flags=zmq.NOBLOCK, copy=False)
except zmq.Again: except zmq.Again:
return 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]) recv_obj: EmbeddingData = safe_pickle_loads(parts[0])
if getattr(recv_obj, "error_msg", None) is not None: if getattr(recv_obj, "error_msg", None) is not None:
@@ -1803,6 +1836,31 @@ class MMReceiverBase(ABC):
) )
obj.need_wait_for_mm_inputs = False 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 # For zmq_to_scheduler
def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs): def _process_waiting_requests(self, recv_reqs, waiting_cls, **extra_kwargs):
new_recv_reqs = [] new_recv_reqs = []
@@ -1861,6 +1919,7 @@ class MMReceiverBase(ABC):
if status_value == WaitingImageRequestStatus.SUCCESS: if status_value == WaitingImageRequestStatus.SUCCESS:
new_recv_reqs.append(waiting_req.recv_req) new_recv_reqs.append(waiting_req.recv_req)
elif status_value == WaitingImageRequestStatus.FAIL: elif status_value == WaitingImageRequestStatus.FAIL:
self._sync_fail_info_across_tp(waiting_req)
logger.error( logger.error(
f"Waiting request {waiting_req.rid} failed: {waiting_req.error_msg} {waiting_req.error_code = }" 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) recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
for req, error_msg, error_code in abort_reqs: for req, error_msg, error_code in abort_reqs:
status_code = ( if error_code is None:
HTTPStatus.BAD_REQUEST status_code = HTTPStatus.INTERNAL_SERVER_ERROR
if error_code == 400 elif isinstance(error_code, HTTPStatus):
else HTTPStatus.INTERNAL_SERVER_ERROR status_code = error_code
) else:
status_code = HTTPStatus(int(error_code))
prepare_abort(req, error_msg, status_code=status_code) prepare_abort(req, error_msg, status_code=status_code)
self.stream_output([req], req.return_logprob) self.stream_output([req], req.return_logprob)
return recv_reqs return recv_reqs