mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-11 04:49:25 -04:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6cece98ec5 | ||
|
|
3249d41098 |
No files matched your search
+2
-4
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"),
|
||||
|
||||
Reference in new issue
Block a user