Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions docs/design/feature/async_diffusion_output.md
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,8 @@ After (async D2H):
- `OUTPUT_READY` → resolves `_output_futures[async_output_id]` (with batch split via `_batch_split_map`)
- Non-async messages → `_sync_result_buffer` for other RPCs

Because this thread is the queue's only reader, dispatch is fault contained: each message is routed inside a `try`, so a failure costs that message rather than every later result. An `OUTPUT_READY` without an `async_output_id` cannot be routed to anyone, so it is logged and dropped instead of disappearing silently, and batch splitting resolves each request separately so one corrupt per-request result does not leave its siblings hanging until their own timeout.

4. **collective_rpc() two-path dispatch** (`multiproc_executor.py`):
- **Path 1**: `execute_model` / `execute_model_batch` → generates `rpc_id`, registers Future, waits for pump to deliver `compute_done`
- **Path 2**: All other RPCs → `_sync_result_buffer` (pump-fed) or `_result_mq` (step-mode, no pump)
Expand Down
98 changes: 98 additions & 0 deletions tests/diffusion/test_result_pump.py
Original file line number Diff line number Diff line change
Expand Up @@ -938,3 +938,101 @@ def test_shutdown_clears_completed_outputs(self):
executor.shutdown()

assert executor._completed_outputs == {}


def _feed_msgs_to_pump(executor, msgs):
"""Run _result_pump in a daemon thread, feed *msgs* in order, then stop."""
pending = list(msgs)

def mock_dequeue(timeout=None):
if pending:
return pending.pop(0)
executor._pump_stop.set()
time.sleep(0.05)
raise TimeoutError

executor._result_mq.dequeue = mock_dequeue
t = threading.Thread(target=executor._result_pump, daemon=True)
t.start()
t.join(timeout=2.0)
return t


class _ExplodingBatchOutput:
"""Batch output whose extraction fails for one request id."""

def __init__(self, results, failing_req_id):
self._results = results
self._failing_req_id = failing_req_id

def get_request_output(self, req_id):
if req_id == self._failing_req_id:
raise RuntimeError("corrupt per-request result")
result = self._results.get(req_id)
if result is None:
return None
return SimpleNamespace(result=result)


class TestResultPumpFaultContainment:
"""One bad message must not cost the queue its only reader."""

def test_dispatch_failure_does_not_stop_the_pump(self, mocker):
executor = _make_executor()
mocker.patch.object(
executor._sync_result_buffer,
"put",
side_effect=RuntimeError("sync buffer is gone"),
)
fut = concurrent.futures.Future()
with executor._futures_lock:
executor._rpc_futures["1"] = fut
good = AsyncDiffusionOutput(kind=AsyncOutputKind.RPC_RESULT, rpc_id="1")

thread = _feed_msgs_to_pump(executor, [object(), good])

# The message after the failure was still delivered, so the pump lived.
assert fut.done()
assert fut.result(timeout=1.0) is good
assert not thread.is_alive()

def test_output_ready_without_id_is_logged_not_silent(self, mocker):
from vllm_omni.diffusion.executor import multiproc_executor

executor = _make_executor()
log_error = mocker.patch.object(multiproc_executor.logger, "error")
orphan = AsyncDiffusionOutput(
kind=AsyncOutputKind.OUTPUT_READY,
async_output_id=None,
output=DiffusionOutput(output="img"),
)

_feed_msgs_to_pump(executor, [orphan])

assert executor._completed_outputs == {}
assert executor._output_futures == {}
assert log_error.call_count == 1
assert "async_output_id" in log_error.call_args.args[0]

def test_one_corrupt_request_does_not_drop_its_siblings(self):
executor = _make_executor()
req_ids = ["r0", "r1", "r2"]
batch_id = "batch-fault"
outputs = {rid: DiffusionOutput(output=f"img-{rid}") for rid in req_ids}
with executor._futures_lock:
executor._batch_split_map[batch_id] = {f"{batch_id}/{rid}": rid for rid in req_ids}
ready = AsyncDiffusionOutput(
kind=AsyncOutputKind.OUTPUT_READY,
async_output_id=batch_id,
output=_ExplodingBatchOutput(outputs, failing_req_id="r1"),
)

_feed_msgs_to_pump(executor, [ready])

for rid in ("r0", "r2"):
fut = executor.wait_output_ready(f"{batch_id}/{rid}")
assert fut.done(), f"request {rid} never resolved"
assert fut.result(timeout=1.0) is outputs[rid]
failed = executor.wait_output_ready(f"{batch_id}/r1")
assert failed.done()
assert "corrupt per-request result" in failed.result(timeout=1.0).error
143 changes: 84 additions & 59 deletions vllm_omni/diffusion/executor/multiproc_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -938,62 +938,79 @@ def _result_pump(self, result_mq: MessageQueue | None = None) -> None:
break
continue

if not isinstance(msg, AsyncDiffusionOutput):
# Non-async message: place into the sync buffer for
# collective_rpc() to consume via Path 2.
self._sync_result_buffer.put(msg)
continue
# This thread is the sole reader of result_mq, so an exception
# escaping the dispatch costs the whole queue rather than one
# message: every later result goes undelivered while the server
# keeps reporting healthy. Contain it to the message that caused it.
try:
self._dispatch_result(msg)
except Exception:
logger.exception("Result pump failed to dispatch a message; dropping it")

def _dispatch_result(self, msg: Any) -> None:
"""Route one dequeued message to whoever is waiting for it."""
if not isinstance(msg, AsyncDiffusionOutput):
# Non-async message: place into the sync buffer for
# collective_rpc() to consume via Path 2.
self._sync_result_buffer.put(msg)
return

