Compare commits

...
Author SHA1 Message Date
Alex Cheema 6cece98ec5 Restrict durable events to state changes 2026-05-03 02:31:15 +01:00
Alex Cheema 3249d41098 Route traces over transient events 2026-05-03 02:27:32 +01:00
13 changed files with 103 additions and 50 deletions

No files matched your search

+2 -4
View File
@@ -1939,16 +1939,14 @@ class API:
continue
self._event_log.append(i_event.event)
self.state = apply(self.state, i_event)
event = i_event.event
if isinstance(event, TracesMerged):
self._save_merged_trace(event)
async def _apply_transient(self) -> None:
with self.transient_event_receiver as events:
async for event in events:
if isinstance(event, ChunkGenerated):
await self._dispatch_chunk(event)
elif isinstance(event, TracesMerged):
self._save_merged_trace(event)
async def _dispatch_chunk(self, event: ChunkGenerated) -> None:
if queue := self._image_generation_queues.get(event.command_id, None):
+8 -5
View File
@@ -169,6 +169,7 @@ class Master:
try:
async with self._tg as tg:
tg.start_soon(self._event_processor)
tg.start_soon(self._transient_event_processor)
tg.start_soon(self._command_processor)
tg.start_soon(self._plan)
finally:
@@ -516,10 +517,6 @@ class Master:
local_event.origin,
)
for event in self._multi_buffer.drain():
if isinstance(event, TracesCollected):
await self._handle_traces_collected(event)
continue
logger.debug(f"Master indexing event: {str(event)[:100]}")
event = event.model_copy(
@@ -572,6 +569,12 @@ class Master:
)
)
async def _transient_event_processor(self) -> None:
with self.transient_event_receiver as transients:
async for event in transients:
if isinstance(event, TracesCollected):
await self._handle_traces_collected(event)
# This function is re-entrant, take care!
async def _send_event(self, event: IndexedEvent):
# Convenience method since this line is ugly
@@ -602,7 +605,7 @@ class Master:
for trace_data in self._pending_traces[task_id].values():
all_trace_data.extend(trace_data)
await self.event_sender.send(
await self.transient_event_sender.send(
TracesMerged(task_id=task_id, traces=all_trace_data)
)
+58 -1
View File
@@ -26,6 +26,9 @@ from exo.shared.types.events import (
LocalForwarderEvent,
NodeGatheredInfo,
TaskCreated,
TraceEventData,
TracesCollected,
TracesMerged,
TransientEvent,
)
from exo.shared.types.memory import Memory
@@ -33,7 +36,7 @@ from exo.shared.types.profiling import (
MemoryUsage,
)
from exo.shared.types.snapshots import SnapshotChunk
from exo.shared.types.tasks import TaskStatus
from exo.shared.types.tasks import TaskId, TaskStatus
from exo.shared.types.tasks import TextGeneration as TextGenerationTask
from exo.shared.types.text_generation import (
InputMessage,
@@ -290,3 +293,57 @@ async def test_master_serves_snapshot_for_current_state():
await master.shutdown()
tg.cancel_scope.cancel()
@pytest.mark.asyncio
async def test_master_merges_traces_from_transient_events() -> None:
node_id = NodeId("master")
session_id = SessionId(master_node_id=node_id, election_clock=0)
ge_sender, _global_event_receiver = channel[GlobalForwarderEvent]()
_command_sender, command_receiver = channel[ForwarderCommand]()
_local_event_sender, local_event_receiver = channel[LocalForwarderEvent]()
download_command_sender, _download_command_receiver = channel[
ForwarderDownloadCommand
]()
event_sender, _event_receiver = channel[Event]()
transient_input_sender, transient_input_receiver = channel[TransientEvent]()
transient_output_sender, transient_output_receiver = channel[TransientEvent]()
snapshot_chunk_sender, _snapshot_chunk_receiver = channel[SnapshotChunk]()
master = Master(
node_id,
session_id,
event_sender=event_sender,
transient_event_receiver=transient_input_receiver,
transient_event_sender=transient_output_sender,
global_event_sender=ge_sender,
local_event_receiver=local_event_receiver,
command_receiver=command_receiver,
snapshot_chunk_sender=snapshot_chunk_sender,
download_command_sender=download_command_sender,
)
task_id = TaskId("task-a")
master._expected_ranks[task_id] = {0, 1} # pyright: ignore[reportPrivateUsage]
trace_a = TraceEventData(
name="rank-0", start_us=1, duration_us=2, rank=0, category="test"
)
trace_b = TraceEventData(
name="rank-1", start_us=3, duration_us=4, rank=1, category="test"
)
async with anyio.create_task_group() as tg:
tg.start_soon(master.run)
await transient_input_sender.send(
TracesCollected(task_id=task_id, rank=0, traces=[trace_a])
)
await transient_input_sender.send(
TracesCollected(task_id=task_id, rank=1, traces=[trace_b])
)
merged = await transient_output_receiver.receive()
assert isinstance(merged, TracesMerged)
assert merged.task_id == task_id
assert merged.traces == [trace_a, trace_b]
tg.cancel_scope.cancel()
+1 -11
View File
@@ -7,7 +7,6 @@ from loguru import logger
from exo.shared.models.model_cards import ModelCard
from exo.shared.types.common import ModelId, NodeId
from exo.shared.types.events import (
ChunkGenerated,
CustomModelCardAdded,
CustomModelCardDeleted,
Event,
@@ -21,7 +20,6 @@ from exo.shared.types.events import (
NodeGatheredInfo,
NodeTimedOut,
RunnerStatusUpdated,
TaskAcknowledged,
TaskCreated,
TaskDeleted,
TaskFailed,
@@ -29,8 +27,6 @@ from exo.shared.types.events import (
TestEvent,
TopologyEdgeCreated,
TopologyEdgeDeleted,
TracesCollected,
TracesMerged,
)
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.profiling import (
@@ -76,13 +72,7 @@ from exo.utils.info_gatherer.info_gatherer import (
def event_apply(event: Event, state: State) -> State:
"""Apply an event to state."""
match event:
case (
TestEvent()
| ChunkGenerated()
| TaskAcknowledged()
| TracesCollected()
| TracesMerged()
): # Pass-through events that don't modify state
case TestEvent():
return state
case CustomModelCardAdded():
return apply_custom_model_card_added(event, state)
+1 -4
View File
@@ -152,19 +152,15 @@ Event = (
| TaskStatusUpdated
| TaskFailed
| TaskDeleted
| TaskAcknowledged
| InstanceCreated
| InstanceDeleted
| RunnerStatusUpdated
| NodeTimedOut
| NodeGatheredInfo
| NodeDownloadProgress
| ChunkGenerated
| InputChunkReceived
| TopologyEdgeCreated
| TopologyEdgeDeleted
| TracesCollected
| TracesMerged
| CustomModelCardAdded
| CustomModelCardDeleted
| InstanceLinkCreated
@@ -173,6 +169,7 @@ Event = (
TransientEvent = TaskAcknowledged | ChunkGenerated | TracesCollected | TracesMerged
RunnerEvent = Event | TransientEvent
class IndexedEvent(FrozenModel):
+4 -4
View File
@@ -12,7 +12,7 @@ from exo.shared.constants import EXO_TRACING_ENABLED
from exo.shared.tracing import clear_trace_buffer, get_trace_buffer
from exo.shared.types.chunks import Chunk, ErrorChunk
from exo.shared.types.events import (
Event,
RunnerEvent,
TraceEventData,
TracesCollected,
)
@@ -66,7 +66,7 @@ def _is_primary_output_node(shard_metadata: ShardMetadata) -> bool:
def _send_traces_if_enabled(
event_sender: MpSender[Event],
event_sender: MpSender[RunnerEvent],
task_id: TaskId,
rank: int,
) -> None:
@@ -97,7 +97,7 @@ def _send_traces_if_enabled(
@dataclass
class MfluxBuilder(Builder):
event_sender: MpSender[Event]
event_sender: MpSender[RunnerEvent]
cancel_receiver: MpReceiver[TaskId]
shard_metadata: ShardMetadata | None = None
image_model: DistributedImageModel | None = None
@@ -137,7 +137,7 @@ class MfluxBuilder(Builder):
class ImageEngine(Engine):
image_model: DistributedImageModel
shard_metadata: ShardMetadata
event_sender: MpSender[Event]
event_sender: MpSender[RunnerEvent]
cancel_receiver: MpReceiver[TaskId]
current_gen: (
Generator[tuple[TaskId, Chunk | FinishedResponse | CancelledResponse]] | None
+2 -2
View File
@@ -7,7 +7,7 @@ import mlx.core as mx
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.types.common import ModelId
from exo.shared.types.events import Event
from exo.shared.types.events import RunnerEvent
from exo.shared.types.tasks import TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import ModelLoadingResponse
@@ -32,7 +32,7 @@ from .vision import VisionProcessor
@dataclass
class MlxBuilder(Builder):
model_id: ModelId
event_sender: MpSender[Event]
event_sender: MpSender[RunnerEvent]
cancel_receiver: MpReceiver[TaskId]
inference_model: Model | None = None
tokenizer: TokenizerWrapper | None = None
+2 -2
View File
@@ -3,7 +3,7 @@ import resource
import loguru
from exo.shared.types.events import Event, RunnerStatusUpdated
from exo.shared.types.events import RunnerEvent, RunnerStatusUpdated
from exo.shared.types.tasks import Task, TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runners import RunnerFailed
@@ -15,7 +15,7 @@ logger: "loguru.Logger" = loguru.logger
def entrypoint(
bound_instance: BoundInstance,
event_sender: MpSender[Event],
event_sender: MpSender[RunnerEvent],
task_receiver: MpReceiver[Task],
cancel_receiver: MpReceiver[TaskId],
_logger: "loguru.Logger",
@@ -11,7 +11,7 @@ from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.constants import EXO_MAX_CONCURRENT_REQUESTS
from exo.shared.types.chunks import ErrorChunk, GenerationChunk, PrefillProgressChunk
from exo.shared.types.common import ModelId
from exo.shared.types.events import ChunkGenerated, Event
from exo.shared.types.events import ChunkGenerated, RunnerEvent
from exo.shared.types.tasks import (
CANCEL_ALL_TASKS,
GenerationTask,
@@ -96,7 +96,7 @@ class SequentialGenerator(Engine):
model_id: ModelId
device_rank: int
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
event_sender: MpSender[RunnerEvent]
vision_processor: VisionProcessor | None = None
check_for_cancel_every: int = 50
@@ -320,7 +320,7 @@ class BatchGenerator(Engine):
model_id: ModelId
device_rank: int
cancel_receiver: MpReceiver[TaskId]
event_sender: MpSender[Event]
event_sender: MpSender[RunnerEvent]
check_for_cancel_every: int = 50
vision_processor: VisionProcessor | None = None
+2 -2
View File
@@ -12,7 +12,7 @@ from exo.shared.types.chunks import Chunk
from exo.shared.types.common import CommandId
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerEvent,
RunnerStatusUpdated,
TaskAcknowledged,
TaskStatusUpdated,
@@ -86,7 +86,7 @@ class Runner:
self,
bound_instance: BoundInstance,
builder: Builder,
event_sender: MpSender[Event],
event_sender: MpSender[RunnerEvent],
task_receiver: MpReceiver[Task],
):
self.event_sender = event_sender
+8 -3
View File
@@ -16,9 +16,12 @@ from exo.shared.types.chunks import ErrorChunk
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerEvent,
RunnerStatusUpdated,
TaskAcknowledged,
TaskStatusUpdated,
TracesCollected,
TracesMerged,
TransientEvent,
)
from exo.shared.types.tasks import (
@@ -56,7 +59,7 @@ class RunnerSupervisor:
bound_instance: BoundInstance
runner_process: mp.Process
initialize_timeout: float
_ev_recv: MpReceiver[Event]
_ev_recv: MpReceiver[RunnerEvent]
_task_sender: MpSender[Task]
_event_sender: Sender[Event]
_transient_event_sender: Sender[TransientEvent]
@@ -80,7 +83,7 @@ class RunnerSupervisor:
transient_event_sender: Sender[TransientEvent],
initialize_timeout: float = 400,
) -> Self:
ev_send, ev_recv = mp_channel[Event]()
ev_send, ev_recv = mp_channel[RunnerEvent]()
task_sender, task_recv = mp_channel[Task]()
cancel_sender, cancel_recv = mp_channel[TaskId]()
@@ -241,7 +244,9 @@ class RunnerSupervisor:
)
self.in_progress.pop(event.task_id, None)
self.completed.add(event.task_id)
if isinstance(event, ChunkGenerated):
if isinstance(
event, (ChunkGenerated, TracesCollected, TracesMerged)
):
await self._transient_event_sender.send(event)
else:
await self._event_sender.send(event)
@@ -11,7 +11,7 @@ import exo.worker.runner.llm_inference.model_output_parsers as mlx_model_output_
from exo.shared.types.chunks import TokenChunk
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerEvent,
RunnerStatusUpdated,
TaskAcknowledged,
TaskStatusUpdated,
@@ -110,7 +110,9 @@ CHAT_TASK = TextGeneration(
)
def assert_events_equal(test_events: Iterable[Event], true_events: Iterable[Event]):
def assert_events_equal(
test_events: Iterable[RunnerEvent], true_events: Iterable[RunnerEvent]
):
for test_event, true_event in zip(test_events, true_events, strict=True):
test_event = test_event.model_copy(update={"event_id": true_event.event_id})
assert test_event == true_event, f"{test_event} != {true_event}"
@@ -200,11 +202,11 @@ class FakeExoBatchGenerator:
# Use a fake event_sender to remove test flakiness.
class EventCollector:
def __init__(self, on_event: Callable[[Event], None] | None = None) -> None:
self.events: list[Event] = []
def __init__(self, on_event: Callable[[RunnerEvent], None] | None = None) -> None:
self.events: list[RunnerEvent] = []
self._on_event = on_event
def send(self, event: Event) -> None:
def send(self, event: RunnerEvent) -> None:
self.events.append(event)
if self._on_event:
self._on_event(event)
@@ -254,11 +256,11 @@ def _run(tasks: Iterable[Task], send_after_ready: list[Task] | None = None):
task_sender, task_receiver = mp_channel[Task]()
_cancel_sender, cancel_receiver = mp_channel[TaskId]()
on_event: Callable[[Event], None] | None = None
on_event: Callable[[RunnerEvent], None] | None = None
if send_after_ready:
_saw_running = False
def _on_event(event: Event) -> None:
def _on_event(event: RunnerEvent) -> None:
nonlocal _saw_running
if isinstance(event, RunnerStatusUpdated):
if isinstance(event.runner_status, RunnerRunning):
@@ -10,6 +10,7 @@ from exo.shared.types.common import CommandId, NodeId
from exo.shared.types.events import (
ChunkGenerated,
Event,
RunnerEvent,
RunnerStatusUpdated,
TransientEvent,
)
@@ -51,7 +52,7 @@ async def test_check_runner_emits_error_chunk_for_inflight_text_generation() ->
transient_sender, transient_receiver = channel[TransientEvent]()
task_sender, _ = mp_channel[Task]()
cancel_sender, _ = mp_channel[TaskId]()
_, ev_recv = mp_channel[Event]()
_, ev_recv = mp_channel[RunnerEvent]()
bound_instance: BoundInstance = get_bound_mlx_ring_instance(
instance_id=InstanceId("instance-a"),
@@ -113,7 +114,7 @@ async def test_forward_events_routes_generated_chunks_to_transient() -> None:
transient_sender, transient_receiver = channel[TransientEvent]()
task_sender, _ = mp_channel[Task]()
cancel_sender, _ = mp_channel[TaskId]()
ev_send, ev_recv = mp_channel[Event]()
ev_send, ev_recv = mp_channel[RunnerEvent]()
bound_instance: BoundInstance = get_bound_mlx_ring_instance(
instance_id=InstanceId("instance-a"),