Add explicit MetaInstance binding, slim MetaInstance to use ModelId
- Add MetaInstanceBound event and meta_instance_backing State field for explicit MetaInstance → Instance binding (prevents ambiguous linking when two MetaInstances have identical constraints) - Replace model_card: ModelCard with model_id: ModelId on MetaInstance (load ModelCard on-demand at placement time) - Add MetaInstance API endpoints (POST /meta_instance, DELETE) - Update dashboard to use MetaInstances as primary primitive with unified display items merging MetaInstances and orphan instances - Dashboard launches via MetaInstance instead of direct Instance creation Co-Authored-By: Claude Opus 4.6 <[email protected]>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
ceb76b8f6c
commit
f4329c72c2
@@ -232,6 +232,19 @@ interface RawStateResponse {
|
||||
>;
|
||||
// Thunderbolt bridge cycles (nodes with bridge enabled forming loops)
|
||||
thunderboltBridgeCycles?: string[][];
|
||||
// MetaInstances (declarative instance constraints)
|
||||
metaInstances?: Record<string, MetaInstanceData>;
|
||||
// Explicit MetaInstance → Instance binding
|
||||
metaInstanceBacking?: Record<string, string>;
|
||||
}
|
||||
|
||||
export interface MetaInstanceData {
|
||||
metaInstanceId: string;
|
||||
modelId: string;
|
||||
sharding: string;
|
||||
instanceMeta: string;
|
||||
minNodes: number;
|
||||
nodeIds: string[] | null;
|
||||
}
|
||||
|
||||
export interface MessageAttachment {
|
||||
@@ -495,6 +508,8 @@ class AppStore {
|
||||
isLoadingPreviews = $state(false);
|
||||
previewNodeFilter = $state<Set<string>>(new Set());
|
||||
lastUpdate = $state<number | null>(null);
|
||||
metaInstances = $state<Record<string, MetaInstanceData>>({});
|
||||
metaInstanceBacking = $state<Record<string, string>>({});
|
||||
thunderboltBridgeCycles = $state<string[][]>([]);
|
||||
nodeThunderboltBridge = $state<
|
||||
Record<
|
||||
@@ -1206,6 +1221,9 @@ class AppStore {
|
||||
if (data.downloads) {
|
||||
this.downloads = data.downloads;
|
||||
}
|
||||
// MetaInstances
|
||||
this.metaInstances = data.metaInstances ?? {};
|
||||
this.metaInstanceBacking = data.metaInstanceBacking ?? {};
|
||||
// Thunderbolt bridge cycles
|
||||
this.thunderboltBridgeCycles = data.thunderboltBridgeCycles ?? [];
|
||||
// Thunderbolt bridge status per node
|
||||
@@ -2956,6 +2974,8 @@ export const tps = () => appStore.tps;
|
||||
export const totalTokens = () => appStore.totalTokens;
|
||||
export const topologyData = () => appStore.topologyData;
|
||||
export const instances = () => appStore.instances;
|
||||
export const metaInstances = () => appStore.metaInstances;
|
||||
export const metaInstanceBacking = () => appStore.metaInstanceBacking;
|
||||
export const runners = () => appStore.runners;
|
||||
export const downloads = () => appStore.downloads;
|
||||
export const placementPreviews = () => appStore.placementPreviews;
|
||||
|
||||
@@ -37,10 +37,13 @@
|
||||
toggleTopologyOnlyMode,
|
||||
chatSidebarVisible,
|
||||
toggleChatSidebarVisible,
|
||||
metaInstances,
|
||||
metaInstanceBacking,
|
||||
thunderboltBridgeCycles,
|
||||
nodeThunderboltBridge,
|
||||
type DownloadProgress,
|
||||
type PlacementPreview,
|
||||
type MetaInstanceData,
|
||||
} from "$lib/stores/app.svelte";
|
||||
import HeaderNav from "$lib/components/HeaderNav.svelte";
|
||||
import { fade, fly } from "svelte/transition";
|
||||
@@ -60,6 +63,8 @@
|
||||
const debugEnabled = $derived(debugMode());
|
||||
const topologyOnlyEnabled = $derived(topologyOnlyMode());
|
||||
const sidebarVisible = $derived(chatSidebarVisible());
|
||||
const metaInstancesData = $derived(metaInstances());
|
||||
const metaInstanceBackingData = $derived(metaInstanceBacking());
|
||||
const tbBridgeCycles = $derived(thunderboltBridgeCycles());
|
||||
const tbBridgeData = $derived(nodeThunderboltBridge());
|
||||
const nodeFilter = $derived(previewNodeFilter());
|
||||
@@ -593,39 +598,22 @@
|
||||
launchingModelId = modelId;
|
||||
|
||||
try {
|
||||
// Use the specific preview if provided, otherwise fall back to filtered preview
|
||||
const preview = specificPreview ?? filteredPreview();
|
||||
|
||||
let instanceData: unknown;
|
||||
|
||||
if (preview?.instance) {
|
||||
// Use the instance from the preview
|
||||
instanceData = preview.instance;
|
||||
} else {
|
||||
// Fallback: GET placement from API
|
||||
const placementResponse = await fetch(
|
||||
`/instance/placement?model_id=${encodeURIComponent(modelId)}&sharding=${selectedSharding}&instance_meta=${selectedInstanceType}&min_nodes=${selectedMinNodes}`,
|
||||
);
|
||||
|
||||
if (!placementResponse.ok) {
|
||||
const errorText = await placementResponse.text();
|
||||
console.error("Failed to get placement:", errorText);
|
||||
return;
|
||||
}
|
||||
|
||||
instanceData = await placementResponse.json();
|
||||
}
|
||||
|
||||
// POST the instance to create it
|
||||
const response = await fetch("/instance", {
|
||||
const response = await fetch("/meta_instance", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({ instance: instanceData }),
|
||||
body: JSON.stringify({
|
||||
model_id: modelId,
|
||||
sharding: preview?.sharding ?? selectedSharding,
|
||||
instance_meta: preview?.instance_meta ?? selectedInstanceType,
|
||||
min_nodes: selectedMinNodes,
|
||||
}),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorText = await response.text();
|
||||
console.error("Failed to launch instance:", errorText);
|
||||
console.error("Failed to create meta instance:", errorText);
|
||||
} else {
|
||||
// Always auto-select the newly launched model so the user chats to what they just launched
|
||||
setSelectedChatModel(modelId);
|
||||
@@ -645,7 +633,7 @@
|
||||
setTimeout(scrollToBottom, 1000);
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Error launching instance:", error);
|
||||
console.error("Error creating meta instance:", error);
|
||||
} finally {
|
||||
launchingModelId = null;
|
||||
}
|
||||
@@ -1148,6 +1136,64 @@
|
||||
}
|
||||
}
|
||||
|
||||
async function deleteMetaInstance(metaInstanceId: string) {
|
||||
const meta = metaInstancesData[metaInstanceId];
|
||||
const modelId = meta?.modelId ?? "unknown";
|
||||
if (!confirm(`Delete model ${modelId}?`)) return;
|
||||
|
||||
const wasSelected = selectedChatModel() === modelId;
|
||||
|
||||
try {
|
||||
const response = await fetch(`/meta_instance/${metaInstanceId}`, {
|
||||
method: "DELETE",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
console.error("Failed to delete meta instance:", response.status);
|
||||
} else if (wasSelected) {
|
||||
// Switch to another available model or clear selection
|
||||
const remainingInstances = Object.entries(instanceData).filter(
|
||||
([id]) => id !== getBackingInstanceId(metaInstanceId),
|
||||
);
|
||||
if (remainingInstances.length > 0) {
|
||||
const [, lastInstance] =
|
||||
remainingInstances[remainingInstances.length - 1];
|
||||
const newModelId = getInstanceModelId(lastInstance);
|
||||
if (
|
||||
newModelId &&
|
||||
newModelId !== "Unknown" &&
|
||||
newModelId !== "Unknown Model"
|
||||
) {
|
||||
setSelectedChatModel(newModelId);
|
||||
} else {
|
||||
setSelectedChatModel("");
|
||||
}
|
||||
} else {
|
||||
setSelectedChatModel("");
|
||||
}
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Error deleting meta instance:", error);
|
||||
}
|
||||
}
|
||||
|
||||
// Find the backing Instance ID for a MetaInstance via the explicit binding map
|
||||
function getBackingInstanceId(metaInstanceId: string): string | null {
|
||||
return metaInstanceBackingData[metaInstanceId] ?? null;
|
||||
}
|
||||
|
||||
// Get the set of Instance IDs that are backing MetaInstances
|
||||
function getBackedInstanceIds(): Set<string> {
|
||||
return new Set(Object.values(metaInstanceBackingData));
|
||||
}
|
||||
|
||||
// Get orphan Instance IDs (not backing any MetaInstance)
|
||||
function getOrphanInstanceIds(): string[] {
|
||||
const backedIds = getBackedInstanceIds();
|
||||
return Object.keys(instanceData).filter((id) => !backedIds.has(id));
|
||||
}
|
||||
|
||||
// Helper to unwrap tagged unions like { MlxRingInstance: {...} }
|
||||
function getTagged(obj: unknown): [string | null, unknown] {
|
||||
if (!obj || typeof obj !== "object") return [null, null];
|
||||
@@ -1540,7 +1586,45 @@
|
||||
}
|
||||
|
||||
const nodeCount = $derived(data ? Object.keys(data.nodes).length : 0);
|
||||
const instanceCount = $derived(Object.keys(instanceData).length);
|
||||
const metaInstanceCount = $derived(Object.keys(metaInstancesData).length);
|
||||
const orphanInstanceIds = $derived(getOrphanInstanceIds());
|
||||
const instanceCount = $derived(metaInstanceCount + orphanInstanceIds.length);
|
||||
|
||||
// Unified display items: MetaInstances first, then orphan Instances
|
||||
interface DisplayItem {
|
||||
id: string; // MetaInstance ID or Instance ID (used as key and displayed)
|
||||
modelId: string;
|
||||
instance: unknown | null; // The backing/orphan instance (tagged union) or null if placing
|
||||
instanceId: string | null; // The actual Instance ID (for topology hover)
|
||||
isMetaInstance: boolean;
|
||||
}
|
||||
|
||||
const unifiedDisplayItems = $derived((): DisplayItem[] => {
|
||||
const items: DisplayItem[] = [];
|
||||
// MetaInstances
|
||||
for (const [metaId, meta] of Object.entries(metaInstancesData)) {
|
||||
const backingId = getBackingInstanceId(metaId);
|
||||
items.push({
|
||||
id: metaId,
|
||||
modelId: meta.modelId,
|
||||
instance: backingId ? instanceData[backingId] : null,
|
||||
instanceId: backingId,
|
||||
isMetaInstance: true,
|
||||
});
|
||||
}
|
||||
// Orphan Instances
|
||||
for (const orphanId of getOrphanInstanceIds()) {
|
||||
const inst = instanceData[orphanId];
|
||||
items.push({
|
||||
id: orphanId,
|
||||
modelId: getInstanceModelId(inst),
|
||||
instance: inst,
|
||||
instanceId: orphanId,
|
||||
isMetaInstance: false,
|
||||
});
|
||||
}
|
||||
return items;
|
||||
});
|
||||
|
||||
// Helper to get the number of nodes in a placement preview
|
||||
function getPreviewNodeCount(preview: PlacementPreview): number {
|
||||
@@ -1985,31 +2069,51 @@
|
||||
bind:this={instancesContainerRef}
|
||||
class="max-h-72 xl:max-h-96 space-y-3 overflow-y-auto overflow-x-hidden py-px"
|
||||
>
|
||||
{#each Object.entries(instanceData) as [id, instance]}
|
||||
{@const downloadInfo = getInstanceDownloadStatus(
|
||||
id,
|
||||
instance,
|
||||
)}
|
||||
{#each unifiedDisplayItems() as item (item.id)}
|
||||
{@const id = item.id}
|
||||
{@const instance = item.instance}
|
||||
{@const downloadInfo = instance
|
||||
? getInstanceDownloadStatus(item.instanceId ?? id, instance)
|
||||
: {
|
||||
statusText: "PLACING",
|
||||
statusClass: "starting",
|
||||
isDownloading: false,
|
||||
isFailed: false,
|
||||
progress: null,
|
||||
perNode: [],
|
||||
errorMessage: null,
|
||||
}}
|
||||
{@const statusText = downloadInfo.statusText}
|
||||
{@const isDownloading = downloadInfo.isDownloading}
|
||||
{@const isFailed = statusText === "FAILED"}
|
||||
{@const isLoading =
|
||||
statusText === "LOADING" ||
|
||||
statusText === "WARMING UP" ||
|
||||
statusText === "WAITING"}
|
||||
statusText === "WAITING" ||
|
||||
statusText === "PLACING"}
|
||||
{@const isReady =
|
||||
statusText === "READY" || statusText === "LOADED"}
|
||||
{@const isRunning = statusText === "RUNNING"}
|
||||
<!-- Instance Card -->
|
||||
{@const instanceModelId = getInstanceModelId(instance)}
|
||||
{@const instanceInfo = getInstanceInfo(instance)}
|
||||
{@const instanceConnections =
|
||||
getInstanceConnections(instance)}
|
||||
{@const instanceModelId = item.modelId}
|
||||
{@const instanceInfo = instance
|
||||
? getInstanceInfo(instance)
|
||||
: {
|
||||
instanceType: "Unknown",
|
||||
sharding: "Unknown",
|
||||
nodeNames: [],
|
||||
nodeIds: [],
|
||||
nodeCount: 0,
|
||||
}}
|
||||
{@const instanceConnections = instance
|
||||
? getInstanceConnections(instance)
|
||||
: []}
|
||||
<div
|
||||
class="relative group cursor-pointer"
|
||||
role="button"
|
||||
tabindex="0"
|
||||
onmouseenter={() => (hoveredInstanceId = id)}
|
||||
onmouseenter={() =>
|
||||
(hoveredInstanceId = item.instanceId ?? id)}
|
||||
onmouseleave={() => (hoveredInstanceId = null)}
|
||||
onclick={() => {
|
||||
if (
|
||||
@@ -2108,7 +2212,10 @@
|
||||
>
|
||||
</div>
|
||||
<button
|
||||
onclick={() => deleteInstance(id)}
|
||||
onclick={() =>
|
||||
item.isMetaInstance
|
||||
? deleteMetaInstance(id)
|
||||
: deleteInstance(id)}
|
||||
class="text-xs px-2 py-1 font-mono tracking-wider uppercase border border-red-500/30 text-red-400 hover:bg-red-500/20 hover:text-red-400 hover:border-red-500/50 transition-all duration-200 cursor-pointer"
|
||||
>
|
||||
DELETE
|
||||
@@ -2118,7 +2225,7 @@
|
||||
<div
|
||||
class="text-exo-yellow text-xs font-mono tracking-wide truncate"
|
||||
>
|
||||
{getInstanceModelId(instance)}
|
||||
{instanceModelId}
|
||||
</div>
|
||||
<div class="text-white/60 text-xs font-mono">
|
||||
Strategy: <span class="text-white/80"
|
||||
@@ -2820,31 +2927,54 @@
|
||||
<div
|
||||
class="space-y-3 max-h-72 xl:max-h-96 overflow-y-auto overflow-x-hidden py-px pr-1"
|
||||
>
|
||||
{#each Object.entries(instanceData) as [id, instance]}
|
||||
{@const downloadInfo = getInstanceDownloadStatus(
|
||||
id,
|
||||
instance,
|
||||
)}
|
||||
{#each unifiedDisplayItems() as item (item.id)}
|
||||
{@const id = item.id}
|
||||
{@const instance = item.instance}
|
||||
{@const downloadInfo = instance
|
||||
? getInstanceDownloadStatus(
|
||||
item.instanceId ?? id,
|
||||
instance,
|
||||
)
|
||||
: {
|
||||
statusText: "PLACING",
|
||||
statusClass: "starting",
|
||||
isDownloading: false,
|
||||
isFailed: false,
|
||||
progress: null,
|
||||
perNode: [],
|
||||
errorMessage: null,
|
||||
}}
|
||||
{@const statusText = downloadInfo.statusText}
|
||||
{@const isDownloading = downloadInfo.isDownloading}
|
||||
{@const isFailed = statusText === "FAILED"}
|
||||
{@const isLoading =
|
||||
statusText === "LOADING" ||
|
||||
statusText === "WARMING UP" ||
|
||||
statusText === "WAITING"}
|
||||
statusText === "WAITING" ||
|
||||
statusText === "PLACING"}
|
||||
{@const isReady =
|
||||
statusText === "READY" || statusText === "LOADED"}
|
||||
{@const isRunning = statusText === "RUNNING"}
|
||||
<!-- Instance Card -->
|
||||
{@const instanceModelId = getInstanceModelId(instance)}
|
||||
{@const instanceInfo = getInstanceInfo(instance)}
|
||||
{@const instanceConnections =
|
||||
getInstanceConnections(instance)}
|
||||
{@const instanceModelId = item.modelId}
|
||||
{@const instanceInfo = instance
|
||||
? getInstanceInfo(instance)
|
||||
: {
|
||||
instanceType: "Unknown",
|
||||
sharding: "Unknown",
|
||||
nodeNames: [],
|
||||
nodeIds: [],
|
||||
nodeCount: 0,
|
||||
}}
|
||||
{@const instanceConnections = instance
|
||||
? getInstanceConnections(instance)
|
||||
: []}
|
||||
<div
|
||||
class="relative group cursor-pointer"
|
||||
role="button"
|
||||
tabindex="0"
|
||||
onmouseenter={() => (hoveredInstanceId = id)}
|
||||
onmouseenter={() =>
|
||||
(hoveredInstanceId = item.instanceId ?? id)}
|
||||
onmouseleave={() => (hoveredInstanceId = null)}
|
||||
onclick={() => {
|
||||
if (
|
||||
@@ -2943,7 +3073,10 @@
|
||||
>
|
||||
</div>
|
||||
<button
|
||||
onclick={() => deleteInstance(id)}
|
||||
onclick={() =>
|
||||
item.isMetaInstance
|
||||
? deleteMetaInstance(id)
|
||||
: deleteInstance(id)}
|
||||
class="text-xs px-2 py-1 font-mono tracking-wider uppercase border border-red-500/30 text-red-400 hover:bg-red-500/20 hover:text-red-400 hover:border-red-500/50 transition-all duration-200 cursor-pointer"
|
||||
>
|
||||
DELETE
|
||||
@@ -2953,7 +3086,7 @@
|
||||
<div
|
||||
class="text-exo-yellow text-xs font-mono tracking-wide truncate"
|
||||
>
|
||||
{getInstanceModelId(instance)}
|
||||
{instanceModelId}
|
||||
</div>
|
||||
<div class="text-white/60 text-xs font-mono">
|
||||
Strategy: <span class="text-white/80"
|
||||
|
||||
@@ -71,8 +71,11 @@ from exo.shared.types.api import (
|
||||
ChatCompletionResponse,
|
||||
CreateInstanceParams,
|
||||
CreateInstanceResponse,
|
||||
CreateMetaInstanceParams,
|
||||
CreateMetaInstanceResponse,
|
||||
DeleteDownloadResponse,
|
||||
DeleteInstanceResponse,
|
||||
DeleteMetaInstanceResponse,
|
||||
ErrorInfo,
|
||||
ErrorResponse,
|
||||
FinishReason,
|
||||
@@ -115,8 +118,10 @@ from exo.shared.types.claude_api import (
|
||||
from exo.shared.types.commands import (
|
||||
Command,
|
||||
CreateInstance,
|
||||
CreateMetaInstance,
|
||||
DeleteDownload,
|
||||
DeleteInstance,
|
||||
DeleteMetaInstance,
|
||||
DownloadCommand,
|
||||
ForwarderCommand,
|
||||
ForwarderDownloadCommand,
|
||||
@@ -137,6 +142,7 @@ from exo.shared.types.events import (
|
||||
TracesMerged,
|
||||
)
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.meta_instance import MetaInstance, MetaInstanceId
|
||||
from exo.shared.types.openai_responses import (
|
||||
ResponsesRequest,
|
||||
ResponsesResponse,
|
||||
@@ -275,6 +281,8 @@ class API:
|
||||
self.app.get("/instance/previews")(self.get_placement_previews)
|
||||
self.app.get("/instance/{instance_id}")(self.get_instance)
|
||||
self.app.delete("/instance/{instance_id}")(self.delete_instance)
|
||||
self.app.post("/meta_instance")(self.create_meta_instance)
|
||||
self.app.delete("/meta_instance/{meta_instance_id}")(self.delete_meta_instance)
|
||||
self.app.get("/models")(self.get_models)
|
||||
self.app.get("/v1/models")(self.get_models)
|
||||
self.app.post("/models/add")(self.add_custom_model)
|
||||
@@ -521,6 +529,44 @@ class API:
|
||||
instance_id=instance_id,
|
||||
)
|
||||
|
||||
async def create_meta_instance(
|
||||
self, payload: CreateMetaInstanceParams
|
||||
) -> CreateMetaInstanceResponse:
|
||||
meta_instance = MetaInstance(
|
||||
model_id=payload.model_id,
|
||||
sharding=payload.sharding,
|
||||
instance_meta=payload.instance_meta,
|
||||
min_nodes=payload.min_nodes,
|
||||
node_ids=frozenset(payload.node_ids) if payload.node_ids else None,
|
||||
)
|
||||
command = CreateMetaInstance(meta_instance=meta_instance)
|
||||
await self._send(command)
|
||||
return CreateMetaInstanceResponse(
|
||||
message="Command received.",
|
||||
command_id=command.command_id,
|
||||
meta_instance_id=meta_instance.meta_instance_id,
|
||||
)
|
||||
|
||||
async def delete_meta_instance(
|
||||
self, meta_instance_id: MetaInstanceId
|
||||
) -> DeleteMetaInstanceResponse:
|
||||
meta = self.state.meta_instances.get(meta_instance_id)
|
||||
if not meta:
|
||||
raise HTTPException(status_code=404, detail="MetaInstance not found")
|
||||
|
||||
# Delete the explicitly bound backing instance
|
||||
backing_id = self.state.meta_instance_backing.get(meta_instance_id)
|
||||
if backing_id:
|
||||
await self._send(DeleteInstance(instance_id=backing_id))
|
||||
|
||||
command = DeleteMetaInstance(meta_instance_id=meta_instance_id)
|
||||
await self._send(command)
|
||||
return DeleteMetaInstanceResponse(
|
||||
message="Command received.",
|
||||
command_id=command.command_id,
|
||||
meta_instance_id=meta_instance_id,
|
||||
)
|
||||
|
||||
async def _token_chunk_stream(
|
||||
self, command_id: CommandId
|
||||
) -> AsyncGenerator[ErrorChunk | ToolCallChunk | TokenChunk, None]:
|
||||
|
||||
@@ -13,12 +13,14 @@ from exo.master.placement import (
|
||||
place_instance,
|
||||
)
|
||||
from exo.master.reconcile import (
|
||||
find_satisfying_instance,
|
||||
find_unsatisfied_meta_instances,
|
||||
instance_connections_healthy,
|
||||
try_place_for_meta_instance,
|
||||
)
|
||||
from exo.shared.apply import apply
|
||||
from exo.shared.constants import EXO_EVENT_LOG_DIR, EXO_TRACING_ENABLED
|
||||
from exo.shared.models.model_cards import ModelCard
|
||||
from exo.shared.types.commands import (
|
||||
CreateInstance,
|
||||
CreateMetaInstance,
|
||||
@@ -42,6 +44,7 @@ from exo.shared.types.events import (
|
||||
IndexedEvent,
|
||||
InputChunkReceived,
|
||||
InstanceDeleted,
|
||||
MetaInstanceBound,
|
||||
MetaInstanceCreated,
|
||||
MetaInstanceDeleted,
|
||||
NodeGatheredInfo,
|
||||
@@ -294,6 +297,20 @@ class Master:
|
||||
generated_events.append(
|
||||
MetaInstanceCreated(meta_instance=command.meta_instance)
|
||||
)
|
||||
# Immediate placement attempt for responsiveness
|
||||
model_card = await ModelCard.load(
|
||||
command.meta_instance.model_id
|
||||
)
|
||||
generated_events.extend(
|
||||
try_place_for_meta_instance(
|
||||
command.meta_instance,
|
||||
model_card,
|
||||
self.state.topology,
|
||||
self.state.instances,
|
||||
self.state.node_memory,
|
||||
self.state.node_network,
|
||||
)
|
||||
)
|
||||
case DeleteMetaInstance():
|
||||
generated_events.append(
|
||||
MetaInstanceDeleted(
|
||||
@@ -391,10 +408,31 @@ class Master:
|
||||
self.state.meta_instances,
|
||||
self.state.instances,
|
||||
self.state.topology,
|
||||
self.state.meta_instance_backing,
|
||||
)
|
||||
# Instances already bound by other MetaInstances
|
||||
already_bound = frozenset(self.state.meta_instance_backing.values())
|
||||
for meta_instance in unsatisfied:
|
||||
# Try to bind to an existing unbound instance first
|
||||
existing = find_satisfying_instance(
|
||||
meta_instance,
|
||||
self.state.instances,
|
||||
self.state.topology,
|
||||
exclude=already_bound,
|
||||
)
|
||||
if existing is not None:
|
||||
await self._apply_and_broadcast(
|
||||
MetaInstanceBound(
|
||||
meta_instance_id=meta_instance.meta_instance_id,
|
||||
instance_id=existing,
|
||||
)
|
||||
)
|
||||
continue
|
||||
# Otherwise, place a new instance
|
||||
model_card = await ModelCard.load(meta_instance.model_id)
|
||||
events = try_place_for_meta_instance(
|
||||
meta_instance,
|
||||
model_card,
|
||||
self.state.topology,
|
||||
self.state.instances,
|
||||
self.state.node_memory,
|
||||
|
||||
+38
-12
@@ -3,10 +3,11 @@ from collections.abc import Mapping, Sequence
|
||||
from loguru import logger
|
||||
|
||||
from exo.master.placement import get_transition_events, place_instance
|
||||
from exo.shared.models.model_cards import ModelCard
|
||||
from exo.shared.topology import Topology
|
||||
from exo.shared.types.commands import PlaceInstance
|
||||
from exo.shared.types.common import NodeId
|
||||
from exo.shared.types.events import Event
|
||||
from exo.shared.types.events import Event, MetaInstanceBound
|
||||
from exo.shared.types.meta_instance import MetaInstance, MetaInstanceId
|
||||
from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo
|
||||
from exo.shared.types.topology import RDMAConnection, SocketConnection
|
||||
@@ -91,7 +92,7 @@ def instance_satisfies_meta_instance(
|
||||
This is a pure constraint check (model, min_nodes, node_ids).
|
||||
Use ``instance_connections_healthy`` separately for topology health.
|
||||
"""
|
||||
if instance.shard_assignments.model_id != meta_instance.model_card.model_id:
|
||||
if instance.shard_assignments.model_id != meta_instance.model_id:
|
||||
return False
|
||||
|
||||
instance_nodes = set(instance.shard_assignments.node_to_runner.keys())
|
||||
@@ -108,9 +109,13 @@ def find_satisfying_instance(
|
||||
meta_instance: MetaInstance,
|
||||
instances: Mapping[InstanceId, Instance],
|
||||
topology: Topology,
|
||||
*,
|
||||
exclude: frozenset[InstanceId] = frozenset(),
|
||||
) -> InstanceId | None:
|
||||
"""Find an existing instance that is healthy and satisfies a meta-instance's constraints."""
|
||||
for instance_id, instance in instances.items():
|
||||
if instance_id in exclude:
|
||||
continue
|
||||
if instance_connections_healthy(
|
||||
instance, topology
|
||||
) and instance_satisfies_meta_instance(meta_instance, instance):
|
||||
@@ -122,17 +127,25 @@ def find_unsatisfied_meta_instances(
|
||||
meta_instances: Mapping[MetaInstanceId, MetaInstance],
|
||||
instances: Mapping[InstanceId, Instance],
|
||||
topology: Topology,
|
||||
meta_instance_backing: Mapping[MetaInstanceId, InstanceId],
|
||||
) -> Sequence[MetaInstance]:
|
||||
"""Return meta-instances that have no healthy, satisfying backing instance."""
|
||||
return [
|
||||
meta_instance
|
||||
for meta_instance in meta_instances.values()
|
||||
if find_satisfying_instance(meta_instance, instances, topology) is None
|
||||
]
|
||||
"""Return meta-instances whose bound backing instance is missing or unhealthy."""
|
||||
unsatisfied: list[MetaInstance] = []
|
||||
for meta_id, meta_instance in meta_instances.items():
|
||||
bound_id = meta_instance_backing.get(meta_id)
|
||||
if bound_id is not None:
|
||||
bound_instance = instances.get(bound_id)
|
||||
if bound_instance is not None and instance_connections_healthy(
|
||||
bound_instance, topology
|
||||
):
|
||||
continue # bound and healthy
|
||||
unsatisfied.append(meta_instance)
|
||||
return unsatisfied
|
||||
|
||||
|
||||
def try_place_for_meta_instance(
|
||||
meta_instance: MetaInstance,
|
||||
model_card: ModelCard,
|
||||
topology: Topology,
|
||||
current_instances: Mapping[InstanceId, Instance],
|
||||
node_memory: Mapping[NodeId, MemoryUsage],
|
||||
@@ -140,10 +153,10 @@ def try_place_for_meta_instance(
|
||||
) -> Sequence[Event]:
|
||||
"""Try to place an instance satisfying the meta-instance constraints.
|
||||
|
||||
Returns InstanceCreated events on success, empty sequence on failure.
|
||||
Returns InstanceCreated + MetaInstanceBound events on success, empty sequence on failure.
|
||||
"""
|
||||
command = PlaceInstance(
|
||||
model_card=meta_instance.model_card,
|
||||
model_card=model_card,
|
||||
sharding=meta_instance.sharding,
|
||||
instance_meta=meta_instance.instance_meta,
|
||||
min_nodes=meta_instance.min_nodes,
|
||||
@@ -159,9 +172,22 @@ def try_place_for_meta_instance(
|
||||
set(meta_instance.node_ids) if meta_instance.node_ids else None
|
||||
),
|
||||
)
|
||||
return list(get_transition_events(current_instances, target_instances))
|
||||
events: list[Event] = list(
|
||||
get_transition_events(current_instances, target_instances)
|
||||
)
|
||||
# Find the newly created instance and bind it
|
||||
new_instance_ids = set(target_instances.keys()) - set(current_instances.keys())
|
||||
if new_instance_ids:
|
||||
new_id = next(iter(new_instance_ids))
|
||||
events.append(
|
||||
MetaInstanceBound(
|
||||
meta_instance_id=meta_instance.meta_instance_id,
|
||||
instance_id=new_id,
|
||||
)
|
||||
)
|
||||
return events
|
||||
except ValueError as e:
|
||||
logger.debug(
|
||||
f"MetaInstance placement not possible for {meta_instance.model_card.model_id}: {e}"
|
||||
f"MetaInstance placement not possible for {meta_instance.model_id}: {e}"
|
||||
)
|
||||
return []
|
||||
|
||||
@@ -10,6 +10,7 @@ from exo.shared.topology import Topology
|
||||
from exo.shared.types.common import Host, NodeId
|
||||
from exo.shared.types.events import (
|
||||
IndexedEvent,
|
||||
InstanceCreated,
|
||||
MetaInstanceCreated,
|
||||
MetaInstanceDeleted,
|
||||
)
|
||||
@@ -84,7 +85,7 @@ def _meta_instance(
|
||||
) -> MetaInstance:
|
||||
return MetaInstance(
|
||||
meta_instance_id=meta_instance_id or MetaInstanceId(),
|
||||
model_card=_model_card(model_id),
|
||||
model_id=ModelId(model_id),
|
||||
min_nodes=min_nodes,
|
||||
node_ids=node_ids,
|
||||
)
|
||||
@@ -339,7 +340,7 @@ def test_find_multiple_could_match():
|
||||
|
||||
|
||||
def test_unsatisfied_no_meta_instances():
|
||||
result = find_unsatisfied_meta_instances({}, {}, Topology())
|
||||
result = find_unsatisfied_meta_instances({}, {}, Topology(), {})
|
||||
assert list(result) == []
|
||||
|
||||
|
||||
@@ -347,8 +348,12 @@ def test_unsatisfied_one_satisfied():
|
||||
meta = _meta_instance()
|
||||
id_a, inst_a = _instance()
|
||||
topology = _topology("node-a")
|
||||
# Bound via backing map
|
||||
result = find_unsatisfied_meta_instances(
|
||||
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
|
||||
{meta.meta_instance_id: meta},
|
||||
{id_a: inst_a},
|
||||
topology,
|
||||
{meta.meta_instance_id: id_a},
|
||||
)
|
||||
assert list(result) == []
|
||||
|
||||
@@ -358,7 +363,7 @@ def test_unsatisfied_one_not_satisfied():
|
||||
id_a, inst_a = _instance("test-org/model-y")
|
||||
topology = _topology("node-a")
|
||||
result = find_unsatisfied_meta_instances(
|
||||
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
|
||||
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology, {}
|
||||
)
|
||||
assert list(result) == [meta]
|
||||
|
||||
@@ -375,6 +380,7 @@ def test_unsatisfied_mix():
|
||||
},
|
||||
{id_a: inst_a},
|
||||
topology,
|
||||
{meta_satisfied.meta_instance_id: id_a},
|
||||
)
|
||||
assert list(result) == [meta_unsatisfied]
|
||||
|
||||
@@ -384,7 +390,10 @@ def test_unsatisfied_node_disconnect():
|
||||
id_a, inst_a = _instance(node_ids=["node-a", "node-b"])
|
||||
topology = _topology("node-a") # node-b disconnected
|
||||
result = find_unsatisfied_meta_instances(
|
||||
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
|
||||
{meta.meta_instance_id: meta},
|
||||
{id_a: inst_a},
|
||||
topology,
|
||||
{meta.meta_instance_id: id_a},
|
||||
)
|
||||
assert list(result) == [meta]
|
||||
|
||||
@@ -395,7 +404,10 @@ def test_unsatisfied_edge_break():
|
||||
id_a, inst_a = _instance(node_ids=["node-a", "node-b"])
|
||||
topology = _topology("node-a", "node-b", connect=False) # nodes present, no edges
|
||||
result = find_unsatisfied_meta_instances(
|
||||
{meta.meta_instance_id: meta}, {id_a: inst_a}, topology
|
||||
{meta.meta_instance_id: meta},
|
||||
{id_a: inst_a},
|
||||
topology,
|
||||
{meta.meta_instance_id: id_a},
|
||||
)
|
||||
assert list(result) == [meta]
|
||||
|
||||
@@ -406,14 +418,152 @@ def test_unsatisfied_idempotent():
|
||||
meta_instances = {meta.meta_instance_id: meta}
|
||||
instances: dict[InstanceId, MlxRingInstance] = {}
|
||||
result_1 = list(
|
||||
find_unsatisfied_meta_instances(meta_instances, instances, topology)
|
||||
find_unsatisfied_meta_instances(meta_instances, instances, topology, {})
|
||||
)
|
||||
result_2 = list(
|
||||
find_unsatisfied_meta_instances(meta_instances, instances, topology)
|
||||
find_unsatisfied_meta_instances(meta_instances, instances, topology, {})
|
||||
)
|
||||
assert result_1 == result_2
|
||||
|
||||
|
||||
def test_unsatisfied_exclusive_binding():
|
||||
"""Two MetaInstances for the same model: one is bound, the other is unsatisfied."""
|
||||
meta_a = _meta_instance("test-org/model-x")
|
||||
meta_b = _meta_instance("test-org/model-x")
|
||||
id_inst, inst = _instance("test-org/model-x")
|
||||
topology = _topology("node-a")
|
||||
# meta_a is bound to the only instance → meta_b is unsatisfied
|
||||
backing = {meta_a.meta_instance_id: id_inst}
|
||||
result = find_unsatisfied_meta_instances(
|
||||
{
|
||||
meta_a.meta_instance_id: meta_a,
|
||||
meta_b.meta_instance_id: meta_b,
|
||||
},
|
||||
{id_inst: inst},
|
||||
topology,
|
||||
backing,
|
||||
)
|
||||
assert list(result) == [meta_b]
|
||||
|
||||
|
||||
def test_find_satisfying_instance_exclude():
|
||||
"""find_satisfying_instance should skip instances in the exclude set."""
|
||||
meta = _meta_instance()
|
||||
id_a, inst_a = _instance()
|
||||
id_b, inst_b = _instance()
|
||||
topology = _topology("node-a")
|
||||
# Without exclude, first match returned
|
||||
result = find_satisfying_instance(meta, {id_a: inst_a, id_b: inst_b}, topology)
|
||||
assert result is not None
|
||||
# Exclude the first result → should get the other one
|
||||
result2 = find_satisfying_instance(
|
||||
meta, {id_a: inst_a, id_b: inst_b}, topology, exclude=frozenset({result})
|
||||
)
|
||||
assert result2 is not None
|
||||
assert result2 != result
|
||||
# Exclude both → nothing found
|
||||
result3 = find_satisfying_instance(
|
||||
meta, {id_a: inst_a, id_b: inst_b}, topology, exclude=frozenset({id_a, id_b})
|
||||
)
|
||||
assert result3 is None
|
||||
|
||||
|
||||
def test_apply_meta_instance_bound():
|
||||
"""MetaInstanceBound event should populate the backing map."""
|
||||
from exo.shared.types.events import MetaInstanceBound
|
||||
|
||||
state = State()
|
||||
meta = _meta_instance()
|
||||
id_a, inst_a = _instance()
|
||||
# First, create meta instance and instance in state
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=0, event=MetaInstanceCreated(meta_instance=meta)),
|
||||
)
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=1, event=InstanceCreated(instance=inst_a)),
|
||||
)
|
||||
# Bind them
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(
|
||||
idx=2,
|
||||
event=MetaInstanceBound(
|
||||
meta_instance_id=meta.meta_instance_id, instance_id=id_a
|
||||
),
|
||||
),
|
||||
)
|
||||
assert state.meta_instance_backing[meta.meta_instance_id] == id_a
|
||||
|
||||
|
||||
def test_apply_instance_deleted_clears_backing():
|
||||
"""Deleting an instance should remove it from the backing map."""
|
||||
from exo.shared.types.events import InstanceDeleted, MetaInstanceBound
|
||||
|
||||
state = State()
|
||||
meta = _meta_instance()
|
||||
id_a, inst_a = _instance()
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=0, event=MetaInstanceCreated(meta_instance=meta)),
|
||||
)
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=1, event=InstanceCreated(instance=inst_a)),
|
||||
)
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(
|
||||
idx=2,
|
||||
event=MetaInstanceBound(
|
||||
meta_instance_id=meta.meta_instance_id, instance_id=id_a
|
||||
),
|
||||
),
|
||||
)
|
||||
assert meta.meta_instance_id in state.meta_instance_backing
|
||||
# Delete the instance
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=3, event=InstanceDeleted(instance_id=id_a)),
|
||||
)
|
||||
assert meta.meta_instance_id not in state.meta_instance_backing
|
||||
|
||||
|
||||
def test_apply_meta_instance_deleted_clears_backing():
|
||||
"""Deleting a MetaInstance should remove its entry from the backing map."""
|
||||
from exo.shared.types.events import MetaInstanceBound
|
||||
|
||||
state = State()
|
||||
meta = _meta_instance()
|
||||
id_a, inst_a = _instance()
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=0, event=MetaInstanceCreated(meta_instance=meta)),
|
||||
)
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(idx=1, event=InstanceCreated(instance=inst_a)),
|
||||
)
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(
|
||||
idx=2,
|
||||
event=MetaInstanceBound(
|
||||
meta_instance_id=meta.meta_instance_id, instance_id=id_a
|
||||
),
|
||||
),
|
||||
)
|
||||
# Delete the meta instance
|
||||
state = apply(
|
||||
state,
|
||||
IndexedEvent(
|
||||
idx=3, event=MetaInstanceDeleted(meta_instance_id=meta.meta_instance_id)
|
||||
),
|
||||
)
|
||||
assert meta.meta_instance_id not in state.meta_instance_backing
|
||||
|
||||
|
||||
# --- apply handlers ---
|
||||
|
||||
|
||||
|
||||
+27
-2
@@ -12,6 +12,7 @@ from exo.shared.types.events import (
|
||||
InputChunkReceived,
|
||||
InstanceCreated,
|
||||
InstanceDeleted,
|
||||
MetaInstanceBound,
|
||||
MetaInstanceCreated,
|
||||
MetaInstanceDeleted,
|
||||
NodeDownloadProgress,
|
||||
@@ -76,6 +77,8 @@ def event_apply(event: Event, state: State) -> State:
|
||||
return apply_meta_instance_created(event, state)
|
||||
case MetaInstanceDeleted():
|
||||
return apply_meta_instance_deleted(event, state)
|
||||
case MetaInstanceBound():
|
||||
return apply_meta_instance_bound(event, state)
|
||||
case NodeTimedOut():
|
||||
return apply_node_timed_out(event, state)
|
||||
case NodeDownloadProgress():
|
||||
@@ -191,7 +194,14 @@ def apply_instance_deleted(event: InstanceDeleted, state: State) -> State:
|
||||
new_instances: Mapping[InstanceId, Instance] = {
|
||||
iid: inst for iid, inst in state.instances.items() if iid != event.instance_id
|
||||
}
|
||||
return state.model_copy(update={"instances": new_instances})
|
||||
new_backing: Mapping[MetaInstanceId, InstanceId] = {
|
||||
mid: iid
|
||||
for mid, iid in state.meta_instance_backing.items()
|
||||
if iid != event.instance_id
|
||||
}
|
||||
return state.model_copy(
|
||||
update={"instances": new_instances, "meta_instance_backing": new_backing}
|
||||
)
|
||||
|
||||
|
||||
def apply_meta_instance_created(event: MetaInstanceCreated, state: State) -> State:
|
||||
@@ -208,7 +218,22 @@ def apply_meta_instance_deleted(event: MetaInstanceDeleted, state: State) -> Sta
|
||||
for mid, mi in state.meta_instances.items()
|
||||
if mid != event.meta_instance_id
|
||||
}
|
||||
return state.model_copy(update={"meta_instances": new_meta})
|
||||
new_backing: Mapping[MetaInstanceId, InstanceId] = {
|
||||
mid: iid
|
||||
for mid, iid in state.meta_instance_backing.items()
|
||||
if mid != event.meta_instance_id
|
||||
}
|
||||
return state.model_copy(
|
||||
update={"meta_instances": new_meta, "meta_instance_backing": new_backing}
|
||||
)
|
||||
|
||||
|
||||
def apply_meta_instance_bound(event: MetaInstanceBound, state: State) -> State:
|
||||
new_backing: Mapping[MetaInstanceId, InstanceId] = {
|
||||
**state.meta_instance_backing,
|
||||
event.meta_instance_id: event.instance_id,
|
||||
}
|
||||
return state.model_copy(update={"meta_instance_backing": new_backing})
|
||||
|
||||
|
||||
def apply_runner_status_updated(event: RunnerStatusUpdated, state: State) -> State:
|
||||
|
||||
@@ -9,6 +9,7 @@ from pydantic_core import PydanticUseDefault
|
||||
from exo.shared.models.model_cards import ModelCard, ModelId
|
||||
from exo.shared.types.common import CommandId, NodeId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.meta_instance import MetaInstanceId
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
|
||||
from exo.shared.types.worker.shards import Sharding, ShardMetadata
|
||||
from exo.utils.pydantic_ext import CamelCaseModel
|
||||
@@ -269,6 +270,33 @@ class DeleteInstanceResponse(BaseModel):
|
||||
instance_id: InstanceId
|
||||
|
||||
|
||||
class CreateMetaInstanceParams(BaseModel):
|
||||
model_id: ModelId
|
||||
sharding: Sharding = Sharding.Pipeline
|
||||
instance_meta: InstanceMeta = InstanceMeta.MlxRing
|
||||
min_nodes: int = 1
|
||||
node_ids: list[NodeId] | None = None
|
||||
|
||||
@field_validator("sharding", "instance_meta", mode="plain")
|
||||
@classmethod
|
||||
def use_default(cls, v: object):
|
||||
if not v or not isinstance(v, (Sharding, InstanceMeta)):
|
||||
raise PydanticUseDefault()
|
||||
return v
|
||||
|
||||
|
||||
class CreateMetaInstanceResponse(BaseModel):
|
||||
message: str
|
||||
command_id: CommandId
|
||||
meta_instance_id: MetaInstanceId
|
||||
|
||||
|
||||
class DeleteMetaInstanceResponse(BaseModel):
|
||||
message: str
|
||||
command_id: CommandId
|
||||
meta_instance_id: MetaInstanceId
|
||||
|
||||
|
||||
class AdvancedImageParams(BaseModel):
|
||||
seed: Annotated[int, Field(ge=0)] | None = None
|
||||
num_inference_steps: Annotated[int, Field(ge=1, le=100)] | None = None
|
||||
|
||||
@@ -77,6 +77,11 @@ class MetaInstanceDeleted(BaseEvent):
|
||||
meta_instance_id: MetaInstanceId
|
||||
|
||||
|
||||
class MetaInstanceBound(BaseEvent):
|
||||
meta_instance_id: MetaInstanceId
|
||||
instance_id: InstanceId
|
||||
|
||||
|
||||
class RunnerStatusUpdated(BaseEvent):
|
||||
runner_id: RunnerId
|
||||
runner_status: RunnerStatus
|
||||
@@ -152,6 +157,7 @@ Event = (
|
||||
| InstanceDeleted
|
||||
| MetaInstanceCreated
|
||||
| MetaInstanceDeleted
|
||||
| MetaInstanceBound
|
||||
| RunnerStatusUpdated
|
||||
| RunnerDeleted
|
||||
| NodeTimedOut
|
||||
|
||||
@@ -2,7 +2,7 @@ from typing import final
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from exo.shared.models.model_cards import ModelCard
|
||||
from exo.shared.models.model_cards import ModelId
|
||||
from exo.shared.types.common import Id, NodeId
|
||||
from exo.shared.types.worker.instances import InstanceMeta
|
||||
from exo.shared.types.worker.shards import Sharding
|
||||
@@ -18,7 +18,7 @@ class MetaInstance(FrozenModel):
|
||||
"""Declarative constraint: ensure an instance matching these parameters always exists."""
|
||||
|
||||
meta_instance_id: MetaInstanceId = Field(default_factory=MetaInstanceId)
|
||||
model_card: ModelCard
|
||||
model_id: ModelId
|
||||
sharding: Sharding = Sharding.Pipeline
|
||||
instance_meta: InstanceMeta = InstanceMeta.MlxRing
|
||||
min_nodes: int = 1
|
||||
|
||||
@@ -41,6 +41,7 @@ class State(CamelCaseModel):
|
||||
)
|
||||
instances: Mapping[InstanceId, Instance] = {}
|
||||
meta_instances: Mapping[MetaInstanceId, MetaInstance] = {}
|
||||
meta_instance_backing: Mapping[MetaInstanceId, InstanceId] = {}
|
||||
runners: Mapping[RunnerId, RunnerStatus] = {}
|
||||
downloads: Mapping[NodeId, Sequence[DownloadProgress]] = {}
|
||||
tasks: Mapping[TaskId, Task] = {}
|
||||
|
||||
Reference in New Issue
Block a user