[object Object]

← back to Exo

Initialise _cancelled_tasks in ImageEngine (#2051)

dbcceaa50c5544019a6ad2ba21a5579db1b010a8 · 2026-05-05 17:27:57 +0100 · Evan Quiney

we yielded nonsense chunks from engines; we didn't initialize the image
engine correctly. mostly rewrite of #2049

---------

Co-authored-by: ciaranbor <ciaranborourke-dev@proton.me>

Files touched

Diff

commit dbcceaa50c5544019a6ad2ba21a5579db1b010a8
Author: Evan Quiney <evanev7@gmail.com>
Date:   Tue May 5 17:27:57 2026 +0100

    Initialise _cancelled_tasks in ImageEngine (#2051)
    
    we yielded nonsense chunks from engines; we didn't initialize the image
    engine correctly. mostly rewrite of #2049
    
    ---------
    
    Co-authored-by: ciaranbor <ciaranborourke-dev@proton.me>
---
 src/exo/worker/engines/image/builder.py                |  7 ++++++-
 src/exo/worker/runner/llm_inference/batch_generator.py | 18 ++++++++++++++----
 src/exo/worker/runner/runner.py                        |  4 ++--
 3 files changed, 22 insertions(+), 7 deletions(-)

diff --git a/src/exo/worker/engines/image/builder.py b/src/exo/worker/engines/image/builder.py
index 2dd037be..4d20fd88 100644
--- a/src/exo/worker/engines/image/builder.py
+++ b/src/exo/worker/engines/image/builder.py
@@ -143,6 +143,7 @@ class ImageEngine(Engine):
         Generator[tuple[TaskId, Chunk | FinishedResponse | CancelledResponse]] | None
     ) = field(init=False, default=None)
     queue: deque[ImageTask] = field(init=False, default_factory=deque)
+    _cancelled_tasks: set[TaskId] = field(init=False, default_factory=set)
 
     def warmup(self) -> None:
         image = warmup_image_generator(model=self.image_model)
@@ -168,7 +169,11 @@ class ImageEngine(Engine):
             task = self.queue.popleft()
             self.current_gen = self._run_image_task(task.task_id, task.task_params)
             resp = next(self.current_gen, None)
-        return (resp,) if resp is not None else ()
+        return (
+            (resp,)
+            if resp is not None and _is_primary_output_node(self.shard_metadata)
+            else ()
+        )
 
     def close(self) -> None:
         with contextlib.suppress(NameError, AttributeError):
diff --git a/src/exo/worker/runner/llm_inference/batch_generator.py b/src/exo/worker/runner/llm_inference/batch_generator.py
index c049161f..3d898d83 100644
--- a/src/exo/worker/runner/llm_inference/batch_generator.py
+++ b/src/exo/worker/runner/llm_inference/batch_generator.py
@@ -197,9 +197,14 @@ class SequentialGenerator(Engine):
             self._active = None
             raise
 
-        return itertools.chain(
-            output,
-            map(lambda task: (task, CancelledResponse()), self._cancelled_tasks),
+        return filter(
+            lambda chunk: (
+                not isinstance(chunk[1], GenerationChunk) or self.device_rank == 0
+            ),
+            itertools.chain(
+                output,
+                map(lambda task: (task, CancelledResponse()), self._cancelled_tasks),
+            ),
         )
 
     def _start_next(self) -> None:
@@ -449,7 +454,12 @@ class BatchGenerator(Engine):
                 output.append((task.task_id, FinishedResponse()))
                 del self._active_tasks[uid]
 
-        return itertools.chain(output, self._apply_cancellations())
+        return filter(
+            lambda chunk: (
+                not isinstance(chunk[1], GenerationChunk) or self.device_rank == 0
+            ),
+            itertools.chain(output, self._apply_cancellations()),
+        )
 
     def _apply_cancellations(
         self,
diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py
index be6e9a0c..ac5d0548 100644
--- a/src/exo/worker/runner/runner.py
+++ b/src/exo/worker/runner/runner.py
@@ -390,5 +390,5 @@ class Runner:
         chunk: Chunk,
         command_id: CommandId,
     ):
-        if self.device_rank == 0:
-            self.event_sender.send(ChunkGenerated(command_id=command_id, chunk=chunk))
+        assert isinstance(self.generator, Engine)
+        self.event_sender.send(ChunkGenerated(command_id=command_id, chunk=chunk))

← 9c6ff4ce feat: update rdma_ctl instructions (#1977)  ·  back to Exo  ·  fix(inference): prevent TP collective deadlock via agree_on_ 89d20c18 →