# If shutdown started while we were dequeuing, drop this delivery
# so it cannot repopulate _completed_outputs after shutdown()
# cleared it (issue #6413 / #6439 review). OUTPUT_READY must still
# flow through the dispatch below: unpack_diffusion_output_shm()
# is the only receive-side path that unlinks named SHM segments,
# and the closed-time re-checks under _futures_lock already
# prevent any cache write after unpack.
if self._closed and msg.kind != AsyncOutputKind.OUTPUT_READY:
continue
# If shutdown started while we were dequeuing, drop this delivery
# so it cannot repopulate _completed_outputs after shutdown()
# cleared it (issue #6413 / #6439 review). OUTPUT_READY must still
# flow through the dispatch below: unpack_diffusion_output_shm()
# is the only receive-side path that unlinks named SHM segments,
# and the closed-time re-checks under _futures_lock already
# prevent any cache write after unpack.
if self._closed and msg.kind != AsyncOutputKind.OUTPUT_READY:
return

if msg.kind in (AsyncOutputKind.RPC_RESULT, AsyncOutputKind.COMPUTE_DONE):
with self._futures_lock:
fut = self._rpc_futures.pop(msg.rpc_id, None) if msg.rpc_id else None
if fut is not None and not fut.done():
if msg.error:
try_set_exception(fut, RuntimeError(msg.error))
else:
try_set_result(fut, msg)
elif msg.kind == AsyncOutputKind.OUTPUT_READY:
batch_id = msg.async_output_id
with self._futures_lock:
per_req_map = self._batch_split_map.pop(batch_id, None) if batch_id else None
if per_req_map is not None:
# Batch result: split into per-request DiffusionOutputs.
if msg.kind in (AsyncOutputKind.RPC_RESULT, AsyncOutputKind.COMPUTE_DONE):
with self._futures_lock:
fut = self._rpc_futures.pop(msg.rpc_id, None) if msg.rpc_id else None
if fut is not None and not fut.done():
if msg.error:
try_set_exception(fut, RuntimeError(msg.error))
else:
try_set_result(fut, msg)
elif msg.kind == AsyncOutputKind.OUTPUT_READY:
batch_id = msg.async_output_id
if not batch_id:
# async_output_id is what routes a delivery back to its waiter.
# Without it there is nobody to resolve and nothing to cache, so
# the request hangs until its own timeout. Say so rather than
# dropping the message silently.
logger.error("Dropping OUTPUT_READY with no async_output_id; its request cannot be resolved")
return
with self._futures_lock:
per_req_map = self._batch_split_map.pop(batch_id, None)
if per_req_map is not None:
# Batch result: split into per-request DiffusionOutputs.
try:
unpack_diffusion_output_shm(msg.output)
except Exception:
logger.exception("SHM unpack failed for batch %s", batch_id)
self._deliver_batch_split(per_req_map, msg.output, msg.error)
else:
# Single-request result: unpack SHM first, then resolve or cache atomically.
output_result: DiffusionOutput | None = None
exc: Exception | None = None
if msg.error:
exc = RuntimeError(msg.error)
else:
try:
unpack_diffusion_output_shm(msg.output)
except Exception:
logger.exception("SHM unpack failed for batch %s", batch_id)
self._deliver_batch_split(per_req_map, msg.output, msg.error)
else:
# Single-request result: unpack SHM first, then resolve or cache atomically.
output_result: DiffusionOutput | None = None
exc: Exception | None = None
if msg.error:
exc = RuntimeError(msg.error)
else:
try:
unpack_diffusion_output_shm(msg.output)
output_result = msg.output
except Exception as e:
logger.exception("SHM unpack failed in result pump")
exc = e

if batch_id:
with self._futures_lock:
if self._closed:
# shutdown() cleared _completed_outputs while
# we were unpacking; drop this delivery.
continue
self._finish_output(batch_id, output_result, exc)
output_result = msg.output
except Exception as e:
logger.exception("SHM unpack failed in result pump")
exc = e

with self._futures_lock:
if self._closed:
# shutdown() cleared _completed_outputs while
# we were unpacking; drop this delivery.
return
self._finish_output(batch_id, output_result, exc)

def _finish_output(
self,
Expand Down Expand Up @@ -1077,14 +1094,22 @@ def _deliver_batch_split(
# shutdown() has taken over; do not touch _completed_outputs.
return
for per_req_id, req_id in per_req_map.items():
req_output = batch_output.get_request_output(req_id) if batch_output is not None else None
per_req_result: DiffusionOutput
if req_output is not None and req_output.result is not None:
per_req_result = req_output.result
elif error:
per_req_result = DiffusionOutput(error=error)
else:
per_req_result = DiffusionOutput(error="No output result for batch request")
try:
req_output = batch_output.get_request_output(req_id) if batch_output is not None else None
if req_output is not None and req_output.result is not None:
per_req_result = req_output.result
elif error:
per_req_result = DiffusionOutput(error=error)
else:
per_req_result = DiffusionOutput(error="No output result for batch request")
except Exception as e:
# The split map has already been popped, so an exception escaping
# here would take every request later in this batch with it: they
# would never be resolved and would hang until their own timeout.
# One corrupt request costs one request.
logger.exception("Failed to extract batch output for request %s", req_id)
per_req_result = DiffusionOutput(error=f"Failed to extract batch output for request {req_id}: {e}")
with self._futures_lock:
if self._closed:
# Belt-and-braces: re-check under the lock so a shutdown
Expand Down
Loading