mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-08 19:41:32 -04:00
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8a28664846 | ||
|
|
9473fd2652 | ||
|
|
edef8004f8 | ||
|
|
a0c00f9dfd | ||
|
|
89d20c1888 | ||
|
|
dbcceaa50c |
No files matched your search
@@ -36,6 +36,7 @@ jobs:
|
||||
with:
|
||||
name: exo
|
||||
authToken: "${{ secrets.CACHIX_AUTH_TOKEN }}"
|
||||
pushFilter: '-docker-image.tar.gz$'
|
||||
|
||||
- name: Build Metal packages (macOS only)
|
||||
if: runner.os == 'macOS'
|
||||
|
||||
+26
-7
@@ -183,7 +183,24 @@ let
|
||||
inherit venv;
|
||||
editablePythonSet = pythonSet.overrideScope editableOverlay;
|
||||
mkPythonScript = path: mkApp ''python ${path} "$@"'';
|
||||
mkExo = mkApp ''exo "$@"'';
|
||||
mkOutputs = name:
|
||||
let package = mkApp ''exo "$@"'' name;
|
||||
in {
|
||||
${name} = package;
|
||||
"${name}-docker-image" = pkgs.dockerTools.buildLayeredImage {
|
||||
name = "${name}-docker-image";
|
||||
config = {
|
||||
Entrypoint = [ (lib.getExe package) ];
|
||||
Env = [
|
||||
"EXO_HOME=/var/lib/${name}"
|
||||
];
|
||||
};
|
||||
|
||||
extraCommands = ''
|
||||
mkdir -p var/lib/${name}
|
||||
'';
|
||||
};
|
||||
};
|
||||
};
|
||||
in
|
||||
{
|
||||
@@ -191,7 +208,7 @@ in
|
||||
{ self', pkgs, unfreePkgs, lib, ... }:
|
||||
let
|
||||
inherit (pkgs.stdenv.hostPlatform) isLinux;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "cpu" ]; }; }) editablePythonSet mkExo;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; members.exo = [ "cpu" ]; }) editablePythonSet mkOutputs;
|
||||
|
||||
# Virtual environment with dev dependencies for testing
|
||||
testVenv = (mkPythonSet {
|
||||
@@ -213,10 +230,13 @@ in
|
||||
text = ''exec python ${path} "$@"'';
|
||||
};
|
||||
|
||||
defaultOutputs = mkOutputs "exo";
|
||||
cuda12Outputs = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; members.exo=["cuda12"]; }).mkOutputs "exo-cuda-12";
|
||||
cuda13Outputs = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; members.exo=["cuda13"]; }).mkOutputs "exo-cuda-13";
|
||||
in
|
||||
{
|
||||
packages = {
|
||||
exo = mkExo "exo";
|
||||
# for devShell
|
||||
editableVenv = editablePythonSet.mkVirtualEnv "exo-dev-env" { exo = [ "dev" ]; };
|
||||
# for running tests in ci
|
||||
exo-test-env = testVenv;
|
||||
@@ -225,10 +245,9 @@ in
|
||||
exo-eval-tool-calls = mkBenchScript "exo-eval-tool-calls" (inputs.self + /bench/eval_tool_calls.py);
|
||||
# used by ./tests/run_exo_on.sh
|
||||
exo-get-all-models-on-cluster = mkSimplePythonScript "exo-get-all-models-on-cluster" (inputs.self + /tests/get_all_models_on_cluster.py);
|
||||
} // lib.optionalAttrs isLinux {
|
||||
exo-cuda-12 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; members = { exo = [ "cuda12" ]; }; }).mkExo "exo-cuda-12";
|
||||
exo-cuda-13 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; members = { exo = [ "cuda13" ]; }; }).mkExo "exo-cuda-13";
|
||||
};
|
||||
} // defaultOutputs
|
||||
// lib.optionalAttrs isLinux cuda12Outputs
|
||||
// lib.optionalAttrs isLinux cuda13Outputs;
|
||||
|
||||
checks = {
|
||||
lint = pkgs.runCommand "ruff-lint" { } ''
|
||||
|
||||
+59
-144
@@ -121,8 +121,6 @@ from exo.api.types.openai_responses import (
|
||||
)
|
||||
from exo.master.image_store import ImageStore
|
||||
from exo.master.placement import place_instance as get_instance_placements
|
||||
from exo.routing.event_router import EventRouter
|
||||
from exo.routing.snapshot_receiver import SnapshotReceiver
|
||||
from exo.shared.apply import apply
|
||||
from exo.shared.constants import (
|
||||
DASHBOARD_DIR,
|
||||
@@ -135,12 +133,10 @@ from exo.shared.constants import (
|
||||
)
|
||||
from exo.shared.election import ElectionMessage
|
||||
from exo.shared.logging import InterceptLogger
|
||||
from exo.shared.models import model_cards
|
||||
from exo.shared.models.model_cards import (
|
||||
ModelCard,
|
||||
ModelId,
|
||||
add_to_card_cache,
|
||||
get_card,
|
||||
get_model_cards,
|
||||
)
|
||||
from exo.shared.tracing import TraceEvent, compute_stats, export_trace, load_trace_file
|
||||
from exo.shared.types.chunks import (
|
||||
@@ -166,7 +162,6 @@ from exo.shared.types.commands import (
|
||||
ImageEdits,
|
||||
ImageGeneration,
|
||||
PlaceInstance,
|
||||
RequestSnapshot,
|
||||
SendInputChunk,
|
||||
SetInstanceLink,
|
||||
StartDownload,
|
||||
@@ -174,17 +169,16 @@ from exo.shared.types.commands import (
|
||||
TaskFinished,
|
||||
TextGeneration,
|
||||
)
|
||||
from exo.shared.types.common import CommandId, Id, NodeId, SessionId, SystemId
|
||||
from exo.shared.types.common import CommandId, Id, NodeId, SystemId
|
||||
from exo.shared.types.events import (
|
||||
ChunkGenerated,
|
||||
Event,
|
||||
IndexedEvent,
|
||||
InstanceDeleted,
|
||||
TracesMerged,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.snapshots import SnapshotChunk
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.tasks import (
|
||||
ImageEdits as ImageEditsTask,
|
||||
@@ -211,8 +205,6 @@ from exo.utils.task_group import TaskGroup
|
||||
_API_EVENT_LOG_DIR = EXO_EVENT_LOG_DIR / "api"
|
||||
ONBOARDING_COMPLETE_FILE = EXO_CACHE_HOME / "onboarding_complete"
|
||||
|
||||
_SNAPSHOT_FETCH_TIMEOUT_SECONDS = 30
|
||||
|
||||
|
||||
def _format_to_content_type(image_format: Literal["png", "jpeg", "webp"] | None) -> str:
|
||||
return f"image/{image_format or 'png'}"
|
||||
@@ -242,13 +234,9 @@ class API:
|
||||
def __init__(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
session_id: SessionId,
|
||||
*,
|
||||
port: int,
|
||||
event_router: EventRouter,
|
||||
event_receiver: Receiver[IndexedEvent],
|
||||
transient_event_receiver: Receiver[TransientEvent],
|
||||
snapshot_chunk_receiver: Receiver[SnapshotChunk],
|
||||
command_sender: Sender[ForwarderCommand],
|
||||
download_command_sender: Sender[ForwarderDownloadCommand],
|
||||
# This lets us pause the API if an election is running
|
||||
@@ -257,17 +245,14 @@ class API:
|
||||
self.state = State()
|
||||
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
|
||||
self._system_id = SystemId()
|
||||
self.session_id = session_id
|
||||
self.event_router = event_router
|
||||
self.command_sender = command_sender
|
||||
self.download_command_sender = download_command_sender
|
||||
self.event_receiver = event_receiver
|
||||
self.transient_event_receiver = transient_event_receiver
|
||||
self.snapshot_chunk_receiver = snapshot_chunk_receiver
|
||||
self.election_receiver = election_receiver
|
||||
self.node_id: NodeId = node_id
|
||||
self.last_completed_election: int = 0
|
||||
self.port = port
|
||||
self._sent_image_hashes: set[str] = set()
|
||||
|
||||
self.paused: bool = False
|
||||
self.paused_ev: anyio.Event = anyio.Event()
|
||||
@@ -302,38 +287,22 @@ class API:
|
||||
self._image_generation_queues: dict[
|
||||
CommandId, Sender[ImageChunk | ErrorChunk]
|
||||
] = {}
|
||||
self._observed_generation_commands: set[CommandId] = set()
|
||||
self._image_store = ImageStore(EXO_IMAGE_CACHE_DIR)
|
||||
self._tg: TaskGroup = TaskGroup()
|
||||
|
||||
def reset(
|
||||
self,
|
||||
result_clock: int,
|
||||
session_id: SessionId,
|
||||
event_router: EventRouter,
|
||||
event_receiver: Receiver[IndexedEvent],
|
||||
transient_event_receiver: Receiver[TransientEvent],
|
||||
snapshot_chunk_receiver: Receiver[SnapshotChunk],
|
||||
):
|
||||
def reset(self, result_clock: int, event_receiver: Receiver[IndexedEvent]):
|
||||
logger.info("Resetting API State")
|
||||
self._event_log.close()
|
||||
self._event_log = DiskEventLog(_API_EVENT_LOG_DIR)
|
||||
self.state = State()
|
||||
self._system_id = SystemId()
|
||||
self.session_id = session_id
|
||||
self.event_router = event_router
|
||||
self._text_generation_queues = {}
|
||||
self._image_generation_queues = {}
|
||||
self._observed_generation_commands = set()
|
||||
self.unpause(result_clock)
|
||||
self.event_receiver.close()
|
||||
self.event_receiver = event_receiver
|
||||
self.transient_event_receiver.close()
|
||||
self.transient_event_receiver = transient_event_receiver
|
||||
self.snapshot_chunk_receiver.close()
|
||||
self.snapshot_chunk_receiver = snapshot_chunk_receiver
|
||||
self._tg.start_soon(self._bootstrap_then_apply_state)
|
||||
self._tg.start_soon(self._apply_transient)
|
||||
self._tg.start_soon(self._apply_state)
|
||||
self._sent_image_hashes = set()
|
||||
|
||||
def unpause(self, result_clock: int):
|
||||
logger.info("Unpausing API")
|
||||
@@ -510,6 +479,7 @@ class API:
|
||||
topology=self.state.topology,
|
||||
current_instances=self.state.instances,
|
||||
download_status=self.state.downloads,
|
||||
node_rdma_ctl=self.state.node_rdma_ctl,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
@@ -573,6 +543,7 @@ class API:
|
||||
current_instances=self.state.instances,
|
||||
required_nodes=required_nodes,
|
||||
download_status=self.state.downloads,
|
||||
node_rdma_ctl=self.state.node_rdma_ctl,
|
||||
)
|
||||
except ValueError as exc:
|
||||
if (model_card.model_id, sharding, instance_meta, 0) not in seen:
|
||||
@@ -855,8 +826,18 @@ class API:
|
||||
)
|
||||
command = TextGeneration(task_params=task_params)
|
||||
|
||||
new_images: list[tuple[int, str]] = []
|
||||
for idx, (img, h) in enumerate(zip(images, hashes, strict=True)):
|
||||
if h not in self._sent_image_hashes:
|
||||
self._sent_image_hashes.add(h)
|
||||
new_images.append((idx, img))
|
||||
|
||||
if not new_images:
|
||||
await self._send(command)
|
||||
return command
|
||||
|
||||
all_chunks: list[tuple[int, str]] = []
|
||||
for img_idx, img_data in enumerate(images):
|
||||
for img_idx, img_data in new_images:
|
||||
for i in range(0, len(img_data), EXO_MAX_CHUNK_SIZE):
|
||||
all_chunks.append((img_idx, img_data[i : i + EXO_MAX_CHUNK_SIZE]))
|
||||
|
||||
@@ -1652,17 +1633,16 @@ class API:
|
||||
async def ollama_tags(self) -> OllamaTagsResponse:
|
||||
"""Returns list of models in Ollama tags format. We return the downloaded ones only."""
|
||||
|
||||
def none_if_empty(value: str) -> str | None:
|
||||
return value or None
|
||||
|
||||
downloaded_model_ids: set[str] = set()
|
||||
downloaded_model_ids: set[ModelId] = set()
|
||||
for node_downloads in self.state.downloads.values():
|
||||
for dl in node_downloads:
|
||||
if isinstance(dl, DownloadCompleted):
|
||||
downloaded_model_ids.add(dl.shard_metadata.model_card.model_id)
|
||||
|
||||
cards = [
|
||||
c for c in await get_model_cards() if c.model_id in downloaded_model_ids
|
||||
c
|
||||
for c in await model_cards.card_cache.list_all()
|
||||
if c.model_id in downloaded_model_ids
|
||||
]
|
||||
|
||||
now = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
|
||||
@@ -1675,8 +1655,8 @@ class API:
|
||||
size=card.storage_size.in_bytes,
|
||||
digest="sha256:000000000000",
|
||||
details=OllamaModelDetails(
|
||||
family=none_if_empty(card.family),
|
||||
quantization_level=none_if_empty(card.quantization),
|
||||
family=card.family or None,
|
||||
quantization_level=card.quantization or None,
|
||||
),
|
||||
)
|
||||
for card in cards
|
||||
@@ -1739,7 +1719,7 @@ class API:
|
||||
|
||||
async def get_models(self, status: str | None = Query(default=None)) -> ModelList:
|
||||
"""Returns list of available models, optionally filtered by being downloaded."""
|
||||
cards = await get_model_cards()
|
||||
cards = await model_cards.card_cache.list_all()
|
||||
|
||||
if status == "downloaded":
|
||||
downloaded_model_ids: set[str] = set()
|
||||
@@ -1790,7 +1770,7 @@ class API:
|
||||
|
||||
# Immediately update the local cache so the subsequent GET /models
|
||||
# returns the new model without waiting for the event round-trip.
|
||||
add_to_card_cache(card)
|
||||
model_cards.card_cache.cc[card.model_id] = card
|
||||
|
||||
return ModelListModel(
|
||||
id=card.model_id,
|
||||
@@ -1806,7 +1786,7 @@ class API:
|
||||
|
||||
async def delete_custom_model(self, model_id: ModelId) -> JSONResponse:
|
||||
"""Delete a user-added custom model card and sync deletion across the cluster."""
|
||||
card = get_card(model_id)
|
||||
card = model_cards.card_cache.get(model_id)
|
||||
if card is None or not card.is_custom:
|
||||
raise HTTPException(status_code=404, detail="Custom model card not found")
|
||||
|
||||
@@ -1867,9 +1847,7 @@ class API:
|
||||
try:
|
||||
async with self._tg as tg:
|
||||
logger.info("Starting API")
|
||||
tg.start_soon(self._bootstrap_then_apply_state)
|
||||
tg.start_soon(self._apply_transient)
|
||||
tg.start_soon(self._reconcile_streams)
|
||||
tg.start_soon(self._apply_state)
|
||||
tg.start_soon(self._pause_on_new_election)
|
||||
tg.start_soon(self._cleanup_expired_images)
|
||||
print_startup_banner(self.port)
|
||||
@@ -1883,8 +1861,6 @@ class API:
|
||||
self._event_log.close()
|
||||
self.command_sender.close()
|
||||
self.event_receiver.close()
|
||||
self.transient_event_receiver.close()
|
||||
self.snapshot_chunk_receiver.close()
|
||||
|
||||
async def run_api(self, ev: anyio.Event):
|
||||
cfg = Config()
|
||||
@@ -1900,108 +1876,47 @@ class API:
|
||||
shutdown_trigger=ev.wait,
|
||||
)
|
||||
|
||||
async def _bootstrap_then_apply_state(self):
|
||||
await self._fetch_snapshot()
|
||||
await self._apply_state()
|
||||
|
||||
async def _fetch_snapshot(self) -> None:
|
||||
receiver = SnapshotReceiver(self.node_id, self.session_id)
|
||||
await self.command_sender.send(
|
||||
ForwarderCommand(
|
||||
origin=self._system_id,
|
||||
command=RequestSnapshot(requester_node_id=self.node_id),
|
||||
)
|
||||
)
|
||||
|
||||
with anyio.move_on_after(_SNAPSHOT_FETCH_TIMEOUT_SECONDS):
|
||||
with self.snapshot_chunk_receiver as chunks:
|
||||
async for chunk in chunks:
|
||||
received = receiver.ingest(chunk)
|
||||
if received is None:
|
||||
continue
|
||||
self.state = received.state
|
||||
self.event_router.set_buffer_start(
|
||||
received.last_event_applied_idx + 1
|
||||
)
|
||||
logger.info(
|
||||
f"API bootstrapped from snapshot at idx "
|
||||
f"{received.last_event_applied_idx}"
|
||||
)
|
||||
return
|
||||
logger.info(
|
||||
"API: no snapshot received before timeout; falling back to full event-log replay"
|
||||
)
|
||||
|
||||
async def _apply_state(self):
|
||||
with self.event_receiver as events:
|
||||
async for i_event in events:
|
||||
if i_event.idx <= self.state.last_event_applied_idx:
|
||||
continue
|
||||
self._event_log.append(i_event.event)
|
||||
self.state = apply(self.state, i_event)
|
||||
event = i_event.event
|
||||
|
||||
if isinstance(event, ChunkGenerated):
|
||||
if queue := self._image_generation_queues.get(
|
||||
event.command_id, None
|
||||
):
|
||||
assert isinstance(event.chunk, ImageChunk)
|
||||
try:
|
||||
await queue.send(event.chunk)
|
||||
except (BrokenResourceError, ClosedResourceError):
|
||||
self._image_generation_queues.pop(event.command_id, None)
|
||||
if queue := self._text_generation_queues.get(
|
||||
event.command_id, None
|
||||
):
|
||||
assert not isinstance(event.chunk, ImageChunk)
|
||||
try:
|
||||
await queue.send(event.chunk)
|
||||
except (BrokenResourceError, ClosedResourceError):
|
||||
self._text_generation_queues.pop(event.command_id, None)
|
||||
if isinstance(event, InstanceDeleted):
|
||||
self._close_streams_for_instance(event.instance_id)
|
||||
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)
|
||||
|
||||
async def _dispatch_chunk(self, event: ChunkGenerated) -> None:
|
||||
if queue := self._image_generation_queues.get(event.command_id, None):
|
||||
assert isinstance(event.chunk, ImageChunk)
|
||||
try:
|
||||
await queue.send(event.chunk)
|
||||
except (BrokenResourceError, ClosedResourceError):
|
||||
self._image_generation_queues.pop(event.command_id, None)
|
||||
if queue := self._text_generation_queues.get(event.command_id, None):
|
||||
assert not isinstance(event.chunk, ImageChunk)
|
||||
try:
|
||||
await queue.send(event.chunk)
|
||||
except (BrokenResourceError, ClosedResourceError):
|
||||
self._text_generation_queues.pop(event.command_id, None)
|
||||
|
||||
async def _reconcile_streams(self) -> None:
|
||||
while True:
|
||||
await anyio.sleep(1)
|
||||
self._reconcile_streams_once()
|
||||
|
||||
def _reconcile_streams_once(self) -> None:
|
||||
generation_tasks = [
|
||||
task
|
||||
for task in self.state.tasks.values()
|
||||
if isinstance(
|
||||
def _close_streams_for_instance(self, instance_id: InstanceId) -> None:
|
||||
"""Close any active generation streams for commands running on the given instance."""
|
||||
for task in self.state.tasks.values():
|
||||
if task.instance_id != instance_id:
|
||||
continue
|
||||
if not isinstance(
|
||||
task, (TextGenerationTask, ImageGenerationTask, ImageEditsTask)
|
||||
)
|
||||
]
|
||||
state_command_ids = {task.command_id for task in generation_tasks}
|
||||
self._observed_generation_commands.update(state_command_ids)
|
||||
|
||||
live_command_ids = {
|
||||
task.command_id
|
||||
for task in generation_tasks
|
||||
if task.instance_id in self.state.instances
|
||||
}
|
||||
queued_command_ids = set(self._text_generation_queues) | set(
|
||||
self._image_generation_queues
|
||||
)
|
||||
stale_command_ids = (
|
||||
self._observed_generation_commands - live_command_ids
|
||||
) & queued_command_ids
|
||||
self._close_streams_for_commands(stale_command_ids)
|
||||
|
||||
self._observed_generation_commands = (
|
||||
self._observed_generation_commands & queued_command_ids
|
||||
) | state_command_ids
|
||||
|
||||
def _close_streams_for_commands(self, command_ids: set[CommandId]) -> None:
|
||||
for command_id in command_ids:
|
||||
if sender := self._text_generation_queues.pop(command_id, None):
|
||||
):
|
||||
continue
|
||||
if sender := self._text_generation_queues.pop(task.command_id, None):
|
||||
sender.close()
|
||||
if sender := self._image_generation_queues.pop(command_id, None):
|
||||
if sender := self._image_generation_queues.pop(task.command_id, None):
|
||||
sender.close()
|
||||
|
||||
def _save_merged_trace(self, event: TracesMerged) -> None:
|
||||
|
||||
@@ -1,191 +0,0 @@
|
||||
# pyright: reportPrivateUsage=false
|
||||
|
||||
import hashlib
|
||||
|
||||
import anyio
|
||||
import pytest
|
||||
import zstandard
|
||||
|
||||
from exo.api.main import API
|
||||
from exo.routing.event_router import EventRouter
|
||||
from exo.shared.models.model_cards import ModelId
|
||||
from exo.shared.types.chunks import (
|
||||
ErrorChunk,
|
||||
PrefillProgressChunk,
|
||||
TokenChunk,
|
||||
ToolCallChunk,
|
||||
)
|
||||
from exo.shared.types.commands import ForwarderCommand, RequestSnapshot
|
||||
from exo.shared.types.common import CommandId, NodeId, SessionId, SystemId
|
||||
from exo.shared.types.events import (
|
||||
ChunkGenerated,
|
||||
Event,
|
||||
GlobalForwarderEvent,
|
||||
IndexedEvent,
|
||||
LocalForwarderEvent,
|
||||
TestEvent,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
|
||||
from exo.shared.types.state import State
|
||||
from exo.utils.channels import Receiver, Sender, channel
|
||||
|
||||
|
||||
class _FakeEventLog:
|
||||
def __init__(self) -> None:
|
||||
self.appended: list[Event] = []
|
||||
|
||||
def append(self, event: Event) -> None:
|
||||
self.appended.append(event)
|
||||
|
||||
|
||||
def _snapshot_chunk(
|
||||
state: State, *, requester_node_id: NodeId, session_id: SessionId
|
||||
) -> SnapshotChunk:
|
||||
body = zstandard.ZstdCompressor().compress(state.model_dump_json().encode("utf-8"))
|
||||
return SnapshotChunk.from_data(
|
||||
data=body,
|
||||
transfer_id=SnapshotTransferId("transfer-1"),
|
||||
requester_node_id=requester_node_id,
|
||||
session_id=session_id,
|
||||
schema_version=state.schema_version,
|
||||
last_event_applied_idx=state.last_event_applied_idx,
|
||||
chunk_index=0,
|
||||
total_chunks=1,
|
||||
sha256_hex=hashlib.sha256(body).hexdigest(),
|
||||
)
|
||||
|
||||
|
||||
def _api(
|
||||
node_id: NodeId, session_id: SessionId
|
||||
) -> tuple[
|
||||
API,
|
||||
EventRouter,
|
||||
Receiver[ForwarderCommand],
|
||||
Sender[SnapshotChunk],
|
||||
Sender[TransientEvent],
|
||||
Sender[IndexedEvent],
|
||||
_FakeEventLog,
|
||||
]:
|
||||
router_command_sender, _router_command_receiver = channel[ForwarderCommand]()
|
||||
_global_event_sender, global_event_receiver = channel[GlobalForwarderEvent]()
|
||||
local_event_sender, _local_event_receiver = channel[LocalForwarderEvent]()
|
||||
event_router = EventRouter(
|
||||
session_id=session_id,
|
||||
command_sender=router_command_sender,
|
||||
external_inbound=global_event_receiver,
|
||||
external_outbound=local_event_sender,
|
||||
)
|
||||
|
||||
event_sender, event_receiver = channel[IndexedEvent]()
|
||||
transient_sender, transient_receiver = channel[TransientEvent]()
|
||||
command_sender, command_receiver = channel[ForwarderCommand]()
|
||||
snapshot_sender, snapshot_receiver = channel[SnapshotChunk]()
|
||||
|
||||
api = object.__new__(API)
|
||||
api.node_id = node_id
|
||||
api.session_id = session_id
|
||||
api.event_router = event_router
|
||||
api.event_receiver = event_receiver
|
||||
api.transient_event_receiver = transient_receiver
|
||||
api.snapshot_chunk_receiver = snapshot_receiver
|
||||
api.command_sender = command_sender
|
||||
api._system_id = SystemId("api-system")
|
||||
api.state = State()
|
||||
event_log = _FakeEventLog()
|
||||
api._event_log = event_log # pyright: ignore[reportAttributeAccessIssue]
|
||||
api._image_generation_queues = {}
|
||||
api._text_generation_queues = {}
|
||||
return (
|
||||
api,
|
||||
event_router,
|
||||
command_receiver,
|
||||
snapshot_sender,
|
||||
transient_sender,
|
||||
event_sender,
|
||||
event_log,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_fetch_snapshot_applies_state_and_fast_forwards_router() -> None:
|
||||
node_id = NodeId("api")
|
||||
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
|
||||
(
|
||||
api,
|
||||
event_router,
|
||||
command_receiver,
|
||||
snapshot_sender,
|
||||
_transient_sender,
|
||||
_event_sender,
|
||||
_event_log,
|
||||
) = _api(node_id, session_id)
|
||||
state = State(last_event_applied_idx=7)
|
||||
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(api._fetch_snapshot)
|
||||
command = await command_receiver.receive()
|
||||
assert isinstance(command.command, RequestSnapshot)
|
||||
assert command.command.requester_node_id == node_id
|
||||
|
||||
await snapshot_sender.send(
|
||||
_snapshot_chunk(state, requester_node_id=node_id, session_id=session_id)
|
||||
)
|
||||
|
||||
assert api.state.last_event_applied_idx == 7
|
||||
assert event_router.event_buffer.next_idx_to_release == 8
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_apply_state_ignores_events_covered_by_snapshot() -> None:
|
||||
node_id = NodeId("api")
|
||||
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
|
||||
(
|
||||
api,
|
||||
_event_router,
|
||||
_command_receiver,
|
||||
_snapshot_sender,
|
||||
_transient_sender,
|
||||
event_sender,
|
||||
event_log,
|
||||
) = _api(node_id, session_id)
|
||||
api.state = State(last_event_applied_idx=7)
|
||||
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(api._apply_state)
|
||||
await event_sender.send(IndexedEvent(idx=7, event=TestEvent()))
|
||||
await event_sender.send(IndexedEvent(idx=8, event=TestEvent()))
|
||||
|
||||
while api.state.last_event_applied_idx != 8:
|
||||
await anyio.sleep(0.001)
|
||||
tg.cancel_scope.cancel()
|
||||
|
||||
assert len(event_log.appended) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_api_apply_transient_dispatches_generated_chunks() -> None:
|
||||
node_id = NodeId("api")
|
||||
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
|
||||
(
|
||||
api,
|
||||
_event_router,
|
||||
_command_receiver,
|
||||
_snapshot_sender,
|
||||
transient_sender,
|
||||
_event_sender,
|
||||
_event_log,
|
||||
) = _api(node_id, session_id)
|
||||
command_id = CommandId("cmd-a")
|
||||
chunk_sender, chunk_receiver = channel[
|
||||
TokenChunk | ErrorChunk | ToolCallChunk | PrefillProgressChunk
|
||||
]()
|
||||
api._text_generation_queues[command_id] = chunk_sender
|
||||
|
||||
chunk = ErrorChunk(model=ModelId("test-model"), error_message="test chunk")
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(api._apply_transient)
|
||||
await transient_sender.send(ChunkGenerated(command_id=command_id, chunk=chunk))
|
||||
|
||||
assert await chunk_receiver.receive() == chunk
|
||||
tg.cancel_scope.cancel()
|
||||
@@ -1,11 +1,11 @@
|
||||
# pyright: reportUnusedFunction=false, reportAny=false
|
||||
"""Tests that streaming queues reconcile against durable State."""
|
||||
"""Tests that InstanceDeleted events close active generation streams."""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from exo.api.main import API
|
||||
from exo.api.types import ImageGenerationTaskParams
|
||||
from exo.shared.types.common import CommandId, ModelId, NodeId
|
||||
from exo.shared.types.common import CommandId, ModelId
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.tasks import ImageGeneration, TextGeneration
|
||||
from exo.shared.types.text_generation import (
|
||||
@@ -13,16 +13,15 @@ from exo.shared.types.text_generation import (
|
||||
InputMessageContent,
|
||||
TextGenerationTaskParams,
|
||||
)
|
||||
from exo.shared.types.worker.instances import InstanceId, MlxRingInstance
|
||||
from exo.shared.types.worker.runners import ShardAssignments
|
||||
from exo.shared.types.worker.instances import InstanceId
|
||||
|
||||
|
||||
def _make_api_with_state(state: State) -> API:
|
||||
"""Create a minimal API instance with pre-set state."""
|
||||
api = object.__new__(API)
|
||||
api.state = state
|
||||
api._text_generation_queues = {} # pyright: ignore[reportPrivateUsage]
|
||||
api._image_generation_queues = {} # pyright: ignore[reportPrivateUsage]
|
||||
api._observed_generation_commands = set() # pyright: ignore[reportPrivateUsage]
|
||||
return api
|
||||
|
||||
|
||||
@@ -39,90 +38,45 @@ def _make_text_gen_task(
|
||||
)
|
||||
|
||||
|
||||
def _make_instance(instance_id: InstanceId) -> MlxRingInstance:
|
||||
return MlxRingInstance(
|
||||
instance_id=instance_id,
|
||||
shard_assignments=ShardAssignments(
|
||||
model_id=ModelId("test-model"),
|
||||
node_to_runner={},
|
||||
runner_to_shard={},
|
||||
),
|
||||
hosts_by_node={NodeId("node-1"): []},
|
||||
ephemeral_port=1,
|
||||
)
|
||||
|
||||
|
||||
def test_reconcile_closes_stream_when_task_instance_is_missing() -> None:
|
||||
def test_close_streams_for_deleted_instance() -> None:
|
||||
"""Deleting an instance closes the text generation sender for commands on that instance."""
|
||||
instance_id = InstanceId("inst-1")
|
||||
command_id = CommandId("cmd-1")
|
||||
task = _make_text_gen_task(instance_id, command_id)
|
||||
api = _make_api_with_state(State(tasks={task.task_id: task}, instances={}))
|
||||
|
||||
state = State(tasks={task.task_id: task})
|
||||
api = _make_api_with_state(state)
|
||||
|
||||
sender = MagicMock()
|
||||
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
|
||||
api._close_streams_for_instance(instance_id) # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
sender.close.assert_called_once()
|
||||
assert command_id not in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
def test_reconcile_keeps_stream_for_live_task_instance() -> None:
|
||||
instance_id = InstanceId("inst-live")
|
||||
command_id = CommandId("cmd-live")
|
||||
task = _make_text_gen_task(instance_id, command_id)
|
||||
api = _make_api_with_state(
|
||||
State(
|
||||
tasks={task.task_id: task},
|
||||
instances={instance_id: _make_instance(instance_id)},
|
||||
)
|
||||
)
|
||||
def test_close_streams_ignores_unrelated_instances() -> None:
|
||||
"""Deleting an instance does NOT close streams for commands on other instances."""
|
||||
target_id = InstanceId("inst-delete")
|
||||
other_id = InstanceId("inst-keep")
|
||||
other_cmd = CommandId("cmd-keep")
|
||||
other_task = _make_text_gen_task(other_id, other_cmd)
|
||||
|
||||
state = State(tasks={other_task.task_id: other_task})
|
||||
api = _make_api_with_state(state)
|
||||
|
||||
sender = MagicMock()
|
||||
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
|
||||
api._text_generation_queues[other_cmd] = sender # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
|
||||
api._close_streams_for_instance(target_id) # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
sender.close.assert_not_called()
|
||||
assert command_id in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
|
||||
assert other_cmd in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
def test_reconcile_does_not_close_command_before_state_observes_it() -> None:
|
||||
command_id = CommandId("cmd-not-created-yet")
|
||||
api = _make_api_with_state(State())
|
||||
|
||||
sender = MagicMock()
|
||||
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
sender.close.assert_not_called()
|
||||
assert command_id in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
def test_reconcile_closes_stream_after_observed_task_leaves_state() -> None:
|
||||
instance_id = InstanceId("inst-live")
|
||||
command_id = CommandId("cmd-deleted")
|
||||
task = _make_text_gen_task(instance_id, command_id)
|
||||
api = _make_api_with_state(
|
||||
State(
|
||||
tasks={task.task_id: task},
|
||||
instances={instance_id: _make_instance(instance_id)},
|
||||
)
|
||||
)
|
||||
|
||||
sender = MagicMock()
|
||||
api._text_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
|
||||
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
api.state = State(instances={instance_id: _make_instance(instance_id)})
|
||||
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
sender.close.assert_called_once()
|
||||
assert command_id not in api._text_generation_queues # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
def test_reconcile_closes_image_stream_when_task_instance_is_missing() -> None:
|
||||
def test_close_streams_for_deleted_instance_image_generation() -> None:
|
||||
"""Deleting an instance closes the image generation sender for commands on that instance."""
|
||||
instance_id = InstanceId("inst-img")
|
||||
command_id = CommandId("cmd-img")
|
||||
task = ImageGeneration(
|
||||
@@ -130,12 +84,14 @@ def test_reconcile_closes_image_stream_when_task_instance_is_missing() -> None:
|
||||
command_id=command_id,
|
||||
task_params=ImageGenerationTaskParams(prompt="a cat", model="test-model"),
|
||||
)
|
||||
api = _make_api_with_state(State(tasks={task.task_id: task}, instances={}))
|
||||
|
||||
state = State(tasks={task.task_id: task})
|
||||
api = _make_api_with_state(state)
|
||||
|
||||
sender = MagicMock()
|
||||
api._image_generation_queues[command_id] = sender # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
api._reconcile_streams_once() # pyright: ignore[reportPrivateUsage]
|
||||
api._close_streams_for_instance(instance_id) # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
sender.close.assert_called_once()
|
||||
assert command_id not in api._image_generation_queues # pyright: ignore[reportPrivateUsage]
|
||||
@@ -16,7 +16,8 @@ from exo.download.download_utils import (
|
||||
)
|
||||
from exo.download.shard_downloader import ShardDownloader
|
||||
from exo.shared.constants import EXO_DEFAULT_MODELS_DIR, EXO_MODELS_READ_ONLY_DIRS
|
||||
from exo.shared.models.model_cards import ModelId, get_model_cards
|
||||
from exo.shared.models import model_cards
|
||||
from exo.shared.models.model_cards import ModelId
|
||||
from exo.shared.types.commands import (
|
||||
CancelDownload,
|
||||
DeleteDownload,
|
||||
@@ -422,7 +423,7 @@ class DownloadCoordinator:
|
||||
)
|
||||
# Scan read-only directories for pre-downloaded models
|
||||
if EXO_MODELS_READ_ONLY_DIRS:
|
||||
for card in await get_model_cards():
|
||||
for card in await model_cards.card_cache.list_all():
|
||||
mid = card.model_id
|
||||
if mid in self.active_downloads:
|
||||
continue
|
||||
|
||||
@@ -11,11 +11,11 @@ from exo.download.download_utils import (
|
||||
download_shard,
|
||||
)
|
||||
from exo.download.shard_downloader import ShardDownloader
|
||||
from exo.shared.models import model_cards
|
||||
from exo.shared.models.model_cards import (
|
||||
ModelCard,
|
||||
ModelId,
|
||||
ModelTask,
|
||||
get_model_cards,
|
||||
)
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.worker.shards import (
|
||||
@@ -258,7 +258,7 @@ class ResumableShardDownloader(ShardDownloader):
|
||||
|
||||
tasks = [
|
||||
create_task(download_with_semaphore(model_card))
|
||||
for model_card in await get_model_cards()
|
||||
for model_card in await model_cards.card_cache.list_all()
|
||||
]
|
||||
|
||||
for task in asyncio.as_completed(tasks):
|
||||
|
||||
+1
-52
@@ -17,7 +17,6 @@ from exo.download.impl_shard_downloader import exo_shard_downloader
|
||||
from exo.master.main import Master
|
||||
from exo.routing.event_router import EventRouter
|
||||
from exo.routing.router import Router, get_node_id_keypair
|
||||
from exo.routing.transient_router import TransientRouter
|
||||
from exo.shared.constants import EXO_LOG
|
||||
from exo.shared.election import Election, ElectionResult
|
||||
from exo.shared.logging import logger_cleanup, logger_setup
|
||||
@@ -32,7 +31,6 @@ from exo.worker.main import Worker
|
||||
class Node:
|
||||
router: Router
|
||||
event_router: EventRouter
|
||||
transient_router: TransientRouter
|
||||
download_coordinator: DownloadCoordinator | None
|
||||
worker: Worker | None
|
||||
election: Election # Every node participates in election, as we do want a node to become master even if it isn't a master candidate if no master candidates are present.
|
||||
@@ -57,24 +55,16 @@ class Node:
|
||||
)
|
||||
await router.register_topic(topics.GLOBAL_EVENTS)
|
||||
await router.register_topic(topics.LOCAL_EVENTS)
|
||||
await router.register_topic(topics.TRANSIENT_EVENTS)
|
||||
await router.register_topic(topics.COMMANDS)
|
||||
await router.register_topic(topics.ELECTION_MESSAGES)
|
||||
await router.register_topic(topics.CONNECTION_MESSAGES)
|
||||
await router.register_topic(topics.DOWNLOAD_COMMANDS)
|
||||
await router.register_topic(topics.SNAPSHOT_RESPONSES)
|
||||
event_router = EventRouter(
|
||||
session_id,
|
||||
command_sender=router.sender(topics.COMMANDS),
|
||||
external_outbound=router.sender(topics.LOCAL_EVENTS),
|
||||
external_inbound=router.receiver(topics.GLOBAL_EVENTS),
|
||||
)
|
||||
transient_router = TransientRouter(
|
||||
node_id=node_id,
|
||||
session_id=session_id,
|
||||
external_outbound=router.sender(topics.TRANSIENT_EVENTS),
|
||||
external_inbound=router.receiver(topics.TRANSIENT_EVENTS),
|
||||
)
|
||||
|
||||
logger.info(f"Starting node {node_id}")
|
||||
|
||||
@@ -93,12 +83,8 @@ class Node:
|
||||
if args.spawn_api:
|
||||
api = API(
|
||||
node_id,
|
||||
session_id,
|
||||
port=args.api_port,
|
||||
event_router=event_router,
|
||||
event_receiver=event_router.receiver(),
|
||||
transient_event_receiver=transient_router.receiver(),
|
||||
snapshot_chunk_receiver=router.receiver(topics.SNAPSHOT_RESPONSES),
|
||||
command_sender=router.sender(topics.COMMANDS),
|
||||
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
|
||||
election_receiver=router.receiver(topics.ELECTION_MESSAGES),
|
||||
@@ -109,13 +95,8 @@ class Node:
|
||||
if not args.no_worker:
|
||||
worker = Worker(
|
||||
node_id,
|
||||
session_id,
|
||||
event_router=event_router,
|
||||
event_receiver=event_router.receiver(),
|
||||
event_sender=event_router.sender(),
|
||||
transient_event_receiver=transient_router.receiver(),
|
||||
transient_event_sender=transient_router.sender(),
|
||||
snapshot_chunk_receiver=router.receiver(topics.SNAPSHOT_RESPONSES),
|
||||
command_sender=router.sender(topics.COMMANDS),
|
||||
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
|
||||
api_port=args.api_port,
|
||||
@@ -128,12 +109,9 @@ class Node:
|
||||
node_id,
|
||||
session_id,
|
||||
event_sender=event_router.sender(),
|
||||
transient_event_receiver=transient_router.receiver(),
|
||||
transient_event_sender=transient_router.sender(),
|
||||
global_event_sender=router.sender(topics.GLOBAL_EVENTS),
|
||||
local_event_receiver=router.receiver(topics.LOCAL_EVENTS),
|
||||
command_receiver=router.receiver(topics.COMMANDS),
|
||||
snapshot_chunk_sender=router.sender(topics.SNAPSHOT_RESPONSES),
|
||||
download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS),
|
||||
)
|
||||
|
||||
@@ -154,7 +132,6 @@ class Node:
|
||||
return cls(
|
||||
router,
|
||||
event_router,
|
||||
transient_router,
|
||||
download_coordinator,
|
||||
worker,
|
||||
election,
|
||||
@@ -172,7 +149,6 @@ class Node:
|
||||
signal.signal(signal.SIGTERM, lambda _, __: self.shutdown())
|
||||
tg.start_soon(self.router.run)
|
||||
tg.start_soon(self.event_router.run)
|
||||
tg.start_soon(self.transient_router.run)
|
||||
tg.start_soon(self.election.run)
|
||||
if self.download_coordinator:
|
||||
tg.start_soon(self.download_coordinator.run)
|
||||
@@ -216,13 +192,6 @@ class Node:
|
||||
self.router.receiver(topics.GLOBAL_EVENTS),
|
||||
self.router.sender(topics.LOCAL_EVENTS),
|
||||
)
|
||||
self.transient_router.shutdown()
|
||||
self.transient_router = TransientRouter(
|
||||
node_id=self.node_id,
|
||||
session_id=result.session_id,
|
||||
external_outbound=self.router.sender(topics.TRANSIENT_EVENTS),
|
||||
external_inbound=self.router.receiver(topics.TRANSIENT_EVENTS),
|
||||
)
|
||||
|
||||
if (
|
||||
result.session_id.master_node_id == self.node_id
|
||||
@@ -238,14 +207,9 @@ class Node:
|
||||
self.node_id,
|
||||
result.session_id,
|
||||
event_sender=self.event_router.sender(),
|
||||
transient_event_receiver=self.transient_router.receiver(),
|
||||
transient_event_sender=self.transient_router.sender(),
|
||||
global_event_sender=self.router.sender(topics.GLOBAL_EVENTS),
|
||||
local_event_receiver=self.router.receiver(topics.LOCAL_EVENTS),
|
||||
command_receiver=self.router.receiver(topics.COMMANDS),
|
||||
snapshot_chunk_sender=self.router.sender(
|
||||
topics.SNAPSHOT_RESPONSES
|
||||
),
|
||||
download_command_sender=self.router.sender(
|
||||
topics.DOWNLOAD_COMMANDS
|
||||
),
|
||||
@@ -282,15 +246,8 @@ class Node:
|
||||
# TODO: add profiling etc to resource monitor
|
||||
self.worker = Worker(
|
||||
self.node_id,
|
||||
result.session_id,
|
||||
event_router=self.event_router,
|
||||
event_receiver=self.event_router.receiver(),
|
||||
event_sender=self.event_router.sender(),
|
||||
transient_event_receiver=self.transient_router.receiver(),
|
||||
transient_event_sender=self.transient_router.sender(),
|
||||
snapshot_chunk_receiver=self.router.receiver(
|
||||
topics.SNAPSHOT_RESPONSES
|
||||
),
|
||||
command_sender=self.router.sender(topics.COMMANDS),
|
||||
download_command_sender=self.router.sender(
|
||||
topics.DOWNLOAD_COMMANDS
|
||||
@@ -299,16 +256,8 @@ class Node:
|
||||
)
|
||||
self._tg.start_soon(self.worker.run)
|
||||
if self.api:
|
||||
self.api.reset(
|
||||
result.won_clock,
|
||||
result.session_id,
|
||||
self.event_router,
|
||||
self.event_router.receiver(),
|
||||
self.transient_router.receiver(),
|
||||
self.router.receiver(topics.SNAPSHOT_RESPONSES),
|
||||
)
|
||||
self.api.reset(result.won_clock, self.event_router.receiver())
|
||||
self._tg.start_soon(self.event_router.run)
|
||||
self._tg.start_soon(self.transient_router.run)
|
||||
else:
|
||||
if self.api:
|
||||
self.api.unpause(result.won_clock)
|
||||
|
||||
+2
-67
@@ -1,8 +1,6 @@
|
||||
import hashlib
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import anyio
|
||||
from anyio import to_thread
|
||||
from loguru import logger
|
||||
|
||||
from exo.master.placement import (
|
||||
@@ -27,7 +25,6 @@ from exo.shared.types.commands import (
|
||||
ImageGeneration,
|
||||
PlaceInstance,
|
||||
RequestEventLog,
|
||||
RequestSnapshot,
|
||||
SendInputChunk,
|
||||
SetInstanceLink,
|
||||
TaskCancelled,
|
||||
@@ -55,10 +52,8 @@ from exo.shared.types.events import (
|
||||
TraceEventData,
|
||||
TracesCollected,
|
||||
TracesMerged,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.instance_link import InstanceLink
|
||||
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.tasks import (
|
||||
ImageEdits as ImageEditsTask,
|
||||
@@ -79,15 +74,6 @@ from exo.utils.disk_event_log import DiskEventLog
|
||||
from exo.utils.event_buffer import MultiSourceBuffer
|
||||
from exo.utils.task_group import TaskGroup
|
||||
|
||||
_SNAPSHOT_CHUNK_BYTES = 512 * 1024
|
||||
_MAX_EVENT_LOG_REPLAY_BATCH = 1000
|
||||
|
||||
|
||||
def _encode_state_for_transfer(state: State) -> bytes:
|
||||
import zstandard
|
||||
|
||||
return zstandard.ZstdCompressor().compress(state.model_dump_json().encode("utf-8"))
|
||||
|
||||
|
||||
def _prefill_endpoint_for(state: State, decode_instance_id: InstanceId) -> str | None:
|
||||
decode = state.instances.get(decode_instance_id)
|
||||
@@ -137,11 +123,8 @@ class Master:
|
||||
*,
|
||||
command_receiver: Receiver[ForwarderCommand],
|
||||
event_sender: Sender[Event],
|
||||
transient_event_receiver: Receiver[TransientEvent],
|
||||
transient_event_sender: Sender[TransientEvent],
|
||||
local_event_receiver: Receiver[LocalForwarderEvent],
|
||||
global_event_sender: Sender[GlobalForwarderEvent],
|
||||
snapshot_chunk_sender: Sender[SnapshotChunk],
|
||||
download_command_sender: Sender[ForwarderDownloadCommand],
|
||||
):
|
||||
self.node_id = node_id
|
||||
@@ -151,10 +134,7 @@ class Master:
|
||||
self.command_task_mapping: dict[CommandId, TaskId] = {}
|
||||
self.command_receiver = command_receiver
|
||||
self.local_event_receiver = local_event_receiver
|
||||
self.transient_event_receiver = transient_event_receiver
|
||||
self.transient_event_sender = transient_event_sender
|
||||
self.global_event_sender = global_event_sender
|
||||
self.snapshot_chunk_sender = snapshot_chunk_sender
|
||||
self.download_command_sender = download_command_sender
|
||||
self.event_sender = event_sender
|
||||
self._system_id = SystemId()
|
||||
@@ -175,9 +155,6 @@ class Master:
|
||||
self._event_log.close()
|
||||
self.global_event_sender.close()
|
||||
self.local_event_receiver.close()
|
||||
self.transient_event_receiver.close()
|
||||
self.transient_event_sender.close()
|
||||
self.snapshot_chunk_sender.close()
|
||||
self.command_receiver.close()
|
||||
|
||||
async def shutdown(self):
|
||||
@@ -388,6 +365,7 @@ class Master:
|
||||
self.state.node_memory,
|
||||
self.state.node_network,
|
||||
download_status=self.state.downloads,
|
||||
node_rdma_ctl=self.state.node_rdma_ctl,
|
||||
)
|
||||
transition_events = get_transition_events(
|
||||
self.state.instances, placement, self.state.tasks
|
||||
@@ -464,19 +442,12 @@ class Master:
|
||||
case RequestEventLog():
|
||||
# We should just be able to send everything, since other buffers will ignore old messages
|
||||
# rate limit to 1000 at a time
|
||||
end = min(
|
||||
command.since_idx + _MAX_EVENT_LOG_REPLAY_BATCH,
|
||||
len(self._event_log),
|
||||
)
|
||||
end = min(command.since_idx + 1000, len(self._event_log))
|
||||
for i, event in enumerate(
|
||||
self._event_log.read_range(command.since_idx, end),
|
||||
start=command.since_idx,
|
||||
):
|
||||
await self._send_event(IndexedEvent(idx=i, event=event))
|
||||
case RequestSnapshot():
|
||||
self._tg.start_soon(
|
||||
self._serve_snapshot, command.requester_node_id
|
||||
)
|
||||
for event in generated_events:
|
||||
await self.event_sender.send(event)
|
||||
except ValueError as e:
|
||||
@@ -536,42 +507,6 @@ class Master:
|
||||
self._event_log.append(event)
|
||||
await self._send_event(indexed)
|
||||
|
||||
async def _serve_snapshot(self, requester_node_id: NodeId) -> None:
|
||||
state = self.state
|
||||
if state.last_event_applied_idx < 0:
|
||||
logger.info(
|
||||
f"RequestSnapshot from {requester_node_id} but master has no events yet"
|
||||
)
|
||||
return
|
||||
|
||||
body = await to_thread.run_sync(_encode_state_for_transfer, state)
|
||||
sha256 = hashlib.sha256(body).hexdigest()
|
||||
chunks = [
|
||||
body[i : i + _SNAPSHOT_CHUNK_BYTES]
|
||||
for i in range(0, len(body), _SNAPSHOT_CHUNK_BYTES)
|
||||
] or [b""]
|
||||
transfer_id = SnapshotTransferId()
|
||||
|
||||
logger.info(
|
||||
f"Serving snapshot to {requester_node_id}: "
|
||||
f"idx={state.last_event_applied_idx}, "
|
||||
f"{len(chunks)} chunk(s), {len(body)} bytes total"
|
||||
)
|
||||
for index, chunk in enumerate(chunks):
|
||||
await self.snapshot_chunk_sender.send(
|
||||
SnapshotChunk.from_data(
|
||||
data=chunk,
|
||||
transfer_id=transfer_id,
|
||||
requester_node_id=requester_node_id,
|
||||
session_id=self.session_id,
|
||||
schema_version=state.schema_version,
|
||||
last_event_applied_idx=state.last_event_applied_idx,
|
||||
chunk_index=index,
|
||||
total_chunks=len(chunks),
|
||||
sha256_hex=sha256,
|
||||
)
|
||||
)
|
||||
|
||||
# This function is re-entrant, take care!
|
||||
async def _send_event(self, event: IndexedEvent):
|
||||
# Convenience method since this line is ugly
|
||||
|
||||
@@ -28,7 +28,7 @@ from exo.shared.types.events import (
|
||||
TaskStatusUpdated,
|
||||
)
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo
|
||||
from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo, NodeRdmaCtlStatus
|
||||
from exo.shared.types.tasks import Task, TaskId, TaskStatus
|
||||
from exo.shared.types.worker.downloads import (
|
||||
DownloadCompleted,
|
||||
@@ -105,6 +105,7 @@ def place_instance(
|
||||
node_network: Mapping[NodeId, NodeNetworkInfo],
|
||||
required_nodes: set[NodeId] | None = None,
|
||||
download_status: Mapping[NodeId, Sequence[DownloadProgress]] | None = None,
|
||||
node_rdma_ctl: Mapping[NodeId, NodeRdmaCtlStatus] | None = None,
|
||||
) -> dict[InstanceId, Instance]:
|
||||
cycles = topology.get_cycles()
|
||||
candidate_cycles = list(filter(lambda it: len(it) >= command.min_nodes, cycles))
|
||||
@@ -166,8 +167,18 @@ def place_instance(
|
||||
|
||||
smallest_cycles = get_smallest_cycles(cycles_with_sufficient_memory)
|
||||
|
||||
rdma_ctl_status = node_rdma_ctl or {}
|
||||
|
||||
def _all_rdma_ctl_enabled(cycle: Cycle) -> bool:
|
||||
return all(
|
||||
((status := rdma_ctl_status.get(node_id)) is not None and status.enabled)
|
||||
for node_id in cycle
|
||||
)
|
||||
|
||||
smallest_rdma_cycles = [
|
||||
cycle for cycle in smallest_cycles if topology.is_rdma_cycle(cycle)
|
||||
cycle
|
||||
for cycle in smallest_cycles
|
||||
if topology.is_rdma_cycle(cycle) and _all_rdma_ctl_enabled(cycle)
|
||||
]
|
||||
|
||||
if command.instance_meta == InstanceMeta.MlxJaccl:
|
||||
|
||||
@@ -7,14 +7,12 @@ from loguru import logger
|
||||
|
||||
from exo.master.main import Master
|
||||
from exo.routing.router import get_node_id_keypair
|
||||
from exo.routing.snapshot_receiver import SnapshotReceiver
|
||||
from exo.shared.models.model_cards import ModelCard, ModelTask
|
||||
from exo.shared.types.commands import (
|
||||
CommandId,
|
||||
ForwarderCommand,
|
||||
ForwarderDownloadCommand,
|
||||
PlaceInstance,
|
||||
RequestSnapshot,
|
||||
TextGeneration,
|
||||
)
|
||||
from exo.shared.types.common import ModelId, NodeId, SessionId, SystemId
|
||||
@@ -26,13 +24,11 @@ from exo.shared.types.events import (
|
||||
LocalForwarderEvent,
|
||||
NodeGatheredInfo,
|
||||
TaskCreated,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.memory import Memory
|
||||
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 TextGeneration as TextGenerationTask
|
||||
from exo.shared.types.text_generation import (
|
||||
@@ -60,8 +56,6 @@ async def test_master():
|
||||
local_event_sender, le_receiver = channel[LocalForwarderEvent]()
|
||||
fcds, _fcdr = channel[ForwarderDownloadCommand]()
|
||||
ev_send, ev_recv = channel[Event]()
|
||||
transient_send, transient_recv = channel[TransientEvent]()
|
||||
snapshot_chunk_send, _snapshot_chunk_recv = channel[SnapshotChunk]()
|
||||
|
||||
async def mock_event_router():
|
||||
idx = 0
|
||||
@@ -95,12 +89,9 @@ async def test_master():
|
||||
node_id,
|
||||
session_id,
|
||||
event_sender=ev_send,
|
||||
transient_event_receiver=transient_recv,
|
||||
transient_event_sender=transient_send,
|
||||
global_event_sender=ge_sender,
|
||||
local_event_receiver=le_receiver,
|
||||
command_receiver=co_receiver,
|
||||
snapshot_chunk_sender=snapshot_chunk_send,
|
||||
download_command_sender=fcds,
|
||||
)
|
||||
logger.info("run the master")
|
||||
@@ -238,55 +229,3 @@ async def test_master():
|
||||
|
||||
ev_send.close()
|
||||
await master.shutdown()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_master_serves_snapshot_for_current_state():
|
||||
node_id = NodeId("master")
|
||||
requester_node_id = NodeId("worker")
|
||||
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_event_sender, transient_event_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_event_receiver,
|
||||
transient_event_sender=transient_event_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,
|
||||
)
|
||||
master.state = master.state.model_copy(update={"last_event_applied_idx": 12})
|
||||
|
||||
receiver = SnapshotReceiver(requester_node_id, session_id)
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(master.run)
|
||||
await command_sender.send(
|
||||
ForwarderCommand(
|
||||
origin=SystemId("api"),
|
||||
command=RequestSnapshot(requester_node_id=requester_node_id),
|
||||
)
|
||||
)
|
||||
|
||||
received = None
|
||||
while received is None:
|
||||
chunk = await snapshot_chunk_receiver.receive()
|
||||
received = receiver.ingest(chunk)
|
||||
|
||||
assert received.last_event_applied_idx == 12
|
||||
assert received.state.last_event_applied_idx == 12
|
||||
|
||||
await master.shutdown()
|
||||
tg.cancel_scope.cancel()
|
||||
@@ -21,7 +21,11 @@ from exo.shared.types.events import (
|
||||
)
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.multiaddr import Multiaddr
|
||||
from exo.shared.types.profiling import NetworkInterfaceInfo, NodeNetworkInfo
|
||||
from exo.shared.types.profiling import (
|
||||
NetworkInterfaceInfo,
|
||||
NodeNetworkInfo,
|
||||
NodeRdmaCtlStatus,
|
||||
)
|
||||
from exo.shared.types.tasks import TaskId, TaskStatus, TextGeneration
|
||||
from exo.shared.types.text_generation import (
|
||||
InputMessage,
|
||||
@@ -439,8 +443,21 @@ def test_tensor_rdma_backend_connectivity_matrix(
|
||||
min_nodes=1,
|
||||
)
|
||||
|
||||
node_rdma_ctl = {
|
||||
node_a: NodeRdmaCtlStatus(enabled=True),
|
||||
node_b: NodeRdmaCtlStatus(enabled=True),
|
||||
node_c: NodeRdmaCtlStatus(enabled=True),
|
||||
}
|
||||
|
||||
# act
|
||||
placements = place_instance(cic, topology, {}, node_memory, node_network)
|
||||
placements = place_instance(
|
||||
cic,
|
||||
topology,
|
||||
{},
|
||||
node_memory,
|
||||
node_network,
|
||||
node_rdma_ctl=node_rdma_ctl,
|
||||
)
|
||||
|
||||
# assert
|
||||
assert len(placements) == 1
|
||||
@@ -482,6 +499,131 @@ def test_tensor_rdma_backend_connectivity_matrix(
|
||||
assert len(ip_part.split(".")) == 4
|
||||
|
||||
|
||||
def _build_three_node_rdma_topology() -> tuple[
|
||||
Topology, NodeId, NodeId, NodeId, dict[NodeId, NodeNetworkInfo]
|
||||
]:
|
||||
topology = Topology()
|
||||
node_a = NodeId()
|
||||
node_b = NodeId()
|
||||
node_c = NodeId()
|
||||
|
||||
ethernet_interface = NetworkInterfaceInfo(name="en0", ip_address="10.0.0.1")
|
||||
ethernet_conn = SocketConnection(
|
||||
sink_multiaddr=Multiaddr(address="/ip4/10.0.0.1/tcp/8000")
|
||||
)
|
||||
node_network = {
|
||||
node_a: NodeNetworkInfo(interfaces=[ethernet_interface]),
|
||||
node_b: NodeNetworkInfo(interfaces=[ethernet_interface]),
|
||||
node_c: NodeNetworkInfo(interfaces=[ethernet_interface]),
|
||||
}
|
||||
|
||||
for n in (node_a, node_b, node_c):
|
||||
topology.add_node(n)
|
||||
|
||||
rdma_pairs = [
|
||||
(node_a, node_b, 3),
|
||||
(node_b, node_a, 3),
|
||||
(node_b, node_c, 4),
|
||||
(node_c, node_b, 4),
|
||||
(node_a, node_c, 5),
|
||||
(node_c, node_a, 5),
|
||||
]
|
||||
for src, sink, iface in rdma_pairs:
|
||||
topology.add_connection(
|
||||
Connection(source=src, sink=sink, edge=create_rdma_connection(iface))
|
||||
)
|
||||
|
||||
socket_pairs = [
|
||||
(node_a, node_b),
|
||||
(node_b, node_c),
|
||||
(node_c, node_a),
|
||||
(node_a, node_c),
|
||||
(node_b, node_a),
|
||||
(node_c, node_b),
|
||||
]
|
||||
for src, sink in socket_pairs:
|
||||
topology.add_connection(Connection(source=src, sink=sink, edge=ethernet_conn))
|
||||
|
||||
return topology, node_a, node_b, node_c, node_network
|
||||
|
||||
|
||||
def test_place_mlx_jaccl_rejects_when_a_node_has_rdma_ctl_disabled(
|
||||
model_card: ModelCard,
|
||||
):
|
||||
# arrange
|
||||
model_card = model_card.model_copy(
|
||||
update={"n_layers": 12, "storage_size": Memory.from_bytes(1500)}
|
||||
)
|
||||
topology, node_a, node_b, node_c, node_network = _build_three_node_rdma_topology()
|
||||
node_memory = {
|
||||
node_a: create_node_memory(500),
|
||||
node_b: create_node_memory(500),
|
||||
node_c: create_node_memory(500),
|
||||
}
|
||||
node_rdma_ctl = {
|
||||
node_a: NodeRdmaCtlStatus(enabled=True),
|
||||
node_b: NodeRdmaCtlStatus(enabled=True),
|
||||
node_c: NodeRdmaCtlStatus(enabled=False),
|
||||
}
|
||||
cic = PlaceInstance(
|
||||
sharding=Sharding.Tensor,
|
||||
instance_meta=InstanceMeta.MlxJaccl,
|
||||
command_id=CommandId(),
|
||||
model_card=model_card,
|
||||
min_nodes=3,
|
||||
)
|
||||
|
||||
# act / assert
|
||||
with pytest.raises(
|
||||
ValueError, match="Requested RDMA \\(MlxJaccl\\) but no RDMA-connected cycles"
|
||||
):
|
||||
place_instance(
|
||||
cic,
|
||||
topology,
|
||||
{},
|
||||
node_memory,
|
||||
node_network,
|
||||
node_rdma_ctl=node_rdma_ctl,
|
||||
)
|
||||
|
||||
|
||||
def test_place_mlx_jaccl_rejects_when_node_rdma_ctl_missing(model_card: ModelCard):
|
||||
"""A node with no observed rdma_ctl status must not participate in RDMA placement."""
|
||||
# arrange
|
||||
model_card = model_card.model_copy(
|
||||
update={"n_layers": 12, "storage_size": Memory.from_bytes(1500)}
|
||||
)
|
||||
topology, node_a, node_b, node_c, node_network = _build_three_node_rdma_topology()
|
||||
node_memory = {
|
||||
node_a: create_node_memory(500),
|
||||
node_b: create_node_memory(500),
|
||||
node_c: create_node_memory(500),
|
||||
}
|
||||
# node_c has no rdma_ctl entry at all
|
||||
node_rdma_ctl = {
|
||||
node_a: NodeRdmaCtlStatus(enabled=True),
|
||||
node_b: NodeRdmaCtlStatus(enabled=True),
|
||||
}
|
||||
cic = PlaceInstance(
|
||||
sharding=Sharding.Tensor,
|
||||
instance_meta=InstanceMeta.MlxJaccl,
|
||||
command_id=CommandId(),
|
||||
model_card=model_card,
|
||||
min_nodes=3,
|
||||
)
|
||||
|
||||
# act / assert
|
||||
with pytest.raises(ValueError):
|
||||
place_instance(
|
||||
cic,
|
||||
topology,
|
||||
{},
|
||||
node_memory,
|
||||
node_network,
|
||||
node_rdma_ctl=node_rdma_ctl,
|
||||
)
|
||||
|
||||
|
||||
def _make_task(
|
||||
instance_id: InstanceId,
|
||||
status: TaskStatus = TaskStatus.Running,
|
||||
|
||||
@@ -80,9 +80,6 @@ class EventRouter:
|
||||
def shutdown(self) -> None:
|
||||
self._tg.cancel_tasks()
|
||||
|
||||
def set_buffer_start(self, idx: int) -> None:
|
||||
self.event_buffer.fast_forward_to(idx)
|
||||
|
||||
async def _ingest(self, system_id: SystemId, recv: Receiver[Event]):
|
||||
idx = 0
|
||||
with recv as events:
|
||||
@@ -98,6 +95,7 @@ class EventRouter:
|
||||
self.out_for_delivery[event.event_id] = (anyio.current_time(), f_ev)
|
||||
|
||||
async def _run_ext_in(self):
|
||||
buf = OrderedBuffer[Event]()
|
||||
with self.external_inbound as events:
|
||||
async for event in events:
|
||||
if event.session != self.session_id:
|
||||
@@ -105,12 +103,12 @@ class EventRouter:
|
||||
if event.origin != self.session_id.master_node_id:
|
||||
continue
|
||||
|
||||
self.event_buffer.ingest(event.origin_idx, event.event)
|
||||
buf.ingest(event.origin_idx, event.event)
|
||||
event_id = event.event.event_id
|
||||
if event_id in self.out_for_delivery:
|
||||
self.out_for_delivery.pop(event_id)
|
||||
|
||||
drained = self.event_buffer.drain_indexed()
|
||||
drained = buf.drain_indexed()
|
||||
if drained:
|
||||
self._nack_attempts = 0
|
||||
if self._nack_cancel_scope:
|
||||
@@ -121,9 +119,7 @@ class EventRouter:
|
||||
or self._nack_cancel_scope.cancel_called
|
||||
):
|
||||
# Request the next index.
|
||||
self._tg.start_soon(
|
||||
self._nack_request, self.event_buffer.next_idx_to_release
|
||||
)
|
||||
self._tg.start_soon(self._nack_request, buf.next_idx_to_release)
|
||||
continue
|
||||
|
||||
for idx, event in drained:
|
||||
|
||||
@@ -1,109 +0,0 @@
|
||||
"""Reassembles a snapshot from a stream of `SnapshotChunk`s.
|
||||
|
||||
A receiver belongs to one node; it ignores chunks addressed to other
|
||||
requesters and chunks from prior sessions. Once a transfer's chunks have
|
||||
all been collected and the SHA-256 checks out, the snapshot is decoded into
|
||||
a `State`. Concurrent transfers (for the same requester) are tolerated:
|
||||
each is keyed by `transfer_id`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from dataclasses import dataclass, field
|
||||
from typing import final
|
||||
|
||||
import zstandard
|
||||
from loguru import logger
|
||||
|
||||
from exo.shared.types.common import NodeId, SessionId
|
||||
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
|
||||
from exo.shared.types.state import State
|
||||
|
||||
|
||||
@final
|
||||
@dataclass
|
||||
class _Assembly:
|
||||
"""Partial state for one in-flight snapshot transfer."""
|
||||
|
||||
total_chunks: int
|
||||
sha256_hex: str
|
||||
schema_version: int
|
||||
last_event_applied_idx: int
|
||||
chunks: dict[int, bytes] = field(default_factory=dict)
|
||||
|
||||
def is_complete(self) -> bool:
|
||||
return len(self.chunks) == self.total_chunks
|
||||
|
||||
def assemble(self) -> bytes:
|
||||
return b"".join(self.chunks[i] for i in range(self.total_chunks))
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReceivedSnapshot:
|
||||
last_event_applied_idx: int
|
||||
state: State
|
||||
|
||||
|
||||
class SnapshotReceiver:
|
||||
"""Filters and reassembles inbound chunks into a `ReceivedSnapshot`.
|
||||
|
||||
Stateless w.r.t. delivery: callers feed `SnapshotChunk`s in via `ingest`
|
||||
and check the return value for completion.
|
||||
"""
|
||||
|
||||
def __init__(self, my_node_id: NodeId, session_id: SessionId) -> None:
|
||||
self._my_node_id = my_node_id
|
||||
self._session_id = session_id
|
||||
self._assemblies: dict[SnapshotTransferId, _Assembly] = {}
|
||||
|
||||
def ingest(self, chunk: SnapshotChunk) -> ReceivedSnapshot | None:
|
||||
"""Absorb a chunk; return the snapshot once a transfer completes.
|
||||
|
||||
Returns None for partial transfers, mismatched recipients, stale
|
||||
sessions, version mismatches, or corrupt payloads.
|
||||
"""
|
||||
if chunk.requester_node_id != self._my_node_id:
|
||||
return None
|
||||
if chunk.session_id != self._session_id:
|
||||
return None
|
||||
|
||||
existing = self._assemblies.get(chunk.transfer_id)
|
||||
if existing is None:
|
||||
existing = _Assembly(
|
||||
total_chunks=chunk.total_chunks,
|
||||
sha256_hex=chunk.sha256_hex,
|
||||
schema_version=chunk.schema_version,
|
||||
last_event_applied_idx=chunk.last_event_applied_idx,
|
||||
)
|
||||
self._assemblies[chunk.transfer_id] = existing
|
||||
existing.chunks[chunk.chunk_index] = chunk.data
|
||||
|
||||
if not existing.is_complete():
|
||||
return None
|
||||
|
||||
# Transfer complete — finalise and remove from the in-flight map.
|
||||
del self._assemblies[chunk.transfer_id]
|
||||
body = existing.assemble()
|
||||
if hashlib.sha256(body).hexdigest() != existing.sha256_hex:
|
||||
logger.warning(f"Snapshot {chunk.transfer_id} failed checksum; discarding")
|
||||
return None
|
||||
try:
|
||||
decompressed = zstandard.ZstdDecompressor().decompress(body)
|
||||
state = State.model_validate_json(decompressed.decode("utf-8"))
|
||||
except (zstandard.ZstdError, ValueError) as e:
|
||||
logger.opt(exception=e).warning(
|
||||
f"Snapshot {chunk.transfer_id} could not be decoded; discarding"
|
||||
)
|
||||
return None
|
||||
if state.schema_version != existing.schema_version:
|
||||
# Should not happen — the master writes schema_version into both
|
||||
# the chunk meta and the State payload — but treat it as corrupt.
|
||||
logger.warning(
|
||||
f"Snapshot {chunk.transfer_id} schema version mismatch "
|
||||
f"(chunk={existing.schema_version}, state={state.schema_version})"
|
||||
)
|
||||
return None
|
||||
return ReceivedSnapshot(
|
||||
last_event_applied_idx=existing.last_event_applied_idx, state=state
|
||||
)
|
||||
@@ -141,28 +141,3 @@ async def test_drain_and_ingest_with_new_sequence(buffer: OrderedBuffer[Event]):
|
||||
assert [e[0] for e in drained] == [2]
|
||||
assert buffer.next_idx_to_release == 3
|
||||
assert 4 in buffer.store
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fast_forward_discards_buffered_stale_events(
|
||||
buffer: OrderedBuffer[Event],
|
||||
):
|
||||
buffer.ingest(*make_indexed_event(0))
|
||||
buffer.ingest(*make_indexed_event(2))
|
||||
buffer.ingest(*make_indexed_event(4))
|
||||
|
||||
buffer.fast_forward_to(3)
|
||||
|
||||
assert buffer.next_idx_to_release == 3
|
||||
assert set(buffer.store) == {4}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fast_forward_only_moves_forward(buffer: OrderedBuffer[Event]):
|
||||
buffer.ingest(*make_indexed_event(0))
|
||||
buffer.ingest(*make_indexed_event(1))
|
||||
buffer.drain()
|
||||
|
||||
buffer.fast_forward_to(1)
|
||||
|
||||
assert buffer.next_idx_to_release == 2
|
||||
@@ -1,151 +0,0 @@
|
||||
import hashlib
|
||||
|
||||
import pytest
|
||||
import zstandard
|
||||
|
||||
from exo.routing.snapshot_receiver import SnapshotReceiver
|
||||
from exo.shared.types.common import NodeId, SessionId
|
||||
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
|
||||
from exo.shared.types.state import State
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def session_id() -> SessionId:
|
||||
return SessionId(master_node_id=NodeId("master"), election_clock=0)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def my_node() -> NodeId:
|
||||
return NodeId("worker-1")
|
||||
|
||||
|
||||
def _encode(state: State) -> bytes:
|
||||
return zstandard.ZstdCompressor().compress(state.model_dump_json().encode("utf-8"))
|
||||
|
||||
|
||||
def _make_chunks(
|
||||
body: bytes,
|
||||
*,
|
||||
chunk_size: int,
|
||||
requester_node_id: NodeId,
|
||||
session_id: SessionId,
|
||||
state: State,
|
||||
transfer_id: SnapshotTransferId | None = None,
|
||||
) -> list[SnapshotChunk]:
|
||||
sha256 = hashlib.sha256(body).hexdigest()
|
||||
transfer_id = transfer_id or SnapshotTransferId()
|
||||
pieces = [body[i : i + chunk_size] for i in range(0, len(body), chunk_size)] or [
|
||||
b""
|
||||
]
|
||||
return [
|
||||
SnapshotChunk.from_data(
|
||||
data=piece,
|
||||
transfer_id=transfer_id,
|
||||
requester_node_id=requester_node_id,
|
||||
session_id=session_id,
|
||||
schema_version=state.schema_version,
|
||||
last_event_applied_idx=state.last_event_applied_idx,
|
||||
chunk_index=i,
|
||||
total_chunks=len(pieces),
|
||||
sha256_hex=sha256,
|
||||
)
|
||||
for i, piece in enumerate(pieces)
|
||||
]
|
||||
|
||||
|
||||
def test_completes_on_full_transfer(my_node: NodeId, session_id: SessionId):
|
||||
state = State(last_event_applied_idx=42)
|
||||
chunks = _make_chunks(
|
||||
_encode(state),
|
||||
chunk_size=64,
|
||||
requester_node_id=my_node,
|
||||
session_id=session_id,
|
||||
state=state,
|
||||
)
|
||||
receiver = SnapshotReceiver(my_node, session_id)
|
||||
|
||||
received = None
|
||||
for chunk in chunks:
|
||||
received = receiver.ingest(chunk)
|
||||
assert received is not None
|
||||
assert received.last_event_applied_idx == 42
|
||||
assert received.state.last_event_applied_idx == 42
|
||||
|
||||
|
||||
def test_handles_out_of_order_chunks(my_node: NodeId, session_id: SessionId):
|
||||
state = State(last_event_applied_idx=99)
|
||||
chunks = _make_chunks(
|
||||
_encode(state),
|
||||
chunk_size=32,
|
||||
requester_node_id=my_node,
|
||||
session_id=session_id,
|
||||
state=state,
|
||||
)
|
||||
receiver = SnapshotReceiver(my_node, session_id)
|
||||
|
||||
# Reverse them.
|
||||
received = None
|
||||
for chunk in reversed(chunks):
|
||||
received = receiver.ingest(chunk)
|
||||
assert received is not None
|
||||
assert received.last_event_applied_idx == 99
|
||||
|
||||
|
||||
def test_ignores_chunks_for_other_recipients(my_node: NodeId, session_id: SessionId):
|
||||
state = State(last_event_applied_idx=1)
|
||||
other = NodeId("worker-2")
|
||||
chunks = _make_chunks(
|
||||
_encode(state),
|
||||
chunk_size=64,
|
||||
requester_node_id=other,
|
||||
session_id=session_id,
|
||||
state=state,
|
||||
)
|
||||
receiver = SnapshotReceiver(my_node, session_id)
|
||||
for chunk in chunks:
|
||||
assert receiver.ingest(chunk) is None
|
||||
|
||||
|
||||
def test_ignores_chunks_from_stale_session(my_node: NodeId, session_id: SessionId):
|
||||
state = State(last_event_applied_idx=1)
|
||||
other_session = SessionId(master_node_id=NodeId("other-master"), election_clock=99)
|
||||
chunks = _make_chunks(
|
||||
_encode(state),
|
||||
chunk_size=64,
|
||||
requester_node_id=my_node,
|
||||
session_id=other_session,
|
||||
state=state,
|
||||
)
|
||||
receiver = SnapshotReceiver(my_node, session_id)
|
||||
for chunk in chunks:
|
||||
assert receiver.ingest(chunk) is None
|
||||
|
||||
|
||||
def test_discards_on_checksum_mismatch(my_node: NodeId, session_id: SessionId):
|
||||
state = State(last_event_applied_idx=1)
|
||||
chunks = _make_chunks(
|
||||
_encode(state),
|
||||
chunk_size=64,
|
||||
requester_node_id=my_node,
|
||||
session_id=session_id,
|
||||
state=state,
|
||||
)
|
||||
# Corrupt the last byte of the last chunk.
|
||||
original = chunks[-1]
|
||||
chunks[-1] = SnapshotChunk.from_data(
|
||||
data=original.data + b"\x00garbage",
|
||||
transfer_id=original.transfer_id,
|
||||
requester_node_id=original.requester_node_id,
|
||||
session_id=original.session_id,
|
||||
schema_version=original.schema_version,
|
||||
last_event_applied_idx=original.last_event_applied_idx,
|
||||
chunk_index=original.chunk_index,
|
||||
total_chunks=original.total_chunks,
|
||||
sha256_hex=original.sha256_hex,
|
||||
)
|
||||
|
||||
receiver = SnapshotReceiver(my_node, session_id)
|
||||
received = None
|
||||
for chunk in chunks:
|
||||
received = receiver.ingest(chunk)
|
||||
assert received is None
|
||||
@@ -1,37 +0,0 @@
|
||||
from exo.routing import topics
|
||||
from exo.shared.types.commands import ForwarderCommand, RequestSnapshot
|
||||
from exo.shared.types.common import NodeId, SessionId, SystemId
|
||||
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
|
||||
|
||||
|
||||
def test_request_snapshot_round_trips_through_forwarder_command() -> None:
|
||||
command = ForwarderCommand(
|
||||
origin=SystemId("system-1"),
|
||||
command=RequestSnapshot(requester_node_id=NodeId("worker-1")),
|
||||
)
|
||||
|
||||
restored = ForwarderCommand.model_validate_json(command.model_dump_json())
|
||||
|
||||
assert isinstance(restored.command, RequestSnapshot)
|
||||
assert restored.command.requester_node_id == NodeId("worker-1")
|
||||
|
||||
|
||||
def test_snapshot_response_topic_round_trips_chunk() -> None:
|
||||
chunk = SnapshotChunk.from_data(
|
||||
data=b"snapshot-bytes",
|
||||
transfer_id=SnapshotTransferId("transfer-1"),
|
||||
requester_node_id=NodeId("worker-1"),
|
||||
session_id=SessionId(master_node_id=NodeId("master"), election_clock=1),
|
||||
schema_version=1,
|
||||
last_event_applied_idx=42,
|
||||
chunk_index=0,
|
||||
total_chunks=1,
|
||||
sha256_hex="unused",
|
||||
)
|
||||
|
||||
restored = topics.SNAPSHOT_RESPONSES.deserialize(
|
||||
topics.SNAPSHOT_RESPONSES.serialize(chunk)
|
||||
)
|
||||
|
||||
assert restored == chunk
|
||||
assert restored.data == b"snapshot-bytes"
|
||||
@@ -1,69 +0,0 @@
|
||||
import anyio
|
||||
import pytest
|
||||
|
||||
from exo.routing import topics
|
||||
from exo.routing.transient_router import TransientRouter
|
||||
from exo.shared.types.common import NodeId, SessionId
|
||||
from exo.shared.types.events import (
|
||||
GlobalForwarderTransientEvent,
|
||||
TaskAcknowledged,
|
||||
)
|
||||
from exo.shared.types.tasks import TaskId
|
||||
from exo.utils.channels import channel
|
||||
|
||||
|
||||
def test_transient_topic_round_trips_wrapped_event() -> None:
|
||||
wrapped = GlobalForwarderTransientEvent(
|
||||
origin=NodeId("worker"),
|
||||
session=SessionId(master_node_id=NodeId("master"), election_clock=1),
|
||||
event=TaskAcknowledged(task_id=TaskId("task-1")),
|
||||
)
|
||||
|
||||
restored = topics.TRANSIENT_EVENTS.deserialize(
|
||||
topics.TRANSIENT_EVENTS.serialize(wrapped)
|
||||
)
|
||||
|
||||
assert restored == wrapped
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_transient_router_publishes_and_dispatches_session_events() -> None:
|
||||
node_id = NodeId("worker")
|
||||
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
|
||||
external_outbound_sender, external_outbound_receiver = channel[
|
||||
GlobalForwarderTransientEvent
|
||||
]()
|
||||
external_inbound_sender, external_inbound_receiver = channel[
|
||||
GlobalForwarderTransientEvent
|
||||
]()
|
||||
router = TransientRouter(
|
||||
node_id=node_id,
|
||||
session_id=session_id,
|
||||
external_outbound=external_outbound_sender,
|
||||
external_inbound=external_inbound_receiver,
|
||||
)
|
||||
local_sender = router.sender()
|
||||
local_receiver = router.receiver()
|
||||
event = TaskAcknowledged(task_id=TaskId("task-1"))
|
||||
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(router.run)
|
||||
await local_sender.send(event)
|
||||
|
||||
wrapped = await external_outbound_receiver.receive()
|
||||
assert wrapped.origin == node_id
|
||||
assert wrapped.session == session_id
|
||||
assert wrapped.event == event
|
||||
|
||||
await external_inbound_sender.send(wrapped)
|
||||
assert await local_receiver.receive() == event
|
||||
|
||||
stale_session = SessionId(master_node_id=NodeId("master"), election_clock=2)
|
||||
await external_inbound_sender.send(
|
||||
wrapped.model_copy(update={"session": stale_session})
|
||||
)
|
||||
await anyio.sleep(0)
|
||||
assert local_receiver.collect() == []
|
||||
|
||||
router.shutdown()
|
||||
tg.cancel_scope.cancel()
|
||||
@@ -6,10 +6,8 @@ from exo.shared.election import ElectionMessage
|
||||
from exo.shared.types.commands import ForwarderCommand, ForwarderDownloadCommand
|
||||
from exo.shared.types.events import (
|
||||
GlobalForwarderEvent,
|
||||
GlobalForwarderTransientEvent,
|
||||
LocalForwarderEvent,
|
||||
)
|
||||
from exo.shared.types.snapshots import SnapshotChunk
|
||||
from exo.utils.pydantic_ext import FrozenModel
|
||||
|
||||
|
||||
@@ -41,9 +39,6 @@ class TypedTopic[T: FrozenModel]:
|
||||
|
||||
GLOBAL_EVENTS = TypedTopic("global_events", PublishPolicy.Always, GlobalForwarderEvent)
|
||||
LOCAL_EVENTS = TypedTopic("local_events", PublishPolicy.Always, LocalForwarderEvent)
|
||||
TRANSIENT_EVENTS = TypedTopic(
|
||||
"transient_events", PublishPolicy.Always, GlobalForwarderTransientEvent
|
||||
)
|
||||
COMMANDS = TypedTopic("commands", PublishPolicy.Always, ForwarderCommand)
|
||||
ELECTION_MESSAGES = TypedTopic(
|
||||
"election_messages", PublishPolicy.Always, ElectionMessage
|
||||
@@ -54,6 +49,3 @@ CONNECTION_MESSAGES = TypedTopic(
|
||||
DOWNLOAD_COMMANDS = TypedTopic(
|
||||
"download_commands", PublishPolicy.Always, ForwarderDownloadCommand
|
||||
)
|
||||
SNAPSHOT_RESPONSES = TypedTopic(
|
||||
"snapshot_responses", PublishPolicy.Always, SnapshotChunk
|
||||
)
|
||||
@@ -1,84 +0,0 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from anyio import BrokenResourceError, ClosedResourceError
|
||||
from loguru import logger
|
||||
|
||||
from exo.shared.types.common import NodeId, SessionId
|
||||
from exo.shared.types.events import (
|
||||
GlobalForwarderTransientEvent,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.utils.channels import Receiver, Sender, channel
|
||||
from exo.utils.task_group import TaskGroup
|
||||
|
||||
|
||||
@dataclass
|
||||
class TransientRouter:
|
||||
"""Routes unordered, non-durable events over the transient-events topic."""
|
||||
|
||||
node_id: NodeId
|
||||
session_id: SessionId
|
||||
external_outbound: Sender[GlobalForwarderTransientEvent]
|
||||
external_inbound: Receiver[GlobalForwarderTransientEvent]
|
||||
_outbound: list[Sender[TransientEvent]] = field(init=False, default_factory=list)
|
||||
_inbound: list[Receiver[TransientEvent]] = field(init=False, default_factory=list)
|
||||
_tg: TaskGroup = field(init=False, default_factory=TaskGroup)
|
||||
|
||||
def sender(self) -> Sender[TransientEvent]:
|
||||
send, recv = channel[TransientEvent]()
|
||||
if self._tg.is_running():
|
||||
self._tg.start_soon(self._publish, recv)
|
||||
else:
|
||||
self._inbound.append(recv)
|
||||
return send
|
||||
|
||||
def receiver(self) -> Receiver[TransientEvent]:
|
||||
assert not self._tg.is_running(), (
|
||||
"TransientRouter receivers must be registered before run()"
|
||||
)
|
||||
send, recv = channel[TransientEvent]()
|
||||
self._outbound.append(send)
|
||||
return recv
|
||||
|
||||
def shutdown(self) -> None:
|
||||
self._tg.cancel_tasks()
|
||||
|
||||
async def run(self) -> None:
|
||||
try:
|
||||
async with self._tg as tg:
|
||||
for recv in self._inbound:
|
||||
tg.start_soon(self._publish, recv)
|
||||
tg.start_soon(self._dispatch_inbound)
|
||||
finally:
|
||||
self.external_outbound.close()
|
||||
for send in self._outbound:
|
||||
send.close()
|
||||
|
||||
async def _publish(self, recv: Receiver[TransientEvent]) -> None:
|
||||
with recv as events:
|
||||
async for event in events:
|
||||
await self.external_outbound.send(
|
||||
GlobalForwarderTransientEvent(
|
||||
origin=self.node_id,
|
||||
session=self.session_id,
|
||||
event=event,
|
||||
)
|
||||
)
|
||||
|
||||
async def _dispatch_inbound(self) -> None:
|
||||
with self.external_inbound as wrapped_events:
|
||||
async for wrapped in wrapped_events:
|
||||
if wrapped.session != self.session_id:
|
||||
continue
|
||||
stale: set[int] = set()
|
||||
for index, send in enumerate(self._outbound):
|
||||
try:
|
||||
await send.send(wrapped.event)
|
||||
except (ClosedResourceError, BrokenResourceError):
|
||||
stale.add(index)
|
||||
if stale:
|
||||
for index in sorted(stale, reverse=True):
|
||||
self._outbound.pop(index)
|
||||
logger.debug(
|
||||
f"TransientRouter dropped {len(stale)} closed receivers"
|
||||
)
|
||||
+27
-32
@@ -41,14 +41,7 @@ from exo.shared.types.profiling import (
|
||||
ThunderboltBridgeStatus,
|
||||
)
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.tasks import (
|
||||
ImageEdits,
|
||||
ImageGeneration,
|
||||
Task,
|
||||
TaskId,
|
||||
TaskStatus,
|
||||
TextGeneration,
|
||||
)
|
||||
from exo.shared.types.tasks import Task, TaskId, TaskStatus
|
||||
from exo.shared.types.topology import Connection, RDMAConnection
|
||||
from exo.shared.types.worker.downloads import DownloadProgress
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId
|
||||
@@ -73,6 +66,18 @@ from exo.utils.info_gatherer.info_gatherer import (
|
||||
)
|
||||
|
||||
|
||||
def _is_rdma_ctl_enabled(
|
||||
node_id: NodeId, node_rdma_ctl: Mapping[NodeId, NodeRdmaCtlStatus]
|
||||
) -> bool:
|
||||
"""A node is RDMA-capable only if rdma_ctl status has been observed as enabled.
|
||||
|
||||
Missing entries default to ``False`` — if we have not yet observed (or the node
|
||||
cannot run) ``rdma_ctl``, it must not participate in an RDMA-backed instance.
|
||||
"""
|
||||
status = node_rdma_ctl.get(node_id)
|
||||
return status is not None and status.enabled
|
||||
|
||||
|
||||
def event_apply(event: Event, state: State) -> State:
|
||||
"""Apply an event to state."""
|
||||
match event:
|
||||
@@ -80,6 +85,7 @@ def event_apply(event: Event, state: State) -> State:
|
||||
TestEvent()
|
||||
| ChunkGenerated()
|
||||
| TaskAcknowledged()
|
||||
| InputChunkReceived()
|
||||
| TracesCollected()
|
||||
| TracesMerged()
|
||||
): # Pass-through events that don't modify state
|
||||
@@ -102,8 +108,6 @@ def event_apply(event: Event, state: State) -> State:
|
||||
return apply_runner_status_updated(event, state)
|
||||
case TaskCreated():
|
||||
return apply_task_created(event, state)
|
||||
case InputChunkReceived():
|
||||
return apply_input_chunk_received(event, state)
|
||||
case TaskDeleted():
|
||||
return apply_task_deleted(event, state)
|
||||
case TaskFailed():
|
||||
@@ -168,32 +172,10 @@ def apply_task_created(event: TaskCreated, state: State) -> State:
|
||||
return state.model_copy(update={"tasks": new_tasks})
|
||||
|
||||
|
||||
def apply_input_chunk_received(event: InputChunkReceived, state: State) -> State:
|
||||
command_chunks = {
|
||||
**state.input_chunks.get(event.command_id, {}),
|
||||
event.chunk.chunk_index: event.chunk,
|
||||
}
|
||||
return state.model_copy(
|
||||
update={
|
||||
"input_chunks": {**state.input_chunks, event.command_id: command_chunks}
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def apply_task_deleted(event: TaskDeleted, state: State) -> State:
|
||||
task = state.tasks.get(event.task_id)
|
||||
new_tasks: Mapping[TaskId, Task] = {
|
||||
tid: task for tid, task in state.tasks.items() if tid != event.task_id
|
||||
}
|
||||
if isinstance(task, (TextGeneration, ImageGeneration, ImageEdits)):
|
||||
new_input_chunks = {
|
||||
command_id: chunks
|
||||
for command_id, chunks in state.input_chunks.items()
|
||||
if command_id != task.command_id
|
||||
}
|
||||
return state.model_copy(
|
||||
update={"tasks": new_tasks, "input_chunks": new_input_chunks}
|
||||
)
|
||||
return state.model_copy(update={"tasks": new_tasks})
|
||||
|
||||
|
||||
@@ -430,6 +412,9 @@ def apply_node_gathered_info(event: NodeGatheredInfo, state: State) -> State:
|
||||
for nid in state.node_thunderbolt
|
||||
for tb_ident in state.node_thunderbolt[nid].interfaces
|
||||
}
|
||||
source_is_rdma_enabled = _is_rdma_ctl_enabled(
|
||||
event.node_id, state.node_rdma_ctl
|
||||
)
|
||||
as_rdma_conns = [
|
||||
Connection(
|
||||
source=event.node_id,
|
||||
@@ -442,6 +427,10 @@ def apply_node_gathered_info(event: NodeGatheredInfo, state: State) -> State:
|
||||
for tb_conn in info.conns
|
||||
if tb_conn.source_uuid in conn_map
|
||||
if tb_conn.sink_uuid in conn_map
|
||||
if source_is_rdma_enabled
|
||||
and _is_rdma_ctl_enabled(
|
||||
conn_map[tb_conn.sink_uuid][0], state.node_rdma_ctl
|
||||
)
|
||||
]
|
||||
topology.replace_all_out_rdma_connections(event.node_id, as_rdma_conns)
|
||||
case ThunderboltBridgeInfo():
|
||||
@@ -465,6 +454,12 @@ def apply_node_gathered_info(event: NodeGatheredInfo, state: State) -> State:
|
||||
**state.node_rdma_ctl,
|
||||
event.node_id: NodeRdmaCtlStatus(enabled=info.enabled),
|
||||
}
|
||||
# If RDMA just got disabled on this node, drop any RDMA edges touching it
|
||||
# so placement / topology consumers cannot pick a disabled node for an
|
||||
# RDMA-backed instance. (Edges will repopulate on the next
|
||||
# MacThunderboltConnections poll once both endpoints are enabled again.)
|
||||
if not info.enabled:
|
||||
topology.remove_all_rdma_connections_touching(event.node_id)
|
||||
|
||||
return state.model_copy(update=update)
|
||||
|
||||
|
||||
@@ -4,13 +4,13 @@ from pathlib import Path
|
||||
|
||||
from exo.utils.dashboard_path import find_dashboard, find_resources
|
||||
|
||||
_EXO_HOME_ENV = os.environ.get("EXO_HOME", None)
|
||||
_EXO_HOME_ENV = os.environ.get("EXO_HOME", "")
|
||||
|
||||
|
||||
def _get_xdg_dir(env_var: str, fallback: str) -> Path:
|
||||
"""Get XDG directory, prioritising EXO_HOME environment variable if its set. On non-Linux platforms, default to ~/.exo."""
|
||||
|
||||
if _EXO_HOME_ENV is not None:
|
||||
if _EXO_HOME_ENV != "":
|
||||
return Path.home() / _EXO_HOME_ENV
|
||||
|
||||
if sys.platform != "linux":
|
||||
|
||||
@@ -39,7 +39,57 @@ _BUILTIN_CARD_DIRS = [
|
||||
Path(RESOURCES_DIR) / "image_model_cards",
|
||||
]
|
||||
|
||||
_card_cache: dict[ModelId, "ModelCard"] = {}
|
||||
|
||||
class _CardCache:
|
||||
def __init__(self):
|
||||
self.cc: dict[ModelId, "ModelCard"] = {}
|
||||
|
||||
def get(self, model_id: ModelId) -> "ModelCard | None":
|
||||
return self.cc.get(model_id)
|
||||
|
||||
async def save(self, card: "ModelCard"):
|
||||
self.cc[card.model_id] = card
|
||||
try:
|
||||
await card.save_to_custom_dir()
|
||||
except OSError as e:
|
||||
logger.warning(f"failed to save custom model card ({e.strerror})")
|
||||
|
||||
async def pop(self, model_id: ModelId) -> "ModelCard | None":
|
||||
"""Delete a user-added custom model card. Returns True if deleted."""
|
||||
card_path = _custom_cards_dir / (ModelId(model_id).normalize() + ".toml")
|
||||
try:
|
||||
if await card_path.exists():
|
||||
await card_path.unlink()
|
||||
return self.cc.pop(model_id, None)
|
||||
except OSError as e:
|
||||
logger.warning(f"failed to delete custom model card ({e.strerror})")
|
||||
|
||||
async def list_all(self) -> list["ModelCard"]:
|
||||
if len(self.cc) == 0:
|
||||
await self.refresh()
|
||||
if EXO_ENABLE_IMAGE_MODELS:
|
||||
return list(self.cc.values())
|
||||
return [c for c in self.cc.values() if not _is_image_card(c)]
|
||||
|
||||
async def _load_cards_from_dir(self, directory: Path, *, is_custom: bool) -> None:
|
||||
"""Load all TOML model cards from a directory into the cache."""
|
||||
async for toml_file in directory.rglob("*.toml"):
|
||||
try:
|
||||
card = await ModelCard.load_from_path(toml_file)
|
||||
if is_custom:
|
||||
card = card.model_copy(update={"is_custom": True})
|
||||
if self.get(card.model_id) is None:
|
||||
self.cc[card.model_id] = card
|
||||
except (ValidationError, TOMLKitError):
|
||||
pass
|
||||
|
||||
async def refresh(self) -> None:
|
||||
for path in _BUILTIN_CARD_DIRS:
|
||||
await self._load_cards_from_dir(path, is_custom=False)
|
||||
await self._load_cards_from_dir(_custom_cards_dir, is_custom=True)
|
||||
|
||||
|
||||
card_cache = _CardCache()
|
||||
|
||||
|
||||
def detect_vision_from_config(model_id: ModelId) -> "VisionCardConfig | None":
|
||||
@@ -59,42 +109,10 @@ def detect_vision_from_config(model_id: ModelId) -> "VisionCardConfig | None":
|
||||
return None
|
||||
|
||||
|
||||
async def _load_cards_from_dir(directory: Path, *, is_custom: bool) -> None:
|
||||
"""Load all TOML model cards from a directory into the cache."""
|
||||
async for toml_file in directory.rglob("*.toml"):
|
||||
try:
|
||||
card = await ModelCard.load_from_path(toml_file)
|
||||
if is_custom:
|
||||
card = card.model_copy(update={"is_custom": True})
|
||||
if card.model_id not in _card_cache:
|
||||
_card_cache[card.model_id] = card
|
||||
except (ValidationError, TOMLKitError):
|
||||
pass
|
||||
|
||||
|
||||
async def _refresh_card_cache() -> None:
|
||||
for path in _BUILTIN_CARD_DIRS:
|
||||
await _load_cards_from_dir(path, is_custom=False)
|
||||
await _load_cards_from_dir(_custom_cards_dir, is_custom=True)
|
||||
|
||||
|
||||
def _is_image_card(card: "ModelCard") -> bool:
|
||||
return any(t in (ModelTask.TextToImage, ModelTask.ImageToImage) for t in card.tasks)
|
||||
|
||||
|
||||
def get_card(model_id: ModelId) -> "ModelCard | None":
|
||||
"""Look up a single model card from the cache by ID."""
|
||||
return _card_cache.get(model_id)
|
||||
|
||||
|
||||
async def get_model_cards() -> list["ModelCard"]:
|
||||
if len(_card_cache) == 0:
|
||||
await _refresh_card_cache()
|
||||
if EXO_ENABLE_IMAGE_MODELS:
|
||||
return list(_card_cache.values())
|
||||
return [c for c in _card_cache.values() if not _is_image_card(c)]
|
||||
|
||||
|
||||
class ModelTask(str, Enum):
|
||||
TextGeneration = "TextGeneration"
|
||||
TextToImage = "TextToImage"
|
||||
@@ -196,14 +214,13 @@ class ModelCard(FrozenModel):
|
||||
# Is it okay that model card.load defaults to network access if the card doesn't exist? do we want to be more explicit here?
|
||||
@staticmethod
|
||||
async def load(model_id: ModelId) -> "ModelCard":
|
||||
if model_id not in _card_cache:
|
||||
await _refresh_card_cache()
|
||||
if (mc := _card_cache.get(model_id)) is not None:
|
||||
if card_cache.get(model_id) is None:
|
||||
await card_cache.refresh()
|
||||
if (mc := card_cache.get(model_id)) is not None:
|
||||
return mc
|
||||
|
||||
mc = await ModelCard.fetch_from_hf(model_id)
|
||||
await mc.save_to_custom_dir()
|
||||
_card_cache[model_id] = mc
|
||||
return mc
|
||||
|
||||
@staticmethod
|
||||
@@ -233,21 +250,6 @@ class ModelCard(FrozenModel):
|
||||
)
|
||||
|
||||
|
||||
def add_to_card_cache(card: "ModelCard") -> None:
|
||||
"""Add or update a model card in the in-memory cache."""
|
||||
_card_cache[card.model_id] = card
|
||||
|
||||
|
||||
async def delete_custom_card(model_id: ModelId) -> bool:
|
||||
"""Delete a user-added custom model card. Returns True if deleted."""
|
||||
card_path = _custom_cards_dir / (ModelId(model_id).normalize() + ".toml")
|
||||
if await card_path.exists():
|
||||
await card_path.unlink()
|
||||
_card_cache.pop(model_id, None)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class ConfigData(BaseModel):
|
||||
model_config = {"extra": "ignore"} # Allow unknown fields
|
||||
|
||||
|
||||
@@ -1,85 +0,0 @@
|
||||
from exo.shared.apply import apply
|
||||
from exo.shared.models.model_cards import ModelId
|
||||
from exo.shared.types.chunks import InputImageChunk
|
||||
from exo.shared.types.common import CommandId
|
||||
from exo.shared.types.events import (
|
||||
IndexedEvent,
|
||||
InputChunkReceived,
|
||||
TaskCreated,
|
||||
TaskDeleted,
|
||||
)
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.tasks import TaskId, TaskStatus, TextGeneration
|
||||
from exo.shared.types.text_generation import (
|
||||
InputMessage,
|
||||
InputMessageContent,
|
||||
TextGenerationTaskParams,
|
||||
)
|
||||
from exo.shared.types.worker.instances import InstanceId
|
||||
|
||||
|
||||
def test_apply_input_chunk_received_stores_chunk_in_state() -> None:
|
||||
command_id = CommandId("command-1")
|
||||
chunk = InputImageChunk(
|
||||
model=ModelId("mlx-community/test-model"),
|
||||
command_id=command_id,
|
||||
data="abc",
|
||||
chunk_index=0,
|
||||
total_chunks=1,
|
||||
image_index=0,
|
||||
)
|
||||
|
||||
state = apply(
|
||||
State(),
|
||||
IndexedEvent(
|
||||
idx=0,
|
||||
event=InputChunkReceived(command_id=command_id, chunk=chunk),
|
||||
),
|
||||
)
|
||||
|
||||
assert state.input_chunks == {command_id: {0: chunk}}
|
||||
|
||||
|
||||
def test_apply_task_deleted_removes_chunks_for_generation_command() -> None:
|
||||
command_id = CommandId("command-1")
|
||||
task_id = TaskId("task-1")
|
||||
chunk = InputImageChunk(
|
||||
model=ModelId("mlx-community/test-model"),
|
||||
command_id=command_id,
|
||||
data="abc",
|
||||
chunk_index=0,
|
||||
total_chunks=1,
|
||||
image_index=0,
|
||||
)
|
||||
task = TextGeneration(
|
||||
task_id=task_id,
|
||||
instance_id=InstanceId("instance-1"),
|
||||
task_status=TaskStatus.Pending,
|
||||
command_id=command_id,
|
||||
task_params=TextGenerationTaskParams(
|
||||
model=ModelId("mlx-community/test-model"),
|
||||
input=[
|
||||
InputMessage(role="user", content=InputMessageContent("hello")),
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
state = State()
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(
|
||||
idx=0,
|
||||
event=InputChunkReceived(command_id=command_id, chunk=chunk),
|
||||
),
|
||||
)
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=1, event=TaskCreated(task_id=task_id, task=task)),
|
||||
)
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=2, event=TaskDeleted(task_id=task_id)),
|
||||
)
|
||||
|
||||
assert state.tasks == {}
|
||||
assert state.input_chunks == {}
|
||||
@@ -0,0 +1,231 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from exo.shared.apply import apply_node_gathered_info
|
||||
from exo.shared.topology import Topology
|
||||
from exo.shared.types.common import NodeId
|
||||
from exo.shared.types.events import NodeGatheredInfo
|
||||
from exo.shared.types.profiling import (
|
||||
NodeRdmaCtlStatus,
|
||||
NodeThunderboltInfo,
|
||||
)
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.thunderbolt import ThunderboltConnection, ThunderboltIdentifier
|
||||
from exo.shared.types.topology import RDMAConnection
|
||||
from exo.utils.info_gatherer.info_gatherer import (
|
||||
MacThunderboltConnections,
|
||||
RdmaCtlStatus,
|
||||
)
|
||||
|
||||
|
||||
def _now() -> str:
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
|
||||
def _make_state_with_thunderbolt_idents(
|
||||
*node_ids_and_uuids: tuple[NodeId, str, str],
|
||||
rdma_ctl: dict[NodeId, NodeRdmaCtlStatus] | None = None,
|
||||
) -> State:
|
||||
"""Build a State with Thunderbolt identifiers per node so the apply MacThunderboltConnections
|
||||
case can resolve uuid -> (node, iface)."""
|
||||
node_thunderbolt = {
|
||||
nid: NodeThunderboltInfo(
|
||||
interfaces=[ThunderboltIdentifier(rdma_interface=iface, domain_uuid=uuid)]
|
||||
)
|
||||
for nid, uuid, iface in node_ids_and_uuids
|
||||
}
|
||||
return State(
|
||||
node_thunderbolt=node_thunderbolt,
|
||||
node_rdma_ctl=rdma_ctl or {},
|
||||
)
|
||||
|
||||
|
||||
def _has_rdma_edge(topology: Topology, source: NodeId, sink: NodeId) -> bool:
|
||||
return any(
|
||||
isinstance(edge, RDMAConnection)
|
||||
for edge in topology.get_all_connections_between(source, sink)
|
||||
)
|
||||
|
||||
|
||||
def test_mac_thunderbolt_connections_emits_rdma_when_both_endpoints_enabled():
|
||||
node_a = NodeId()
|
||||
node_b = NodeId()
|
||||
state = _make_state_with_thunderbolt_idents(
|
||||
(node_a, "uuid-a", "rdma_en1"),
|
||||
(node_b, "uuid-b", "rdma_en1"),
|
||||
rdma_ctl={
|
||||
node_a: NodeRdmaCtlStatus(enabled=True),
|
||||
node_b: NodeRdmaCtlStatus(enabled=True),
|
||||
},
|
||||
)
|
||||
|
||||
event = NodeGatheredInfo(
|
||||
node_id=node_a,
|
||||
when=_now(),
|
||||
info=MacThunderboltConnections(
|
||||
conns=[ThunderboltConnection(source_uuid="uuid-a", sink_uuid="uuid-b")]
|
||||
),
|
||||
)
|
||||
|
||||
new_state = apply_node_gathered_info(event, state)
|
||||
|
||||
assert _has_rdma_edge(new_state.topology, node_a, node_b)
|
||||
|
||||
|
||||
def test_mac_thunderbolt_connections_skips_rdma_when_source_rdma_ctl_disabled():
|
||||
node_a = NodeId()
|
||||
node_b = NodeId()
|
||||
state = _make_state_with_thunderbolt_idents(
|
||||
(node_a, "uuid-a", "rdma_en1"),
|
||||
(node_b, "uuid-b", "rdma_en1"),
|
||||
rdma_ctl={
|
||||
node_a: NodeRdmaCtlStatus(enabled=False),
|
||||
node_b: NodeRdmaCtlStatus(enabled=True),
|
||||
},
|
||||
)
|
||||
|
||||
event = NodeGatheredInfo(
|
||||
node_id=node_a,
|
||||
when=_now(),
|
||||
info=MacThunderboltConnections(
|
||||
conns=[ThunderboltConnection(source_uuid="uuid-a", sink_uuid="uuid-b")]
|
||||
),
|
||||
)
|
||||
|
||||
new_state = apply_node_gathered_info(event, state)
|
||||
|
||||
assert not _has_rdma_edge(new_state.topology, node_a, node_b)
|
||||
|
||||
|
||||
def test_mac_thunderbolt_connections_skips_rdma_when_sink_rdma_ctl_disabled():
|
||||
node_a = NodeId()
|
||||
node_b = NodeId()
|
||||
state = _make_state_with_thunderbolt_idents(
|
||||
(node_a, "uuid-a", "rdma_en1"),
|
||||
(node_b, "uuid-b", "rdma_en1"),
|
||||
rdma_ctl={
|
||||
node_a: NodeRdmaCtlStatus(enabled=True),
|
||||
node_b: NodeRdmaCtlStatus(enabled=False),
|
||||
},
|
||||
)
|
||||
|
||||
event = NodeGatheredInfo(
|
||||
node_id=node_a,
|
||||
when=_now(),
|
||||
info=MacThunderboltConnections(
|
||||
conns=[ThunderboltConnection(source_uuid="uuid-a", sink_uuid="uuid-b")]
|
||||
),
|
||||
)
|
||||
|
||||
new_state = apply_node_gathered_info(event, state)
|
||||
|
||||
assert not _has_rdma_edge(new_state.topology, node_a, node_b)
|
||||
|
||||
|
||||
def test_mac_thunderbolt_connections_skips_rdma_when_rdma_ctl_status_missing():
|
||||
"""Missing rdma_ctl status defaults to not-enabled — node is RDMA-incapable."""
|
||||
node_a = NodeId()
|
||||
node_b = NodeId()
|
||||
state = _make_state_with_thunderbolt_idents(
|
||||
(node_a, "uuid-a", "rdma_en1"),
|
||||
(node_b, "uuid-b", "rdma_en1"),
|
||||
rdma_ctl={
|
||||
node_a: NodeRdmaCtlStatus(enabled=True),
|
||||
# node_b intentionally absent
|
||||
},
|
||||
)
|
||||
|
||||
event = NodeGatheredInfo(
|
||||
node_id=node_a,
|
||||
when=_now(),
|
||||
info=MacThunderboltConnections(
|
||||
conns=[ThunderboltConnection(source_uuid="uuid-a", sink_uuid="uuid-b")]
|
||||
),
|
||||
)
|
||||
|
||||
new_state = apply_node_gathered_info(event, state)
|
||||
|
||||
assert not _has_rdma_edge(new_state.topology, node_a, node_b)
|
||||
|
||||
|
||||
def test_rdma_ctl_status_disabled_purges_existing_rdma_edges():
|
||||
"""When a node reports rdma_ctl disabled, all RDMA edges touching it must be removed."""
|
||||
node_a = NodeId()
|
||||
node_b = NodeId()
|
||||
|
||||
# Start with both nodes RDMA-enabled and existing RDMA edges in the topology.
|
||||
state = _make_state_with_thunderbolt_idents(
|
||||
(node_a, "uuid-a", "rdma_en1"),
|
||||
(node_b, "uuid-b", "rdma_en1"),
|
||||
rdma_ctl={
|
||||
node_a: NodeRdmaCtlStatus(enabled=True),
|
||||
node_b: NodeRdmaCtlStatus(enabled=True),
|
||||
},
|
||||
)
|
||||
state = apply_node_gathered_info(
|
||||
NodeGatheredInfo(
|
||||
node_id=node_a,
|
||||
when=_now(),
|
||||
info=MacThunderboltConnections(
|
||||
conns=[ThunderboltConnection(source_uuid="uuid-a", sink_uuid="uuid-b")]
|
||||
),
|
||||
),
|
||||
state,
|
||||
)
|
||||
state = apply_node_gathered_info(
|
||||
NodeGatheredInfo(
|
||||
node_id=node_b,
|
||||
when=_now(),
|
||||
info=MacThunderboltConnections(
|
||||
conns=[ThunderboltConnection(source_uuid="uuid-b", sink_uuid="uuid-a")]
|
||||
),
|
||||
),
|
||||
state,
|
||||
)
|
||||
assert _has_rdma_edge(state.topology, node_a, node_b)
|
||||
assert _has_rdma_edge(state.topology, node_b, node_a)
|
||||
|
||||
# Now node_a flips to rdma_ctl disabled — both directions of RDMA edge must drop.
|
||||
state = apply_node_gathered_info(
|
||||
NodeGatheredInfo(
|
||||
node_id=node_a, when=_now(), info=RdmaCtlStatus(enabled=False)
|
||||
),
|
||||
state,
|
||||
)
|
||||
|
||||
assert not _has_rdma_edge(state.topology, node_a, node_b)
|
||||
assert not _has_rdma_edge(state.topology, node_b, node_a)
|
||||
assert state.node_rdma_ctl[node_a].enabled is False
|
||||
|
||||
|
||||
def test_topology_remove_all_rdma_connections_touching_keeps_socket_edges():
|
||||
"""Purging RDMA edges for a disabled node must not affect non-RDMA edges."""
|
||||
from exo.shared.types.multiaddr import Multiaddr
|
||||
from exo.shared.types.topology import Connection, SocketConnection
|
||||
|
||||
topology = Topology()
|
||||
node_a = NodeId()
|
||||
node_b = NodeId()
|
||||
topology.add_node(node_a)
|
||||
topology.add_node(node_b)
|
||||
topology.add_connection(
|
||||
Connection(
|
||||
source=node_a,
|
||||
sink=node_b,
|
||||
edge=RDMAConnection(
|
||||
source_rdma_iface="rdma_en1", sink_rdma_iface="rdma_en1"
|
||||
),
|
||||
)
|
||||
)
|
||||
socket_edge = SocketConnection(
|
||||
sink_multiaddr=Multiaddr(address="/ip4/10.0.0.1/tcp/8000")
|
||||
)
|
||||
topology.add_connection(Connection(source=node_a, sink=node_b, edge=socket_edge))
|
||||
|
||||
topology.remove_all_rdma_connections_touching(node_a)
|
||||
|
||||
assert not _has_rdma_edge(topology, node_a, node_b)
|
||||
# Socket edge survives.
|
||||
assert any(
|
||||
isinstance(edge, SocketConnection)
|
||||
for edge in topology.get_all_connections_between(node_a, node_b)
|
||||
)
|
||||
@@ -25,7 +25,6 @@ def test_state_serialization_roundtrip() -> None:
|
||||
json_repr = state.model_dump_json()
|
||||
restored_state = State.model_validate_json(json_repr)
|
||||
|
||||
assert restored_state.schema_version == state.schema_version
|
||||
assert (
|
||||
state.topology.to_snapshot().nodes
|
||||
== restored_state.topology.to_snapshot().nodes
|
||||
|
||||
@@ -169,6 +169,22 @@ class Topology:
|
||||
for conn in new_connections:
|
||||
self.add_connection(conn)
|
||||
|
||||
def remove_all_rdma_connections_touching(self, node_id: NodeId) -> None:
|
||||
"""Remove every RDMA edge incident to ``node_id`` (incoming or outgoing)."""
|
||||
if node_id not in self._vertex_indices:
|
||||
return
|
||||
rx_idx = self._vertex_indices[node_id]
|
||||
rdma_edge_idxs = [
|
||||
edge_idx
|
||||
for edge_idx in (
|
||||
*self._graph.out_edge_indices(rx_idx),
|
||||
*self._graph.in_edge_indices(rx_idx),
|
||||
)
|
||||
if isinstance(self._graph.get_edge_data_by_index(edge_idx), RDMAConnection)
|
||||
]
|
||||
for edge_idx in rdma_edge_idxs:
|
||||
self._graph.remove_edge_from_index(edge_idx)
|
||||
|
||||
def remove_connection(self, conn: Connection) -> None:
|
||||
if (
|
||||
conn.source not in self._vertex_indices
|
||||
|
||||
@@ -67,12 +67,6 @@ class RequestEventLog(BaseCommand):
|
||||
since_idx: int
|
||||
|
||||
|
||||
class RequestSnapshot(BaseCommand):
|
||||
"""Ask the current master to send a State snapshot to this node."""
|
||||
|
||||
requester_node_id: NodeId
|
||||
|
||||
|
||||
class StartDownload(BaseCommand):
|
||||
target_node_id: NodeId
|
||||
shard_metadata: ShardMetadata
|
||||
@@ -112,7 +106,6 @@ DownloadCommand = StartDownload | DeleteDownload | CancelDownload
|
||||
Command = (
|
||||
TestCommand
|
||||
| RequestEventLog
|
||||
| RequestSnapshot
|
||||
| TextGeneration
|
||||
| ImageGeneration
|
||||
| ImageEdits
|
||||
|
||||
@@ -172,9 +172,6 @@ Event = (
|
||||
)
|
||||
|
||||
|
||||
TransientEvent = TaskAcknowledged | ChunkGenerated | TracesCollected | TracesMerged
|
||||
|
||||
|
||||
class IndexedEvent(FrozenModel):
|
||||
"""An event indexed by the master, with a globally unique index"""
|
||||
|
||||
@@ -198,11 +195,3 @@ class LocalForwarderEvent(FrozenModel):
|
||||
origin: SystemId
|
||||
session: SessionId
|
||||
event: Event
|
||||
|
||||
|
||||
class GlobalForwarderTransientEvent(FrozenModel):
|
||||
"""An unordered, non-durable event published to the cluster."""
|
||||
|
||||
origin: NodeId
|
||||
session: SessionId
|
||||
event: TransientEvent
|
||||
@@ -1,54 +0,0 @@
|
||||
"""Wire types for snapshot transfer between master and a joining node.
|
||||
|
||||
Snapshots can be tens of MB; the gossipsub message ceiling is around 1 MB.
|
||||
We slice the compressed snapshot body into chunks and publish each chunk on
|
||||
the SNAPSHOT_RESPONSES topic. The receiver collects chunks for its own
|
||||
`requester_node_id`, validates the SHA-256 of the reassembled body, and
|
||||
materialises the State.
|
||||
"""
|
||||
|
||||
import base64
|
||||
|
||||
from exo.shared.types.common import Id, NodeId, SessionId
|
||||
from exo.utils.pydantic_ext import FrozenModel
|
||||
|
||||
|
||||
class SnapshotTransferId(Id):
|
||||
"""Identifies a single snapshot transfer (one master response to one
|
||||
`RequestSnapshot`). Distinct transfers may interleave; the id lets
|
||||
receivers keep them apart."""
|
||||
|
||||
|
||||
class SnapshotChunk(FrozenModel):
|
||||
"""One slice of a snapshot in flight.
|
||||
|
||||
`data_b64` carries a base64-encoded slice of the zstd-compressed JSON
|
||||
dump of State. Concatenating the *decoded* bytes of all chunks for a
|
||||
`transfer_id` in order of `chunk_index` yields the full compressed
|
||||
body; `sha256_hex` is the SHA-256 of that decoded blob.
|
||||
|
||||
We use base64 explicitly because the topic layer JSON-encodes messages,
|
||||
and JSON can't carry raw binary. Helpers `from_data` / `data` keep the
|
||||
base64 detail at the boundaries.
|
||||
"""
|
||||
|
||||
transfer_id: SnapshotTransferId
|
||||
requester_node_id: NodeId
|
||||
session_id: SessionId
|
||||
schema_version: int
|
||||
last_event_applied_idx: int
|
||||
chunk_index: int
|
||||
total_chunks: int
|
||||
sha256_hex: str
|
||||
data_b64: str
|
||||
|
||||
@classmethod
|
||||
def from_data(cls, *, data: bytes, **kwargs: object) -> "SnapshotChunk":
|
||||
return cls(data_b64=base64.b64encode(data).decode("ascii"), **kwargs) # pyright: ignore[reportArgumentType]
|
||||
|
||||
@property
|
||||
def data(self) -> bytes:
|
||||
return base64.b64decode(self.data_b64)
|
||||
|
||||
|
||||
__all__ = ["SnapshotChunk", "SnapshotTransferId"]
|
||||
@@ -7,8 +7,7 @@ from pydantic.alias_generators import to_camel
|
||||
|
||||
from exo.shared.models.model_cards import ModelCard
|
||||
from exo.shared.topology import Topology, TopologySnapshot
|
||||
from exo.shared.types.chunks import InputImageChunk
|
||||
from exo.shared.types.common import CommandId, ModelId, NodeId
|
||||
from exo.shared.types.common import ModelId, NodeId
|
||||
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
|
||||
from exo.shared.types.profiling import (
|
||||
DiskUsage,
|
||||
@@ -43,16 +42,10 @@ class State(FrozenModel):
|
||||
strict=True,
|
||||
arbitrary_types_allowed=True,
|
||||
)
|
||||
# Bump when a State change makes older snapshots unsafe to restore.
|
||||
schema_version: int = Field(default=1, ge=1)
|
||||
|
||||
instances: Mapping[InstanceId, Instance] = {}
|
||||
runners: Mapping[RunnerId, RunnerStatus] = {}
|
||||
downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
|
||||
tasks: Mapping[TaskId, Task] = {}
|
||||
# Durable request input chunks for active image requests. Workers rebuild
|
||||
# local image caches from this state instead of reading events directly.
|
||||
input_chunks: Mapping[CommandId, Mapping[int, InputImageChunk]] = {}
|
||||
last_seen: Mapping[NodeId, datetime] = {}
|
||||
topology: Topology = Field(default_factory=Topology)
|
||||
last_event_applied_idx: int = Field(default=-1, ge=-1)
|
||||
@@ -73,8 +66,7 @@ class State(FrozenModel):
|
||||
instance_links: Mapping[InstanceLinkId, InstanceLink] = {}
|
||||
prefill_server_ports: Mapping[RunnerId, int] = {}
|
||||
|
||||
# User-added model cards. Workers can reconcile their on-disk custom card
|
||||
# cache from this state after snapshot bootstrap.
|
||||
# User-added model cards. Workers can reconcile their on-disk custom card cache
|
||||
custom_model_cards: Mapping[ModelId, ModelCard] = {}
|
||||
|
||||
@field_serializer("topology", mode="plain")
|
||||
|
||||
@@ -135,9 +135,9 @@ class TextGenerationTaskParams(BaseModel, frozen=True):
|
||||
prefill_endpoint: str | None = None
|
||||
|
||||
def with_card_sampling_defaults(self) -> "TextGenerationTaskParams":
|
||||
from exo.shared.models.model_cards import get_card
|
||||
from exo.shared.models import model_cards
|
||||
|
||||
card = get_card(self.model)
|
||||
card = model_cards.card_cache.get(self.model)
|
||||
if card is None:
|
||||
return self
|
||||
|
||||
|
||||
@@ -47,18 +47,6 @@ class OrderedBuffer[T]:
|
||||
logger.trace(f"Releasing event {ret}")
|
||||
return ret
|
||||
|
||||
def fast_forward_to(self, idx: int) -> None:
|
||||
"""Skip every event before idx.
|
||||
|
||||
Snapshot restore uses this after applying state that already includes
|
||||
events before idx. Any buffered or future event below idx is stale.
|
||||
"""
|
||||
if idx <= self.next_idx_to_release:
|
||||
return
|
||||
self.next_idx_to_release = idx
|
||||
for stale_idx in [i for i in self.store if i < idx]:
|
||||
del self.store[stale_idx]
|
||||
|
||||
|
||||
class MultiSourceBuffer[SourceId, T]:
|
||||
"""
|
||||
|
||||
@@ -29,10 +29,6 @@ class KeyedBackoff[K]:
|
||||
"""Return the number of recorded attempts for a key."""
|
||||
return self._attempts.get(key, 0)
|
||||
|
||||
def tracked_keys(self) -> set[K]:
|
||||
"""Return keys that currently have recorded backoff state."""
|
||||
return set(self._attempts) | set(self._last_time)
|
||||
|
||||
def reset(self, key: K) -> None:
|
||||
"""Reset backoff state for a key (e.g., on success)."""
|
||||
self._attempts.pop(key, None)
|
||||
|
||||
@@ -1,13 +0,0 @@
|
||||
from exo.utils.keyed_backoff import KeyedBackoff
|
||||
|
||||
|
||||
def test_tracked_keys_reports_and_resets_backoff_state() -> None:
|
||||
backoff = KeyedBackoff[str]()
|
||||
|
||||
backoff.record_attempt("instance-a")
|
||||
|
||||
assert backoff.tracked_keys() == {"instance-a"}
|
||||
|
||||
backoff.reset("instance-a")
|
||||
|
||||
assert backoff.tracked_keys() == set()
|
||||
@@ -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):
|
||||
|
||||
+62
-126
@@ -8,38 +8,30 @@ from loguru import logger
|
||||
|
||||
from exo.api.types import ImageEditsTaskParams
|
||||
from exo.download.download_utils import is_read_only_model_dir, resolve_existing_model
|
||||
from exo.routing.event_router import EventRouter
|
||||
from exo.routing.snapshot_receiver import SnapshotReceiver
|
||||
from exo.shared.apply import apply
|
||||
from exo.shared.constants import EXO_MAX_INSTANCE_RETRIES
|
||||
from exo.shared.models.model_cards import (
|
||||
ModelCard,
|
||||
ModelId,
|
||||
add_to_card_cache,
|
||||
delete_custom_card,
|
||||
)
|
||||
from exo.shared.models.model_cards import ModelId, card_cache
|
||||
from exo.shared.types.chunks import InputImageChunk
|
||||
from exo.shared.types.commands import (
|
||||
DeleteInstance,
|
||||
ForwarderCommand,
|
||||
ForwarderDownloadCommand,
|
||||
RequestSnapshot,
|
||||
StartDownload,
|
||||
)
|
||||
from exo.shared.types.common import CommandId, NodeId, SessionId, SystemId
|
||||
from exo.shared.types.common import CommandId, NodeId, SystemId
|
||||
from exo.shared.types.events import (
|
||||
Event,
|
||||
IndexedEvent,
|
||||
InputChunkReceived,
|
||||
InstanceDeleted,
|
||||
NodeDownloadProgress,
|
||||
NodeGatheredInfo,
|
||||
TaskCreated,
|
||||
TaskStatusUpdated,
|
||||
TopologyEdgeCreated,
|
||||
TopologyEdgeDeleted,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.multiaddr import Multiaddr
|
||||
from exo.shared.types.snapshots import SnapshotChunk
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.tasks import (
|
||||
CancelTask,
|
||||
@@ -65,21 +57,14 @@ from exo.utils.task_group import TaskGroup
|
||||
from exo.worker.plan import plan
|
||||
from exo.worker.runner.supervisor import RunnerSupervisor
|
||||
|
||||
_SNAPSHOT_FETCH_TIMEOUT_SECONDS = 30
|
||||
|
||||
|
||||
class Worker:
|
||||
def __init__(
|
||||
self,
|
||||
node_id: NodeId,
|
||||
session_id: SessionId,
|
||||
*,
|
||||
event_router: EventRouter,
|
||||
event_receiver: Receiver[IndexedEvent],
|
||||
event_sender: Sender[Event],
|
||||
transient_event_receiver: Receiver[TransientEvent],
|
||||
transient_event_sender: Sender[TransientEvent],
|
||||
snapshot_chunk_receiver: Receiver[SnapshotChunk],
|
||||
# This is for requesting updates. It doesn't need to be a general command sender right now,
|
||||
# but I think it's the correct way to be thinking about commands
|
||||
command_sender: Sender[ForwarderCommand],
|
||||
@@ -87,13 +72,8 @@ class Worker:
|
||||
api_port: int,
|
||||
):
|
||||
self.node_id: NodeId = node_id
|
||||
self.session_id: SessionId = session_id
|
||||
self.event_router = event_router
|
||||
self.event_receiver = event_receiver
|
||||
self.event_sender = event_sender
|
||||
self.transient_event_receiver = transient_event_receiver
|
||||
self.transient_event_sender = transient_event_sender
|
||||
self.snapshot_chunk_receiver = snapshot_chunk_receiver
|
||||
self.command_sender = command_sender
|
||||
self.download_command_sender = download_command_sender
|
||||
self.api_port = api_port
|
||||
@@ -113,7 +93,6 @@ class Worker:
|
||||
self._instance_backoff: KeyedBackoff[InstanceId] = KeyedBackoff(
|
||||
base=0.5, cap=10.0
|
||||
)
|
||||
self._synced_custom_cards: dict[ModelId, ModelCard] = {}
|
||||
self._stopped: anyio.Event = anyio.Event()
|
||||
|
||||
async def run(self):
|
||||
@@ -124,61 +103,23 @@ class Worker:
|
||||
|
||||
try:
|
||||
async with self._tg as tg:
|
||||
tg.start_soon(self._bootstrap_then_run, info_gatherer, info_recv)
|
||||
tg.start_soon(info_gatherer.run)
|
||||
tg.start_soon(self._forward_info, info_recv)
|
||||
tg.start_soon(self.plan_step)
|
||||
tg.start_soon(self._event_applier)
|
||||
tg.start_soon(self._poll_connection_updates)
|
||||
tg.start_soon(self._reconcile_custom_cards)
|
||||
|
||||
finally:
|
||||
# Actual shutdown code - waits for all tasks to complete before executing.
|
||||
logger.info("Stopping Worker")
|
||||
self.event_sender.close()
|
||||
self.transient_event_receiver.close()
|
||||
self.transient_event_sender.close()
|
||||
self.snapshot_chunk_receiver.close()
|
||||
self.command_sender.close()
|
||||
self.download_command_sender.close()
|
||||
for runner in self.runners.values():
|
||||
runner.shutdown()
|
||||
self._stopped.set()
|
||||
|
||||
async def _bootstrap_then_run(
|
||||
self, info_gatherer: InfoGatherer, info_recv: Receiver[GatheredInfo]
|
||||
) -> None:
|
||||
await self._fetch_snapshot()
|
||||
self._sync_input_views_from_state()
|
||||
self._tg.start_soon(info_gatherer.run)
|
||||
self._tg.start_soon(self._forward_info, info_recv)
|
||||
self._tg.start_soon(self.plan_step)
|
||||
self._tg.start_soon(self._event_applier)
|
||||
self._tg.start_soon(self._reconcile_instance_backoff)
|
||||
self._tg.start_soon(self._reconcile_custom_cards)
|
||||
self._tg.start_soon(self._poll_connection_updates)
|
||||
|
||||
async def _fetch_snapshot(self) -> None:
|
||||
receiver = SnapshotReceiver(self.node_id, self.session_id)
|
||||
await self.command_sender.send(
|
||||
ForwarderCommand(
|
||||
origin=self._system_id,
|
||||
command=RequestSnapshot(requester_node_id=self.node_id),
|
||||
)
|
||||
)
|
||||
|
||||
with anyio.move_on_after(_SNAPSHOT_FETCH_TIMEOUT_SECONDS):
|
||||
with self.snapshot_chunk_receiver as chunks:
|
||||
async for chunk in chunks:
|
||||
received = receiver.ingest(chunk)
|
||||
if received is None:
|
||||
continue
|
||||
self.state = received.state
|
||||
self.event_router.set_buffer_start(
|
||||
received.last_event_applied_idx + 1
|
||||
)
|
||||
logger.info(
|
||||
f"Worker bootstrapped from snapshot at idx "
|
||||
f"{received.last_event_applied_idx}"
|
||||
)
|
||||
return
|
||||
logger.info(
|
||||
"No snapshot received before timeout; falling back to full event-log replay"
|
||||
)
|
||||
|
||||
async def _forward_info(self, recv: Receiver[GatheredInfo]):
|
||||
with recv as info_stream:
|
||||
async for info in info_stream:
|
||||
@@ -193,69 +134,55 @@ class Worker:
|
||||
async def _event_applier(self):
|
||||
with self.event_receiver as events:
|
||||
async for event in events:
|
||||
if event.idx <= self.state.last_event_applied_idx:
|
||||
continue
|
||||
# 2. for each event, apply it to the state
|
||||
self.state = apply(self.state, event=event)
|
||||
self._sync_input_views_from_state()
|
||||
event = event.event
|
||||
|
||||
async def _reconcile_instance_backoff(self) -> None:
|
||||
while True:
|
||||
await anyio.sleep(1)
|
||||
self._reconcile_instance_backoff_once()
|
||||
if isinstance(event, InstanceDeleted):
|
||||
self._instance_backoff.reset(event.instance_id)
|
||||
|
||||
def _reconcile_instance_backoff_once(self) -> None:
|
||||
live_instances = set(self.state.instances)
|
||||
for instance_id in self._instance_backoff.tracked_keys():
|
||||
if instance_id not in live_instances:
|
||||
self._instance_backoff.reset(instance_id)
|
||||
# Buffer input image chunks for image editing
|
||||
if isinstance(event, InputChunkReceived):
|
||||
cmd_id = event.command_id
|
||||
if cmd_id not in self.input_chunk_buffer:
|
||||
self.input_chunk_buffer[cmd_id] = {}
|
||||
self.input_chunk_counts[cmd_id] = event.chunk.total_chunks
|
||||
|
||||
self.input_chunk_buffer[cmd_id][event.chunk.chunk_index] = (
|
||||
event.chunk
|
||||
)
|
||||
if (
|
||||
len(self.input_chunk_buffer[cmd_id])
|
||||
== self.input_chunk_counts[cmd_id]
|
||||
):
|
||||
per_image: defaultdict[int, list[InputImageChunk]] = (
|
||||
defaultdict(list)
|
||||
)
|
||||
for chunk in self.input_chunk_buffer[cmd_id].values():
|
||||
per_image[chunk.image_index].append(chunk)
|
||||
for chunks_for_image in per_image.values():
|
||||
sorted_chunks = sorted(
|
||||
chunks_for_image, key=lambda c: c.chunk_index
|
||||
)
|
||||
img = Base64Image("".join(c.data for c in sorted_chunks))
|
||||
self.image_cache[
|
||||
Base64ImageHash(
|
||||
hashlib.sha256(img.encode("ascii")).hexdigest()
|
||||
)
|
||||
] = img
|
||||
|
||||
async def _reconcile_custom_cards(self) -> None:
|
||||
while True:
|
||||
await anyio.sleep(1)
|
||||
await self._sync_custom_cards_from_state()
|
||||
target = dict(self.state.custom_model_cards)
|
||||
for model_id, card in target.items():
|
||||
if card_cache.get(model_id) == card:
|
||||
continue
|
||||
await card_cache.save(card)
|
||||
|
||||
async def _sync_custom_cards_from_state(self) -> None:
|
||||
target = dict(self.state.custom_model_cards)
|
||||
for model_id, card in target.items():
|
||||
if self._synced_custom_cards.get(model_id) == card:
|
||||
continue
|
||||
await card.save_to_custom_dir()
|
||||
add_to_card_cache(card)
|
||||
self._synced_custom_cards[model_id] = card
|
||||
|
||||
for model_id in list(self._synced_custom_cards):
|
||||
if model_id in target:
|
||||
continue
|
||||
await delete_custom_card(model_id)
|
||||
self._synced_custom_cards.pop(model_id, None)
|
||||
|
||||
def _sync_input_views_from_state(self) -> None:
|
||||
self.input_chunk_buffer = {
|
||||
command_id: dict(chunks)
|
||||
for command_id, chunks in self.state.input_chunks.items()
|
||||
}
|
||||
self.input_chunk_counts = {
|
||||
command_id: next(iter(chunks.values())).total_chunks
|
||||
for command_id, chunks in self.input_chunk_buffer.items()
|
||||
if chunks
|
||||
}
|
||||
|
||||
self.image_cache = {}
|
||||
for command_id, chunks in self.input_chunk_buffer.items():
|
||||
expected_chunks = self.input_chunk_counts.get(command_id)
|
||||
if expected_chunks is None or len(chunks) != expected_chunks:
|
||||
continue
|
||||
|
||||
per_image: defaultdict[int, list[InputImageChunk]] = defaultdict(list)
|
||||
for chunk in chunks.values():
|
||||
per_image[chunk.image_index].append(chunk)
|
||||
for chunks_for_image in per_image.values():
|
||||
sorted_chunks = sorted(chunks_for_image, key=lambda c: c.chunk_index)
|
||||
image = Base64Image("".join(chunk.data for chunk in sorted_chunks))
|
||||
self.image_cache[
|
||||
Base64ImageHash(hashlib.sha256(image.encode("ascii")).hexdigest())
|
||||
] = image
|
||||
for card in await card_cache.list_all():
|
||||
if card.model_id not in target:
|
||||
await card_cache.pop(card.model_id)
|
||||
|
||||
async def plan_step(self):
|
||||
while True:
|
||||
@@ -267,7 +194,7 @@ class Worker:
|
||||
self.state.instances,
|
||||
self.state.runners,
|
||||
self.state.tasks,
|
||||
self.state.input_chunks,
|
||||
self.input_chunk_buffer,
|
||||
self.image_cache,
|
||||
self._instance_backoff,
|
||||
self._download_backoff,
|
||||
@@ -399,9 +326,15 @@ class Worker:
|
||||
advanced_params=task.task_params.advanced_params,
|
||||
),
|
||||
)
|
||||
# Cleanup buffers
|
||||
if cmd_id in self.input_chunk_buffer:
|
||||
del self.input_chunk_buffer[cmd_id]
|
||||
if cmd_id in self.input_chunk_counts:
|
||||
del self.input_chunk_counts[cmd_id]
|
||||
await self._start_runner_task(modified_task)
|
||||
|
||||
case TextGeneration() if task.task_params.image_hashes:
|
||||
cmd_id = task.command_id
|
||||
resolved_images = [
|
||||
self.image_cache[h]
|
||||
for _, h in sorted(task.task_params.image_hashes.items())
|
||||
@@ -413,6 +346,10 @@ class Worker:
|
||||
)
|
||||
}
|
||||
)
|
||||
if cmd_id in self.input_chunk_buffer:
|
||||
del self.input_chunk_buffer[cmd_id]
|
||||
if cmd_id in self.input_chunk_counts:
|
||||
del self.input_chunk_counts[cmd_id]
|
||||
await self._start_runner_task(modified_task)
|
||||
case LoadModel(instance_id=instance_id):
|
||||
if (instance := self.state.instances.get(instance_id)) is not None:
|
||||
@@ -438,7 +375,6 @@ class Worker:
|
||||
runner = RunnerSupervisor.create(
|
||||
bound_instance=task.bound_instance,
|
||||
event_sender=self.event_sender.clone(),
|
||||
transient_event_sender=self.transient_event_sender.clone(),
|
||||
)
|
||||
self.runners[task.bound_instance.bound_runner_id] = runner
|
||||
self._tg.start_soon(runner.run)
|
||||
|
||||
@@ -138,8 +138,10 @@ class SequentialGenerator(Engine):
|
||||
def agree_on_tasks(self) -> None:
|
||||
"""Agree between all ranks about the task ordering (some may have received in different order or not at all)."""
|
||||
agreed, different = mx_all_gather_tasks(self._maybe_queue, self.group)
|
||||
self._queue.extend(task for task in self._maybe_queue if task in agreed)
|
||||
self._maybe_queue = [task for task in self._maybe_queue if task in different]
|
||||
# Extend from `agreed` (sorted by task_id on all ranks) to guarantee every
|
||||
# rank enqueues tasks in the same order, preventing TP collective deadlocks.
|
||||
self._queue.extend(agreed)
|
||||
self._maybe_queue = list(different)
|
||||
|
||||
def agree_on_cancellations(self) -> None:
|
||||
"""Agree between all ranks about which tasks to cancel."""
|
||||
@@ -197,9 +199,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:
|
||||
@@ -368,8 +375,10 @@ class BatchGenerator(Engine):
|
||||
def agree_on_tasks(self) -> None:
|
||||
"""Agree between all ranks about the task ordering (some may have received in different order or not at all)."""
|
||||
agreed, different = mx_all_gather_tasks(self._maybe_queue, self.group)
|
||||
self._queue.extend(task for task in self._maybe_queue if task in agreed)
|
||||
self._maybe_queue = [task for task in self._maybe_queue if task in different]
|
||||
# Extend from `agreed` (sorted by task_id on all ranks) to guarantee every
|
||||
# rank enqueues tasks in the same order, preventing TP collective deadlocks.
|
||||
self._queue.extend(agreed)
|
||||
self._maybe_queue = list(different)
|
||||
|
||||
def agree_on_cancellations(self) -> None:
|
||||
"""Agree between all ranks about which tasks to cancel."""
|
||||
@@ -449,7 +458,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,
|
||||
|
||||
@@ -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))
|
||||
@@ -19,7 +19,6 @@ from exo.shared.types.events import (
|
||||
RunnerStatusUpdated,
|
||||
TaskAcknowledged,
|
||||
TaskStatusUpdated,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.tasks import (
|
||||
CANCEL_ALL_TASKS,
|
||||
@@ -59,7 +58,6 @@ class RunnerSupervisor:
|
||||
_ev_recv: MpReceiver[Event]
|
||||
_task_sender: MpSender[Task]
|
||||
_event_sender: Sender[Event]
|
||||
_transient_event_sender: Sender[TransientEvent]
|
||||
_cancel_sender: MpSender[TaskId]
|
||||
_tg: TaskGroup = field(default_factory=TaskGroup, init=False)
|
||||
status: RunnerStatus = field(default_factory=RunnerIdle, init=False)
|
||||
@@ -77,7 +75,6 @@ class RunnerSupervisor:
|
||||
*,
|
||||
bound_instance: BoundInstance,
|
||||
event_sender: Sender[Event],
|
||||
transient_event_sender: Sender[TransientEvent],
|
||||
initialize_timeout: float = 400,
|
||||
) -> Self:
|
||||
ev_send, ev_recv = mp_channel[Event]()
|
||||
@@ -107,7 +104,6 @@ class RunnerSupervisor:
|
||||
_task_sender=task_sender,
|
||||
_cancel_sender=cancel_sender,
|
||||
_event_sender=event_sender,
|
||||
_transient_event_sender=transient_event_sender,
|
||||
)
|
||||
|
||||
return self
|
||||
@@ -128,8 +124,6 @@ class RunnerSupervisor:
|
||||
self._task_sender.close()
|
||||
with contextlib.suppress(ClosedResourceError):
|
||||
self._event_sender.close()
|
||||
with contextlib.suppress(ClosedResourceError):
|
||||
self._transient_event_sender.close()
|
||||
with contextlib.suppress(ClosedResourceError):
|
||||
self._cancel_sender.send(CANCEL_ALL_TASKS)
|
||||
with contextlib.suppress(ClosedResourceError):
|
||||
@@ -241,10 +235,7 @@ class RunnerSupervisor:
|
||||
)
|
||||
self.in_progress.pop(event.task_id, None)
|
||||
self.completed.add(event.task_id)
|
||||
if isinstance(event, ChunkGenerated):
|
||||
await self._transient_event_sender.send(event)
|
||||
else:
|
||||
await self._event_sender.send(event)
|
||||
await self._event_sender.send(event)
|
||||
except (ClosedResourceError, BrokenResourceError) as e:
|
||||
await self._check_runner(e)
|
||||
finally:
|
||||
@@ -284,7 +275,7 @@ class RunnerSupervisor:
|
||||
for task in self.in_progress.values():
|
||||
if isinstance(task, (TextGeneration, ImageGeneration, ImageEdits)):
|
||||
with anyio.CancelScope(shield=True):
|
||||
await self._transient_event_sender.send(
|
||||
await self._event_sender.send(
|
||||
ChunkGenerated(
|
||||
command_id=task.command_id,
|
||||
chunk=ErrorChunk(
|
||||
|
||||
@@ -16,7 +16,7 @@ from exo.download.download_utils import (
|
||||
fetch_file_list_with_cache,
|
||||
resolve_model_dir,
|
||||
)
|
||||
from exo.shared.models.model_cards import ModelCard, ModelId, get_model_cards
|
||||
from exo.shared.models.model_cards import ModelCard, ModelId, card_cache
|
||||
from exo.worker.engines.mlx.utils_mlx import (
|
||||
get_eos_token_ids_for_model,
|
||||
load_tokenizer_for_model_id,
|
||||
@@ -76,7 +76,7 @@ def get_test_models() -> list[ModelCard]:
|
||||
"""Get a representative sample of models to test."""
|
||||
# Pick one model from each family to test
|
||||
families: dict[str, ModelCard] = {}
|
||||
for card in asyncio.run(get_model_cards()):
|
||||
for card in asyncio.run(card_cache.list_all()):
|
||||
# Extract family name (e.g., "llama-3.1" from "llama-3.1-8b")
|
||||
parts = card.model_id.short().split("-")
|
||||
family = "-".join(parts[:2]) if len(parts) >= 2 else parts[0]
|
||||
@@ -298,7 +298,7 @@ async def test_tokenizer_special_tokens(model_card: ModelCard) -> None:
|
||||
async def test_kimi_tokenizer_specifically():
|
||||
"""Test Kimi tokenizer with its specific patches and quirks."""
|
||||
kimi_models = [
|
||||
card for card in await get_model_cards() if "kimi" in card.model_id.lower()
|
||||
card for card in await card_cache.list_all() if "kimi" in card.model_id.lower()
|
||||
]
|
||||
|
||||
if not kimi_models:
|
||||
@@ -350,7 +350,7 @@ async def test_glm_tokenizer_specifically():
|
||||
|
||||
glm_model_cards = [
|
||||
card
|
||||
for card in await get_model_cards()
|
||||
for card in await card_cache.list_all()
|
||||
if contains(card, "glm")
|
||||
and not contains(card, "-5")
|
||||
and not contains(card, "4.7")
|
||||
|
||||
@@ -7,12 +7,7 @@ import pytest
|
||||
from exo.shared.models.model_cards import ModelId
|
||||
from exo.shared.types.chunks import ErrorChunk
|
||||
from exo.shared.types.common import CommandId, NodeId
|
||||
from exo.shared.types.events import (
|
||||
ChunkGenerated,
|
||||
Event,
|
||||
RunnerStatusUpdated,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.events import ChunkGenerated, Event, RunnerStatusUpdated
|
||||
from exo.shared.types.tasks import Task, TaskId, TextGeneration
|
||||
from exo.shared.types.text_generation import (
|
||||
InputMessage,
|
||||
@@ -48,7 +43,6 @@ class _DeadProcess:
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_runner_emits_error_chunk_for_inflight_text_generation() -> None:
|
||||
event_sender, event_receiver = channel[Event]()
|
||||
transient_sender, transient_receiver = channel[TransientEvent]()
|
||||
task_sender, _ = mp_channel[Task]()
|
||||
cancel_sender, _ = mp_channel[TaskId]()
|
||||
_, ev_recv = mp_channel[Event]()
|
||||
@@ -68,7 +62,6 @@ async def test_check_runner_emits_error_chunk_for_inflight_text_generation() ->
|
||||
_ev_recv=ev_recv,
|
||||
_task_sender=task_sender,
|
||||
_event_sender=event_sender,
|
||||
_transient_event_sender=transient_sender,
|
||||
_cancel_sender=cancel_sender,
|
||||
)
|
||||
|
||||
@@ -88,7 +81,7 @@ async def test_check_runner_emits_error_chunk_for_inflight_text_generation() ->
|
||||
|
||||
await supervisor._check_runner(RuntimeError("boom")) # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
got_chunk = await transient_receiver.receive()
|
||||
got_chunk = await event_receiver.receive()
|
||||
got_status = await event_receiver.receive()
|
||||
|
||||
assert isinstance(got_chunk, ChunkGenerated)
|
||||
@@ -100,57 +93,5 @@ async def test_check_runner_emits_error_chunk_for_inflight_text_generation() ->
|
||||
assert isinstance(got_status.runner_status, RunnerFailed)
|
||||
|
||||
event_sender.close()
|
||||
transient_sender.close()
|
||||
with anyio.move_on_after(0.1):
|
||||
await event_receiver.aclose()
|
||||
with anyio.move_on_after(0.1):
|
||||
await transient_receiver.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forward_events_routes_generated_chunks_to_transient() -> None:
|
||||
event_sender, event_receiver = channel[Event]()
|
||||
transient_sender, transient_receiver = channel[TransientEvent]()
|
||||
task_sender, _ = mp_channel[Task]()
|
||||
cancel_sender, _ = mp_channel[TaskId]()
|
||||
ev_send, ev_recv = mp_channel[Event]()
|
||||
|
||||
bound_instance: BoundInstance = get_bound_mlx_ring_instance(
|
||||
instance_id=InstanceId("instance-a"),
|
||||
model_id=ModelId("mlx-community/Llama-3.2-1B-Instruct-4bit"),
|
||||
runner_id=RunnerId("runner-a"),
|
||||
node_id=NodeId("node-a"),
|
||||
)
|
||||
|
||||
supervisor = RunnerSupervisor(
|
||||
shard_metadata=bound_instance.bound_shard,
|
||||
bound_instance=bound_instance,
|
||||
runner_process=cast("mp.Process", cast(object, _DeadProcess())),
|
||||
initialize_timeout=400,
|
||||
_ev_recv=ev_recv,
|
||||
_task_sender=task_sender,
|
||||
_event_sender=event_sender,
|
||||
_transient_event_sender=transient_sender,
|
||||
_cancel_sender=cancel_sender,
|
||||
)
|
||||
|
||||
command_id = CommandId("cmd-a")
|
||||
chunk = ChunkGenerated(
|
||||
command_id=command_id,
|
||||
chunk=ErrorChunk(
|
||||
model=bound_instance.bound_shard.model_card.model_id,
|
||||
error_message="test chunk",
|
||||
),
|
||||
)
|
||||
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(supervisor._forward_events) # pyright: ignore[reportPrivateUsage]
|
||||
ev_send.send(chunk)
|
||||
assert await transient_receiver.receive() == chunk
|
||||
tg.cancel_scope.cancel()
|
||||
|
||||
assert event_receiver.collect() == []
|
||||
|
||||
event_sender.close()
|
||||
transient_sender.close()
|
||||
ev_send.close()
|
||||
@@ -1,78 +0,0 @@
|
||||
# pyright: reportPrivateUsage=false
|
||||
|
||||
import pytest
|
||||
|
||||
import exo.worker.main as worker_main
|
||||
from exo.shared.models.model_cards import ModelCard, ModelTask
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.state import State
|
||||
from exo.worker.main import Worker
|
||||
|
||||
|
||||
def _model_card(model_id: ModelId) -> ModelCard:
|
||||
return ModelCard(
|
||||
model_id=model_id,
|
||||
n_layers=1,
|
||||
storage_size=Memory.from_bytes(1),
|
||||
hidden_size=1,
|
||||
supports_tensor=True,
|
||||
tasks=[ModelTask.TextGeneration],
|
||||
)
|
||||
|
||||
|
||||
def _worker() -> Worker:
|
||||
worker = object.__new__(Worker)
|
||||
worker.state = State()
|
||||
worker._synced_custom_cards = {}
|
||||
return worker
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_syncs_custom_cards_from_state(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
saved: list[ModelId] = []
|
||||
cached: list[ModelId] = []
|
||||
|
||||
async def save_to_custom_dir(card: ModelCard) -> None:
|
||||
saved.append(card.model_id)
|
||||
|
||||
def add_to_card_cache(card: ModelCard) -> None:
|
||||
cached.append(card.model_id)
|
||||
|
||||
monkeypatch.setattr(ModelCard, "save_to_custom_dir", save_to_custom_dir)
|
||||
monkeypatch.setattr(worker_main, "add_to_card_cache", add_to_card_cache)
|
||||
|
||||
card = _model_card(ModelId("custom/model"))
|
||||
worker = _worker()
|
||||
worker.state = State(custom_model_cards={card.model_id: card})
|
||||
|
||||
await worker._sync_custom_cards_from_state()
|
||||
await worker._sync_custom_cards_from_state()
|
||||
|
||||
assert saved == [card.model_id]
|
||||
assert cached == [card.model_id]
|
||||
assert worker._synced_custom_cards == {card.model_id: card}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_deletes_custom_cards_missing_from_state(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
deleted: list[ModelId] = []
|
||||
|
||||
async def delete_custom_card(model_id: ModelId) -> bool:
|
||||
deleted.append(model_id)
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(worker_main, "delete_custom_card", delete_custom_card)
|
||||
|
||||
card = _model_card(ModelId("custom/model"))
|
||||
worker = _worker()
|
||||
worker._synced_custom_cards = {card.model_id: card}
|
||||
|
||||
await worker._sync_custom_cards_from_state()
|
||||
|
||||
assert deleted == [card.model_id]
|
||||
assert worker._synced_custom_cards == {}
|
||||
@@ -1,36 +0,0 @@
|
||||
# pyright: reportPrivateUsage=false
|
||||
|
||||
from exo.shared.types.common import ModelId, NodeId
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.worker.instances import InstanceId, MlxRingInstance
|
||||
from exo.shared.types.worker.runners import ShardAssignments
|
||||
from exo.utils.keyed_backoff import KeyedBackoff
|
||||
from exo.worker.main import Worker
|
||||
|
||||
|
||||
def _make_instance(instance_id: InstanceId) -> MlxRingInstance:
|
||||
return MlxRingInstance(
|
||||
instance_id=instance_id,
|
||||
shard_assignments=ShardAssignments(
|
||||
model_id=ModelId("test-model"),
|
||||
node_to_runner={},
|
||||
runner_to_shard={},
|
||||
),
|
||||
hosts_by_node={NodeId("node-1"): []},
|
||||
ephemeral_port=1,
|
||||
)
|
||||
|
||||
|
||||
def test_worker_reconciles_instance_backoff_from_state() -> None:
|
||||
live_instance_id = InstanceId("inst-live")
|
||||
deleted_instance_id = InstanceId("inst-deleted")
|
||||
worker = object.__new__(Worker)
|
||||
worker.state = State(instances={live_instance_id: _make_instance(live_instance_id)})
|
||||
worker._instance_backoff = KeyedBackoff[InstanceId]()
|
||||
worker._instance_backoff.record_attempt(live_instance_id)
|
||||
worker._instance_backoff.record_attempt(deleted_instance_id)
|
||||
|
||||
worker._reconcile_instance_backoff_once()
|
||||
|
||||
assert worker._instance_backoff.attempts(live_instance_id) == 1
|
||||
assert worker._instance_backoff.attempts(deleted_instance_id) == 0
|
||||
@@ -1,130 +0,0 @@
|
||||
# pyright: reportPrivateUsage=false
|
||||
|
||||
import hashlib
|
||||
|
||||
import anyio
|
||||
import pytest
|
||||
import zstandard
|
||||
|
||||
from exo.routing.event_router import EventRouter
|
||||
from exo.shared.types.commands import (
|
||||
ForwarderCommand,
|
||||
ForwarderDownloadCommand,
|
||||
RequestSnapshot,
|
||||
)
|
||||
from exo.shared.types.common import NodeId, SessionId
|
||||
from exo.shared.types.events import (
|
||||
Event,
|
||||
GlobalForwarderEvent,
|
||||
IndexedEvent,
|
||||
LocalForwarderEvent,
|
||||
TestEvent,
|
||||
TransientEvent,
|
||||
)
|
||||
from exo.shared.types.snapshots import SnapshotChunk, SnapshotTransferId
|
||||
from exo.shared.types.state import State
|
||||
from exo.utils.channels import Receiver, Sender, channel
|
||||
from exo.worker.main import Worker
|
||||
|
||||
|
||||
def _snapshot_chunk(
|
||||
state: State, *, requester_node_id: NodeId, session_id: SessionId
|
||||
) -> SnapshotChunk:
|
||||
body = zstandard.ZstdCompressor().compress(state.model_dump_json().encode("utf-8"))
|
||||
return SnapshotChunk.from_data(
|
||||
data=body,
|
||||
transfer_id=SnapshotTransferId("transfer-1"),
|
||||
requester_node_id=requester_node_id,
|
||||
session_id=session_id,
|
||||
schema_version=state.schema_version,
|
||||
last_event_applied_idx=state.last_event_applied_idx,
|
||||
chunk_index=0,
|
||||
total_chunks=1,
|
||||
sha256_hex=hashlib.sha256(body).hexdigest(),
|
||||
)
|
||||
|
||||
|
||||
def _worker(
|
||||
node_id: NodeId, session_id: SessionId
|
||||
) -> tuple[
|
||||
Worker,
|
||||
EventRouter,
|
||||
Receiver[ForwarderCommand],
|
||||
Sender[SnapshotChunk],
|
||||
Sender[IndexedEvent],
|
||||
]:
|
||||
router_command_sender, _router_command_receiver = channel[ForwarderCommand]()
|
||||
_global_event_sender, global_event_receiver = channel[GlobalForwarderEvent]()
|
||||
local_event_sender, _local_event_receiver = channel[LocalForwarderEvent]()
|
||||
event_router = EventRouter(
|
||||
session_id=session_id,
|
||||
command_sender=router_command_sender,
|
||||
external_inbound=global_event_receiver,
|
||||
external_outbound=local_event_sender,
|
||||
)
|
||||
|
||||
event_sender, event_receiver = channel[IndexedEvent]()
|
||||
local_event_output_sender, _local_event_output_receiver = channel[Event]()
|
||||
transient_sender, transient_receiver = channel[TransientEvent]()
|
||||
command_sender, command_receiver = channel[ForwarderCommand]()
|
||||
download_command_sender, _download_command_receiver = channel[
|
||||
ForwarderDownloadCommand
|
||||
]()
|
||||
snapshot_sender, snapshot_receiver = channel[SnapshotChunk]()
|
||||
|
||||
worker = Worker(
|
||||
node_id,
|
||||
session_id,
|
||||
event_router=event_router,
|
||||
event_receiver=event_receiver,
|
||||
event_sender=local_event_output_sender,
|
||||
transient_event_receiver=transient_receiver,
|
||||
transient_event_sender=transient_sender,
|
||||
snapshot_chunk_receiver=snapshot_receiver,
|
||||
command_sender=command_sender,
|
||||
download_command_sender=download_command_sender,
|
||||
api_port=52415,
|
||||
)
|
||||
return worker, event_router, command_receiver, snapshot_sender, event_sender
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_fetch_snapshot_applies_state_and_fast_forwards_router() -> None:
|
||||
node_id = NodeId("worker")
|
||||
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
|
||||
worker, event_router, command_receiver, snapshot_sender, _event_sender = _worker(
|
||||
node_id, session_id
|
||||
)
|
||||
state = State(last_event_applied_idx=7)
|
||||
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(worker._fetch_snapshot)
|
||||
command = await command_receiver.receive()
|
||||
assert isinstance(command.command, RequestSnapshot)
|
||||
assert command.command.requester_node_id == node_id
|
||||
|
||||
await snapshot_sender.send(
|
||||
_snapshot_chunk(state, requester_node_id=node_id, session_id=session_id)
|
||||
)
|
||||
|
||||
assert worker.state.last_event_applied_idx == 7
|
||||
assert event_router.event_buffer.next_idx_to_release == 8
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_worker_event_applier_ignores_events_covered_by_snapshot() -> None:
|
||||
node_id = NodeId("worker")
|
||||
session_id = SessionId(master_node_id=NodeId("master"), election_clock=1)
|
||||
worker, _event_router, _command_receiver, _snapshot_sender, event_sender = _worker(
|
||||
node_id, session_id
|
||||
)
|
||||
worker.state = State(last_event_applied_idx=7)
|
||||
|
||||
async with anyio.create_task_group() as tg:
|
||||
tg.start_soon(worker._event_applier)
|
||||
await event_sender.send(IndexedEvent(idx=7, event=TestEvent()))
|
||||
await event_sender.send(IndexedEvent(idx=8, event=TestEvent()))
|
||||
|
||||
while worker.state.last_event_applied_idx != 8:
|
||||
await anyio.sleep(0.001)
|
||||
tg.cancel_scope.cancel()
|
||||
Reference in new issue
Block a user