dynamic type registry
This commit is contained in:
@@ -46,8 +46,8 @@ from exo.shared.types.api import (
|
||||
)
|
||||
from exo.shared.types.chunks import TokenChunk
|
||||
from exo.shared.types.commands import (
|
||||
BaseCommand,
|
||||
ChatCompletion,
|
||||
Command,
|
||||
CreateInstance,
|
||||
DeleteInstance,
|
||||
ForwarderCommand,
|
||||
@@ -66,6 +66,7 @@ from exo.shared.types.models import ModelId, ModelMetadata
|
||||
from exo.shared.types.state import State
|
||||
from exo.shared.types.tasks import ChatCompletionTaskParams
|
||||
from exo.shared.types.worker.instances import (
|
||||
BaseInstance,
|
||||
Instance,
|
||||
InstanceId,
|
||||
InstanceMeta,
|
||||
@@ -317,7 +318,7 @@ class API:
|
||||
sharding: Sharding = Sharding.Pipeline,
|
||||
instance_meta: InstanceMeta = InstanceMeta.MlxRing,
|
||||
min_nodes: int = 1,
|
||||
) -> Instance:
|
||||
) -> BaseInstance:
|
||||
model_meta = await resolve_model_meta(model_id)
|
||||
|
||||
try:
|
||||
@@ -449,7 +450,7 @@ class API:
|
||||
model_id=card.model_id,
|
||||
sharding=sharding,
|
||||
instance_meta=instance_meta,
|
||||
instance=instance,
|
||||
instance=cast(Instance, instance),
|
||||
memory_delta_by_node=memory_delta_by_node or None,
|
||||
error=None,
|
||||
)
|
||||
@@ -458,7 +459,7 @@ class API:
|
||||
|
||||
return PlacementPreviewResponse(previews=previews)
|
||||
|
||||
def get_instance(self, instance_id: InstanceId) -> Instance:
|
||||
def get_instance(self, instance_id: InstanceId) -> BaseInstance:
|
||||
if instance_id not in self.state.instances:
|
||||
raise HTTPException(status_code=404, detail="Instance not found")
|
||||
return self.state.instances[instance_id]
|
||||
@@ -808,7 +809,7 @@ class API:
|
||||
if message.clock > self.last_completed_election:
|
||||
self.paused = True
|
||||
|
||||
async def _send(self, command: Command):
|
||||
async def _send(self, command: BaseCommand):
|
||||
while self.paused:
|
||||
await self.paused_ev.wait()
|
||||
await self.command_sender.send(
|
||||
|
||||
@@ -216,7 +216,8 @@ class Master:
|
||||
IndexedEvent(idx=i, event=self._event_log[i])
|
||||
)
|
||||
case _:
|
||||
# Plugin-managed commands are handled above
|
||||
# Plugin commands should be handled by registry above;
|
||||
# this is a safety fallback for unhandled commands
|
||||
pass
|
||||
for event in generated_events:
|
||||
await self.event_sender.send(event)
|
||||
|
||||
@@ -26,7 +26,7 @@ from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.models import ModelId
|
||||
from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo
|
||||
from exo.shared.types.worker.instances import (
|
||||
Instance,
|
||||
BaseInstance,
|
||||
InstanceId,
|
||||
InstanceMeta,
|
||||
MlxJacclInstance,
|
||||
@@ -43,8 +43,8 @@ def random_ephemeral_port() -> int:
|
||||
def add_instance_to_placements(
|
||||
command: CreateInstance,
|
||||
topology: Topology,
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
) -> Mapping[InstanceId, Instance]:
|
||||
current_instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> Mapping[InstanceId, BaseInstance]:
|
||||
# TODO: validate against topology
|
||||
|
||||
return {**current_instances, command.instance.instance_id: command.instance}
|
||||
@@ -53,10 +53,10 @@ def add_instance_to_placements(
|
||||
def place_instance(
|
||||
command: PlaceInstance,
|
||||
topology: Topology,
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
current_instances: Mapping[InstanceId, BaseInstance],
|
||||
node_memory: Mapping[NodeId, MemoryUsage],
|
||||
node_network: Mapping[NodeId, NodeNetworkInfo],
|
||||
) -> dict[InstanceId, Instance]:
|
||||
) -> dict[InstanceId, BaseInstance]:
|
||||
cycles = topology.get_cycles()
|
||||
candidate_cycles = list(filter(lambda it: len(it) >= command.min_nodes, cycles))
|
||||
cycles_with_sufficient_memory = filter_cycles_by_memory(
|
||||
@@ -159,19 +159,14 @@ def place_instance(
|
||||
hosts_by_node=hosts_by_node,
|
||||
ephemeral_port=ephemeral_port,
|
||||
)
|
||||
case _:
|
||||
# Plugin-managed instance types have their own placement functions
|
||||
raise ValueError(
|
||||
f"Instance type {command.instance_meta} must use plugin placement"
|
||||
)
|
||||
|
||||
return target_instances
|
||||
|
||||
|
||||
def delete_instance(
|
||||
command: DeleteInstance,
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
) -> dict[InstanceId, Instance]:
|
||||
current_instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> dict[InstanceId, BaseInstance]:
|
||||
target_instances = dict(deepcopy(current_instances))
|
||||
if command.instance_id in target_instances:
|
||||
del target_instances[command.instance_id]
|
||||
@@ -180,8 +175,8 @@ def delete_instance(
|
||||
|
||||
|
||||
def get_transition_events(
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
target_instances: Mapping[InstanceId, Instance],
|
||||
current_instances: Mapping[InstanceId, BaseInstance],
|
||||
target_instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> Sequence[Event]:
|
||||
events: list[Event] = []
|
||||
|
||||
|
||||
+10
-24
@@ -4,11 +4,13 @@ This module provides the plugin architecture for extending exo with custom
|
||||
workload types (simulations, ML frameworks, etc.) without modifying core code.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from exo.plugins.base import EXOPlugin, PluginCommand, PluginInstance
|
||||
from exo.plugins.registry import PluginRegistry, discover_plugins
|
||||
from exo.plugins.base import EXOPlugin, PluginCommand, PluginInstance
|
||||
from exo.plugins.registry import PluginRegistry, discover_plugins
|
||||
from exo.plugins.type_registry import (
|
||||
command_registry,
|
||||
event_registry,
|
||||
instance_registry,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"EXOPlugin",
|
||||
@@ -16,23 +18,7 @@ __all__ = [
|
||||
"PluginInstance",
|
||||
"PluginRegistry",
|
||||
"discover_plugins",
|
||||
"command_registry",
|
||||
"event_registry",
|
||||
"instance_registry",
|
||||
]
|
||||
|
||||
|
||||
def __getattr__(name: str) -> Any: # pyright: ignore[reportAny]
|
||||
"""Lazy import to avoid circular dependencies."""
|
||||
if name in ("EXOPlugin", "PluginCommand", "PluginInstance"):
|
||||
from exo.plugins.base import EXOPlugin, PluginCommand, PluginInstance
|
||||
|
||||
return {
|
||||
"EXOPlugin": EXOPlugin,
|
||||
"PluginCommand": PluginCommand,
|
||||
"PluginInstance": PluginInstance,
|
||||
}[name]
|
||||
if name in ("PluginRegistry", "discover_plugins"):
|
||||
from exo.plugins.registry import PluginRegistry, discover_plugins
|
||||
|
||||
return {"PluginRegistry": PluginRegistry, "discover_plugins": discover_plugins}[
|
||||
name
|
||||
]
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
@@ -15,7 +15,7 @@ from exo.utils.pydantic_ext import TaggedModel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from exo.shared.topology import Topology
|
||||
from exo.shared.types.worker.instances import BoundInstance, Instance
|
||||
from exo.shared.types.worker.instances import BaseInstance, BoundInstance
|
||||
from exo.utils.channels import MpReceiver, MpSender
|
||||
from exo.worker.runner.runner_supervisor import RunnerSupervisor
|
||||
|
||||
@@ -113,7 +113,7 @@ class EXOPlugin(ABC):
|
||||
self,
|
||||
command: Any, # pyright: ignore[reportAny]
|
||||
topology: "Topology",
|
||||
current_instances: Mapping[InstanceId, "Instance"],
|
||||
current_instances: Mapping[InstanceId, "BaseInstance"],
|
||||
) -> Sequence[Event]:
|
||||
"""Process a command and return events to emit.
|
||||
|
||||
@@ -140,7 +140,7 @@ class EXOPlugin(ABC):
|
||||
def plan_task(
|
||||
self,
|
||||
runners: Mapping[RunnerId, "RunnerSupervisor"],
|
||||
instances: Mapping[InstanceId, "Instance"],
|
||||
instances: Mapping[InstanceId, "BaseInstance"],
|
||||
) -> Task | None:
|
||||
"""Plan the next task for plugin instances.
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
|
||||
from exo.shared.types.commands import Command
|
||||
from exo.shared.types.commands import BaseCommand
|
||||
from exo.shared.types.common import NodeId
|
||||
from exo.shared.types.state import State
|
||||
|
||||
@@ -17,5 +17,5 @@ class PluginContext:
|
||||
"""
|
||||
|
||||
state: State
|
||||
send_command: Callable[[Command], Awaitable[None]]
|
||||
send_command: Callable[[BaseCommand], Awaitable[None]]
|
||||
node_id: NodeId
|
||||
|
||||
@@ -1,32 +1,15 @@
|
||||
"""FLASH Plugin - MPI-based simulation support for Exo."""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
# Import types directly (these don't cause circular imports)
|
||||
from exo.plugins.implementations.flash.plugin import FLASHPlugin
|
||||
from exo.plugins.implementations.flash.types import (
|
||||
FLASHInstance,
|
||||
LaunchFLASH,
|
||||
StopFLASH,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from exo.plugins.implementations.flash.plugin import FLASHPlugin
|
||||
|
||||
__all__ = ["FLASHPlugin", "FLASHInstance", "LaunchFLASH", "StopFLASH", "register"]
|
||||
|
||||
|
||||
def register() -> "FLASHPlugin":
|
||||
def register() -> FLASHPlugin:
|
||||
"""Entry point for plugin discovery."""
|
||||
# Lazy import to avoid circular imports during module loading
|
||||
from exo.plugins.implementations.flash.plugin import FLASHPlugin
|
||||
|
||||
return FLASHPlugin()
|
||||
|
||||
|
||||
# For backwards compatibility, allow importing FLASHPlugin from this module
|
||||
def __getattr__(name: str) -> Any: # pyright: ignore[reportAny]
|
||||
if name == "FLASHPlugin":
|
||||
from exo.plugins.implementations.flash.plugin import FLASHPlugin
|
||||
|
||||
return FLASHPlugin
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
@@ -11,7 +11,7 @@ from exo.shared.types.common import Host, NodeId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.models import ModelId, ModelMetadata
|
||||
from exo.shared.types.topology import SocketConnection
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId
|
||||
from exo.shared.types.worker.instances import BaseInstance, InstanceId
|
||||
from exo.shared.types.worker.runners import (
|
||||
RunnerId,
|
||||
ShardAssignments,
|
||||
@@ -22,8 +22,8 @@ from exo.shared.types.worker.shards import PipelineShardMetadata
|
||||
def place_flash_instance(
|
||||
command: LaunchFLASH,
|
||||
topology: Topology,
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
) -> dict[InstanceId, Instance]:
|
||||
current_instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> dict[InstanceId, BaseInstance]:
|
||||
"""Place a FLASH simulation instance across available nodes.
|
||||
|
||||
Unlike MLX instances which use ring/JACCL topology for tensor parallelism,
|
||||
@@ -31,7 +31,7 @@ def place_flash_instance(
|
||||
node IPs so the runner can generate an MPI hostfile.
|
||||
"""
|
||||
instance_id = InstanceId()
|
||||
target_instances: dict[InstanceId, Instance] = dict(deepcopy(current_instances))
|
||||
target_instances: dict[InstanceId, BaseInstance] = dict(deepcopy(current_instances))
|
||||
|
||||
all_nodes = list(topology.list_nodes())
|
||||
|
||||
|
||||
@@ -4,14 +4,14 @@ from collections.abc import Mapping
|
||||
|
||||
from exo.plugins.implementations.flash.types import FLASHInstance
|
||||
from exo.shared.types.tasks import LoadModel, Task
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId
|
||||
from exo.shared.types.worker.instances import BaseInstance, InstanceId
|
||||
from exo.shared.types.worker.runners import RunnerId, RunnerIdle
|
||||
from exo.worker.runner.runner_supervisor import RunnerSupervisor
|
||||
|
||||
|
||||
def plan_flash(
|
||||
runners: Mapping[RunnerId, RunnerSupervisor],
|
||||
instances: Mapping[InstanceId, Instance],
|
||||
instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> Task | None:
|
||||
"""Plan tasks specifically for FLASH instances.
|
||||
|
||||
|
||||
@@ -21,7 +21,7 @@ from exo.shared.topology import Topology
|
||||
from exo.shared.types.commands import DeleteInstance
|
||||
from exo.shared.types.events import Event
|
||||
from exo.shared.types.tasks import Task
|
||||
from exo.shared.types.worker.instances import BoundInstance, Instance, InstanceId
|
||||
from exo.shared.types.worker.instances import BaseInstance, BoundInstance, InstanceId
|
||||
from exo.shared.types.worker.runners import RunnerId
|
||||
from exo.utils.channels import MpReceiver, MpSender
|
||||
from exo.worker.runner.runner_supervisor import RunnerSupervisor
|
||||
@@ -60,7 +60,7 @@ class FLASHPlugin(EXOPlugin):
|
||||
self,
|
||||
command: Any, # pyright: ignore[reportAny]
|
||||
topology: Topology,
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
current_instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> Sequence[Event]:
|
||||
from exo.master.placement import delete_instance, get_transition_events
|
||||
|
||||
@@ -81,7 +81,7 @@ class FLASHPlugin(EXOPlugin):
|
||||
def plan_task(
|
||||
self,
|
||||
runners: Mapping[RunnerId, RunnerSupervisor],
|
||||
instances: Mapping[InstanceId, Instance],
|
||||
instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> Task | None:
|
||||
return plan_flash(runners, instances)
|
||||
|
||||
|
||||
@@ -1,34 +1,21 @@
|
||||
"""FLASH plugin types - commands and instances."""
|
||||
# ruff: noqa: I001 - Import order intentional for Pydantic model_rebuild
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from exo.shared.types.common import CommandId, Host, NodeId
|
||||
from exo.shared.types.worker.runners import ShardAssignments
|
||||
from exo.utils.pydantic_ext import TaggedModel
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from exo.shared.types.worker.instances import InstanceId
|
||||
from exo.shared.types.worker.runners import RunnerId
|
||||
from exo.shared.types.worker.shards import (
|
||||
PipelineShardMetadata,
|
||||
TensorShardMetadata,
|
||||
)
|
||||
|
||||
from exo.plugins.type_registry import command_registry, instance_registry
|
||||
from exo.shared.types.commands import BaseCommand
|
||||
from exo.shared.types.common import Host, NodeId
|
||||
from exo.shared.types.worker.instances import BaseInstance, InstanceId
|
||||
from exo.shared.types.worker.runners import RunnerId
|
||||
from exo.shared.types.worker.shards import ShardMetadata
|
||||
|
||||
# ============================================================================
|
||||
# Commands
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class LaunchFLASH(TaggedModel):
|
||||
@command_registry.register
|
||||
class LaunchFLASH(BaseCommand):
|
||||
"""Command to launch a FLASH MPI simulation."""
|
||||
|
||||
command_id: CommandId = Field(default_factory=CommandId)
|
||||
simulation_name: str
|
||||
flash_executable_path: str
|
||||
parameter_file_path: str
|
||||
@@ -40,11 +27,11 @@ class LaunchFLASH(TaggedModel):
|
||||
hosts: str = ""
|
||||
|
||||
|
||||
class StopFLASH(TaggedModel):
|
||||
@command_registry.register
|
||||
class StopFLASH(BaseCommand):
|
||||
"""Command to stop a running FLASH simulation."""
|
||||
|
||||
command_id: CommandId = Field(default_factory=CommandId)
|
||||
instance_id: "InstanceId"
|
||||
instance_id: InstanceId
|
||||
|
||||
|
||||
# ============================================================================
|
||||
@@ -52,7 +39,8 @@ class StopFLASH(TaggedModel):
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class FLASHInstance(TaggedModel):
|
||||
@instance_registry.register
|
||||
class FLASHInstance(BaseInstance):
|
||||
"""Instance for FLASH MPI simulation.
|
||||
|
||||
Unlike MLX instances which do tensor parallelism, FLASH instances
|
||||
@@ -60,8 +48,6 @@ class FLASHInstance(TaggedModel):
|
||||
MPI ranks of the FLASH simulation.
|
||||
"""
|
||||
|
||||
instance_id: "InstanceId"
|
||||
shard_assignments: ShardAssignments
|
||||
hosts_by_node: dict[NodeId, list[Host]]
|
||||
flash_executable_path: str
|
||||
parameter_file_path: str
|
||||
@@ -72,20 +58,5 @@ class FLASHInstance(TaggedModel):
|
||||
coordinator_ip: str
|
||||
network_interface: str = "en0" # Network interface for MPI (e.g., en0, eth0)
|
||||
|
||||
def shard(
|
||||
self, runner_id: "RunnerId"
|
||||
) -> "PipelineShardMetadata | TensorShardMetadata | None":
|
||||
def shard(self, runner_id: RunnerId) -> ShardMetadata | None:
|
||||
return self.shard_assignments.runner_to_shard.get(runner_id, None)
|
||||
|
||||
|
||||
# Import types into module namespace for Pydantic model_rebuild() to resolve forward refs
|
||||
from exo.shared.types.worker.instances import InstanceId as InstanceId # noqa: E402, I001
|
||||
from exo.shared.types.worker.runners import RunnerId as RunnerId # noqa: E402, I001
|
||||
from exo.shared.types.worker.shards import ( # noqa: E402, I001
|
||||
PipelineShardMetadata as PipelineShardMetadata,
|
||||
TensorShardMetadata as TensorShardMetadata,
|
||||
)
|
||||
|
||||
# Rebuild models to resolve forward references
|
||||
StopFLASH.model_rebuild()
|
||||
FLASHInstance.model_rebuild()
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
"""Dynamic type registry for plugin types.
|
||||
|
||||
This module provides a registry system that allows plugins to register their
|
||||
command and instance types dynamically, eliminating the need for static union
|
||||
types and avoiding circular imports.
|
||||
"""
|
||||
|
||||
from typing import TypeVar
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from exo.utils.pydantic_ext import CamelCaseModel
|
||||
|
||||
# TypeVar for preserving exact types through the register decorator
|
||||
_TCls = TypeVar("_TCls", bound=type[CamelCaseModel])
|
||||
|
||||
|
||||
class TypeRegistry[T: CamelCaseModel]:
|
||||
"""Registry for dynamically registered Pydantic types.
|
||||
|
||||
Enables plugins to register their types at import time. Deserialization
|
||||
uses the class name from the tagged JSON format to look up the correct type.
|
||||
"""
|
||||
|
||||
def __init__(self, name: str) -> None:
|
||||
self._name = name
|
||||
self._types: dict[str, type[T]] = {}
|
||||
|
||||
def register(self, cls: _TCls) -> _TCls:
|
||||
"""Decorator to register a type with this registry.
|
||||
|
||||
Preserves the exact type through the decorator for proper type checking.
|
||||
"""
|
||||
self._types[cls.__name__] = cls # type: ignore[assignment]
|
||||
logger.debug(f"{self._name}: registered {cls.__name__}")
|
||||
return cls
|
||||
|
||||
def get(self, name: str) -> type[T] | None:
|
||||
"""Look up a type by class name."""
|
||||
return self._types.get(name)
|
||||
|
||||
def all_types(self) -> dict[str, type[T]]:
|
||||
"""Return all registered types."""
|
||||
return dict(self._types)
|
||||
|
||||
def deserialize(self, data: dict[str, dict[str, object]] | CamelCaseModel) -> T:
|
||||
"""Deserialize dict to the appropriate registered type.
|
||||
|
||||
Supports two formats:
|
||||
1. Tagged format: {"ClassName": {...fields...}} - used for network serialization
|
||||
2. Flat format: {...fields...} - used for API requests, tries each type
|
||||
"""
|
||||
# If already deserialized (e.g., from Pydantic), return as-is
|
||||
if isinstance(data, CamelCaseModel):
|
||||
return data # type: ignore[return-value]
|
||||
|
||||
# Check for tagged format: single key that matches a registered type
|
||||
if len(data) == 1:
|
||||
class_name: str = next(iter(data.keys()))
|
||||
cls = self._types.get(class_name)
|
||||
if cls is not None:
|
||||
return cls.model_validate(data[class_name], strict=False)
|
||||
|
||||
# Flat format: try each registered type, use first that validates
|
||||
errors: list[str] = []
|
||||
for type_name, cls in self._types.items():
|
||||
try:
|
||||
return cls.model_validate(data, strict=False)
|
||||
except Exception as e: # noqa: BLE001
|
||||
errors.append(f"{type_name}: {e}")
|
||||
|
||||
# None matched - provide helpful error
|
||||
available = ", ".join(self._types.keys())
|
||||
raise ValueError(
|
||||
f"{self._name}: could not deserialize data. "
|
||||
f"Available types: {available}. Errors: {'; '.join(errors[:3])}"
|
||||
)
|
||||
|
||||
|
||||
# Global registries for commands, instances, events, and tasks
|
||||
command_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("CommandRegistry")
|
||||
instance_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("InstanceRegistry")
|
||||
event_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("EventRegistry")
|
||||
task_registry: TypeRegistry[CamelCaseModel] = TypeRegistry("TaskRegistry")
|
||||
@@ -30,7 +30,7 @@ class TypedTopic[T: CamelCaseModel]:
|
||||
|
||||
@staticmethod
|
||||
def serialize(t: T) -> bytes:
|
||||
return t.model_dump_json().encode("utf-8")
|
||||
return t.model_dump_json(by_alias=True, serialize_as_any=True).encode("utf-8")
|
||||
|
||||
def deserialize(self, b: bytes) -> T:
|
||||
return self.model_type.model_validate_json(b.decode("utf-8"))
|
||||
|
||||
@@ -34,7 +34,7 @@ from exo.shared.types.state import State
|
||||
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
|
||||
from exo.shared.types.worker.instances import BaseInstance, InstanceId
|
||||
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
|
||||
from exo.utils.info_gatherer.info_gatherer import (
|
||||
MacmonMetrics,
|
||||
@@ -81,6 +81,10 @@ def event_apply(event: Event, state: State) -> State:
|
||||
return apply_topology_edge_created(event, state)
|
||||
case TopologyEdgeDeleted():
|
||||
return apply_topology_edge_deleted(event, state)
|
||||
case _:
|
||||
# Unknown event types from plugins are ignored
|
||||
logger.debug(f"Ignoring unknown event type: {type(event).__name__}")
|
||||
return state
|
||||
|
||||
|
||||
def apply(state: State, event: IndexedEvent) -> State:
|
||||
@@ -163,7 +167,7 @@ def apply_task_failed(event: TaskFailed, state: State) -> State:
|
||||
|
||||
def apply_instance_created(event: InstanceCreated, state: State) -> State:
|
||||
instance = event.instance
|
||||
new_instances: Mapping[InstanceId, Instance] = {
|
||||
new_instances: Mapping[InstanceId, BaseInstance] = {
|
||||
**state.instances,
|
||||
instance.instance_id: instance,
|
||||
}
|
||||
@@ -171,7 +175,7 @@ def apply_instance_created(event: InstanceCreated, state: State) -> State:
|
||||
|
||||
|
||||
def apply_instance_deleted(event: InstanceDeleted, state: State) -> State:
|
||||
new_instances: Mapping[InstanceId, Instance] = {
|
||||
new_instances: Mapping[InstanceId, BaseInstance] = {
|
||||
iid: inst for iid, inst in state.instances.items() if iid != event.instance_id
|
||||
}
|
||||
return state.model_copy(update={"instances": new_instances})
|
||||
|
||||
@@ -1,13 +1,19 @@
|
||||
import time
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from pydantic_core import PydanticUseDefault
|
||||
|
||||
from exo.plugins.type_registry import instance_registry
|
||||
from exo.shared.types.common import CommandId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.models import ModelId, ModelMetadata
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
|
||||
from exo.shared.types.worker.instances import (
|
||||
BaseInstance,
|
||||
Instance,
|
||||
InstanceId,
|
||||
InstanceMeta,
|
||||
)
|
||||
from exo.shared.types.worker.shards import Sharding
|
||||
|
||||
FinishReason = Literal[
|
||||
@@ -184,6 +190,12 @@ class PlaceInstanceParams(BaseModel):
|
||||
class CreateInstanceParams(BaseModel):
|
||||
instance: Instance
|
||||
|
||||
@field_validator("instance", mode="before")
|
||||
@classmethod
|
||||
def validate_instance(cls, v: Any) -> BaseInstance: # noqa: ANN401 # pyright: ignore[reportAny]
|
||||
"""Validate instance using registry to handle both tagged and flat formats."""
|
||||
return cast(BaseInstance, instance_registry.deserialize(v)) # pyright: ignore[reportAny]
|
||||
|
||||
|
||||
class PlacementPreview(BaseModel):
|
||||
model_id: ModelId
|
||||
|
||||
@@ -1,6 +1,14 @@
|
||||
# ruff: noqa: I001 - Import order intentional to avoid circular imports
|
||||
from pydantic import Field
|
||||
"""Command types for exo.
|
||||
|
||||
Commands are registered dynamically via the command_registry, allowing plugins
|
||||
to add their own command types without modifying this file.
|
||||
"""
|
||||
|
||||
from typing import Any, cast
|
||||
|
||||
from pydantic import Field, field_validator
|
||||
|
||||
from exo.plugins.type_registry import command_registry
|
||||
from exo.shared.types.api import ChatCompletionTaskParams
|
||||
from exo.shared.types.common import CommandId, NodeId
|
||||
from exo.shared.types.models import ModelMetadata
|
||||
@@ -8,22 +16,24 @@ from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
|
||||
from exo.shared.types.worker.shards import Sharding
|
||||
from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
|
||||
|
||||
# Import FLASH commands from plugin (for serialization compatibility)
|
||||
from exo.plugins.implementations.flash.types import LaunchFLASH, StopFLASH # noqa: E402, I001
|
||||
|
||||
|
||||
class BaseCommand(TaggedModel):
|
||||
"""Base class for all commands."""
|
||||
|
||||
command_id: CommandId = Field(default_factory=CommandId)
|
||||
|
||||
|
||||
@command_registry.register
|
||||
class TestCommand(BaseCommand):
|
||||
__test__ = False
|
||||
|
||||
|
||||
@command_registry.register
|
||||
class ChatCompletion(BaseCommand):
|
||||
request_params: ChatCompletionTaskParams
|
||||
|
||||
|
||||
@command_registry.register
|
||||
class PlaceInstance(BaseCommand):
|
||||
model_meta: ModelMetadata
|
||||
sharding: Sharding
|
||||
@@ -31,22 +41,27 @@ class PlaceInstance(BaseCommand):
|
||||
min_nodes: int
|
||||
|
||||
|
||||
@command_registry.register
|
||||
class CreateInstance(BaseCommand):
|
||||
instance: Instance
|
||||
|
||||
|
||||
@command_registry.register
|
||||
class DeleteInstance(BaseCommand):
|
||||
instance_id: InstanceId
|
||||
|
||||
|
||||
@command_registry.register
|
||||
class TaskFinished(BaseCommand):
|
||||
finished_command_id: CommandId
|
||||
|
||||
|
||||
@command_registry.register
|
||||
class RequestEventLog(BaseCommand):
|
||||
since_idx: int
|
||||
|
||||
|
||||
# Union type for core commands - used by ForwarderCommand for network deserialization
|
||||
Command = (
|
||||
TestCommand
|
||||
| RequestEventLog
|
||||
@@ -54,12 +69,19 @@ Command = (
|
||||
| PlaceInstance
|
||||
| CreateInstance
|
||||
| DeleteInstance
|
||||
| LaunchFLASH
|
||||
| StopFLASH
|
||||
| TaskFinished
|
||||
)
|
||||
|
||||
|
||||
class ForwarderCommand(CamelCaseModel):
|
||||
"""Wrapper for commands that includes origin node."""
|
||||
|
||||
origin: NodeId
|
||||
command: Command
|
||||
command: BaseCommand
|
||||
|
||||
@field_validator("command", mode="before")
|
||||
@classmethod
|
||||
def validate_command(cls, v: Any) -> BaseCommand: # noqa: ANN401 # pyright: ignore[reportAny]
|
||||
"""Validate command, using registry for plugin commands not in Command union."""
|
||||
# First try the registry (handles both core and plugin commands)
|
||||
return cast(BaseCommand, command_registry.deserialize(v)) # pyright: ignore[reportAny]
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
from datetime import datetime
|
||||
from typing import Any, cast
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic import Field, field_validator
|
||||
|
||||
from exo.plugins.type_registry import event_registry, instance_registry, task_registry
|
||||
from exo.shared.topology import Connection
|
||||
from exo.shared.types.chunks import GenerationChunk
|
||||
from exo.shared.types.common import CommandId, Id, NodeId, SessionId
|
||||
from exo.shared.types.tasks import Task, TaskId, TaskStatus
|
||||
from exo.shared.types.tasks import BaseTask, TaskId, TaskStatus
|
||||
from exo.shared.types.worker.downloads import DownloadProgress
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId
|
||||
from exo.shared.types.worker.instances import BaseInstance, InstanceId
|
||||
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
|
||||
from exo.utils.info_gatherer.info_gatherer import GatheredInfo
|
||||
from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
|
||||
@@ -25,36 +27,53 @@ class BaseEvent(TaggedModel):
|
||||
_master_time_stamp: None | datetime = None
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TestEvent(BaseEvent):
|
||||
__test__ = False
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TaskCreated(BaseEvent):
|
||||
task_id: TaskId
|
||||
task: Task
|
||||
task: BaseTask
|
||||
|
||||
@field_validator("task", mode="before")
|
||||
@classmethod
|
||||
def validate_task(cls, v: Any) -> BaseTask: # noqa: ANN401 # pyright: ignore[reportAny]
|
||||
return cast(BaseTask, task_registry.deserialize(v)) # pyright: ignore[reportAny]
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TaskAcknowledged(BaseEvent):
|
||||
task_id: TaskId
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TaskDeleted(BaseEvent):
|
||||
task_id: TaskId
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TaskStatusUpdated(BaseEvent):
|
||||
task_id: TaskId
|
||||
task_status: TaskStatus
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TaskFailed(BaseEvent):
|
||||
task_id: TaskId
|
||||
error_type: str
|
||||
error_message: str
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class InstanceCreated(BaseEvent):
|
||||
instance: Instance
|
||||
instance: BaseInstance
|
||||
|
||||
@field_validator("instance", mode="before")
|
||||
@classmethod
|
||||
def validate_instance(cls, v: Any) -> BaseInstance: # noqa: ANN401 # pyright: ignore[reportAny]
|
||||
return cast(BaseInstance, instance_registry.deserialize(v)) # pyright: ignore[reportAny]
|
||||
|
||||
def __eq__(self, other: object) -> bool:
|
||||
if isinstance(other, InstanceCreated):
|
||||
@@ -63,72 +82,71 @@ class InstanceCreated(BaseEvent):
|
||||
return False
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class InstanceDeleted(BaseEvent):
|
||||
instance_id: InstanceId
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class RunnerStatusUpdated(BaseEvent):
|
||||
runner_id: RunnerId
|
||||
runner_status: RunnerStatus
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class RunnerDeleted(BaseEvent):
|
||||
runner_id: RunnerId
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class NodeTimedOut(BaseEvent):
|
||||
node_id: NodeId
|
||||
|
||||
|
||||
# TODO: bikeshed this name
|
||||
@event_registry.register
|
||||
class NodeGatheredInfo(BaseEvent):
|
||||
node_id: NodeId
|
||||
when: str # this is a manually cast datetime overrode by the master when the event is indexed, rather than the local time on the device
|
||||
info: GatheredInfo
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class NodeDownloadProgress(BaseEvent):
|
||||
download_progress: DownloadProgress
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class ChunkGenerated(BaseEvent):
|
||||
command_id: CommandId
|
||||
chunk: GenerationChunk
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TopologyEdgeCreated(BaseEvent):
|
||||
conn: Connection
|
||||
|
||||
|
||||
@event_registry.register
|
||||
class TopologyEdgeDeleted(BaseEvent):
|
||||
conn: Connection
|
||||
|
||||
|
||||
Event = (
|
||||
TestEvent
|
||||
| TaskCreated
|
||||
| TaskStatusUpdated
|
||||
| TaskFailed
|
||||
| TaskDeleted
|
||||
| TaskAcknowledged
|
||||
| InstanceCreated
|
||||
| InstanceDeleted
|
||||
| RunnerStatusUpdated
|
||||
| RunnerDeleted
|
||||
| NodeTimedOut
|
||||
| NodeGatheredInfo
|
||||
| NodeDownloadProgress
|
||||
| ChunkGenerated
|
||||
| TopologyEdgeCreated
|
||||
| TopologyEdgeDeleted
|
||||
)
|
||||
# Type alias for backward compatibility - use BaseEvent for type hints
|
||||
# Actual deserialization uses event_registry
|
||||
Event = BaseEvent
|
||||
|
||||
|
||||
class IndexedEvent(CamelCaseModel):
|
||||
"""An event indexed by the master, with a globally unique index"""
|
||||
|
||||
idx: int = Field(ge=0)
|
||||
event: Event
|
||||
event: BaseEvent
|
||||
|
||||
@field_validator("event", mode="before")
|
||||
@classmethod
|
||||
def validate_event(cls, v: Any) -> BaseEvent: # noqa: ANN401 # pyright: ignore[reportAny]
|
||||
return cast(BaseEvent, event_registry.deserialize(v)) # pyright: ignore[reportAny]
|
||||
|
||||
|
||||
class ForwarderEvent(CamelCaseModel):
|
||||
@@ -137,4 +155,9 @@ class ForwarderEvent(CamelCaseModel):
|
||||
origin_idx: int = Field(ge=0)
|
||||
origin: NodeId
|
||||
session: SessionId
|
||||
event: Event
|
||||
event: BaseEvent
|
||||
|
||||
@field_validator("event", mode="before")
|
||||
@classmethod
|
||||
def validate_event(cls, v: Any) -> BaseEvent: # noqa: ANN401 # pyright: ignore[reportAny]
|
||||
return cast(BaseEvent, event_registry.deserialize(v)) # pyright: ignore[reportAny]
|
||||
|
||||
@@ -16,7 +16,7 @@ from exo.shared.types.profiling import (
|
||||
)
|
||||
from exo.shared.types.tasks import Task, TaskId
|
||||
from exo.shared.types.worker.downloads import DownloadProgress
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId
|
||||
from exo.shared.types.worker.instances import BaseInstance, InstanceId
|
||||
from exo.shared.types.worker.runners import RunnerId, RunnerStatus
|
||||
from exo.utils.pydantic_ext import CamelCaseModel
|
||||
|
||||
@@ -37,7 +37,7 @@ class State(CamelCaseModel):
|
||||
strict=True,
|
||||
arbitrary_types_allowed=True,
|
||||
)
|
||||
instances: Mapping[InstanceId, Instance] = {}
|
||||
instances: Mapping[InstanceId, BaseInstance] = {}
|
||||
runners: Mapping[RunnerId, RunnerStatus] = {}
|
||||
downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
|
||||
tasks: Mapping[TaskId, Task] = {}
|
||||
@@ -52,6 +52,16 @@ class State(CamelCaseModel):
|
||||
node_network: Mapping[NodeId, NodeNetworkInfo] = {}
|
||||
node_thunderbolt: Mapping[NodeId, NodeThunderboltInfo] = {}
|
||||
|
||||
@field_serializer("instances", mode="plain")
|
||||
def _encode_instances(
|
||||
self, value: Mapping[InstanceId, BaseInstance]
|
||||
) -> dict[str, Any]:
|
||||
"""Serialize instances with full subclass fields."""
|
||||
return {
|
||||
str(k): v.model_dump(by_alias=True, serialize_as_any=True)
|
||||
for k, v in value.items()
|
||||
}
|
||||
|
||||
@field_serializer("topology", mode="plain")
|
||||
def _encode_topology(self, value: Topology) -> TopologySnapshot:
|
||||
return value.to_snapshot()
|
||||
|
||||
@@ -2,6 +2,7 @@ from enum import Enum
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from exo.plugins.type_registry import task_registry
|
||||
from exo.shared.types.api import ChatCompletionTaskParams
|
||||
from exo.shared.types.common import CommandId, Id
|
||||
from exo.shared.types.worker.instances import BoundInstance, InstanceId
|
||||
@@ -28,26 +29,32 @@ class BaseTask(TaggedModel):
|
||||
instance_id: InstanceId
|
||||
|
||||
|
||||
@task_registry.register
|
||||
class CreateRunner(BaseTask): # emitted by Worker
|
||||
bound_instance: BoundInstance
|
||||
|
||||
|
||||
@task_registry.register
|
||||
class DownloadModel(BaseTask): # emitted by Worker
|
||||
shard_metadata: ShardMetadata
|
||||
|
||||
|
||||
@task_registry.register
|
||||
class LoadModel(BaseTask): # emitted by Worker
|
||||
pass
|
||||
|
||||
|
||||
@task_registry.register
|
||||
class ConnectToGroup(BaseTask): # emitted by Worker
|
||||
pass
|
||||
|
||||
|
||||
@task_registry.register
|
||||
class StartWarmup(BaseTask): # emitted by Worker
|
||||
pass
|
||||
|
||||
|
||||
@task_registry.register
|
||||
class ChatCompletion(BaseTask): # emitted by Master
|
||||
command_id: CommandId
|
||||
task_params: ChatCompletionTaskParams
|
||||
@@ -56,16 +63,11 @@ class ChatCompletion(BaseTask): # emitted by Master
|
||||
error_message: str | None = Field(default=None)
|
||||
|
||||
|
||||
@task_registry.register
|
||||
class Shutdown(BaseTask): # emitted by Worker
|
||||
runner_id: RunnerId
|
||||
|
||||
|
||||
Task = (
|
||||
CreateRunner
|
||||
| DownloadModel
|
||||
| ConnectToGroup
|
||||
| LoadModel
|
||||
| StartWarmup
|
||||
| ChatCompletion
|
||||
| Shutdown
|
||||
)
|
||||
# Type alias for backward compatibility - use BaseTask for type hints
|
||||
# Actual deserialization uses task_registry
|
||||
Task = BaseTask
|
||||
|
||||
@@ -1,8 +1,15 @@
|
||||
# ruff: noqa: I001 - Import order intentional to avoid circular imports
|
||||
"""Instance types for exo.
|
||||
|
||||
Instances are registered dynamically via the instance_registry, allowing plugins
|
||||
to add their own instance types without modifying this file.
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
from typing import Any, cast
|
||||
|
||||
from pydantic import model_validator
|
||||
from pydantic import field_validator, model_validator
|
||||
|
||||
from exo.plugins.type_registry import instance_registry
|
||||
from exo.shared.types.common import Host, Id, NodeId
|
||||
from exo.shared.types.worker.runners import RunnerId, ShardAssignments, ShardMetadata
|
||||
from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel
|
||||
@@ -15,10 +22,11 @@ class InstanceId(Id):
|
||||
class InstanceMeta(str, Enum):
|
||||
MlxRing = "MlxRing"
|
||||
MlxJaccl = "MlxJaccl"
|
||||
FLASH = "FLASH"
|
||||
|
||||
|
||||
class BaseInstance(TaggedModel):
|
||||
"""Base class for all instance types."""
|
||||
|
||||
instance_id: InstanceId
|
||||
shard_assignments: ShardAssignments
|
||||
|
||||
@@ -26,29 +34,36 @@ class BaseInstance(TaggedModel):
|
||||
return self.shard_assignments.runner_to_shard.get(runner_id, None)
|
||||
|
||||
|
||||
@instance_registry.register
|
||||
class MlxRingInstance(BaseInstance):
|
||||
hosts_by_node: dict[NodeId, list[Host]]
|
||||
ephemeral_port: int
|
||||
|
||||
|
||||
@instance_registry.register
|
||||
class MlxJacclInstance(BaseInstance):
|
||||
jaccl_devices: list[list[str | None]]
|
||||
jaccl_coordinators: dict[NodeId, str]
|
||||
|
||||
|
||||
# Import FLASHInstance from plugin (for serialization compatibility)
|
||||
from exo.plugins.implementations.flash.types import FLASHInstance # noqa: E402, I001
|
||||
|
||||
|
||||
# TODO: Single node instance
|
||||
Instance = MlxRingInstance | MlxJacclInstance | FLASHInstance
|
||||
# Union type for Pydantic validation - tries each type in order
|
||||
# This is used by API endpoints (dashboard) which send flat format
|
||||
Instance = MlxRingInstance | MlxJacclInstance
|
||||
|
||||
|
||||
class BoundInstance(CamelCaseModel):
|
||||
instance: Instance
|
||||
"""An instance bound to a specific runner on a specific node."""
|
||||
|
||||
instance: BaseInstance
|
||||
bound_runner_id: RunnerId
|
||||
bound_node_id: NodeId
|
||||
|
||||
@field_validator("instance", mode="before")
|
||||
@classmethod
|
||||
def validate_instance(cls, v: Any) -> BaseInstance: # noqa: ANN401 # pyright: ignore[reportAny]
|
||||
"""Validate instance using registry to handle both tagged and flat formats."""
|
||||
return cast(BaseInstance, instance_registry.deserialize(v)) # pyright: ignore[reportAny]
|
||||
|
||||
@property
|
||||
def bound_shard(self) -> ShardMetadata:
|
||||
shard = self.instance.shard(self.bound_runner_id)
|
||||
|
||||
@@ -22,8 +22,8 @@ from exo.shared.types.worker.downloads import (
|
||||
DownloadProgress,
|
||||
)
|
||||
from exo.shared.types.worker.instances import (
|
||||
BaseInstance,
|
||||
BoundInstance,
|
||||
Instance,
|
||||
InstanceId,
|
||||
)
|
||||
from exo.shared.types.worker.runners import (
|
||||
@@ -50,7 +50,7 @@ def plan(
|
||||
download_status: Mapping[ModelId, DownloadProgress],
|
||||
# gdls is not expected to be fresh
|
||||
global_download_status: Mapping[NodeId, Sequence[DownloadProgress]],
|
||||
instances: Mapping[InstanceId, Instance],
|
||||
instances: Mapping[InstanceId, BaseInstance],
|
||||
all_runners: Mapping[RunnerId, RunnerStatus], # all global
|
||||
tasks: Mapping[TaskId, Task],
|
||||
) -> Task | None:
|
||||
@@ -79,7 +79,7 @@ def plan(
|
||||
def _kill_runner(
|
||||
runners: Mapping[RunnerId, RunnerSupervisor],
|
||||
all_runners: Mapping[RunnerId, RunnerStatus],
|
||||
instances: Mapping[InstanceId, Instance],
|
||||
instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> Shutdown | None:
|
||||
for runner in runners.values():
|
||||
runner_id = runner.bound_instance.bound_runner_id
|
||||
@@ -102,7 +102,7 @@ def _kill_runner(
|
||||
def _create_runner(
|
||||
node_id: NodeId,
|
||||
runners: Mapping[RunnerId, RunnerSupervisor],
|
||||
instances: Mapping[InstanceId, Instance],
|
||||
instances: Mapping[InstanceId, BaseInstance],
|
||||
) -> CreateRunner | None:
|
||||
for instance in instances.values():
|
||||
runner_id = instance.shard_assignments.node_to_runner.get(node_id, None)
|
||||
|
||||
Reference in New Issue
Block a user