diff --git a/.github/workflows/pipeline.yml b/.github/workflows/pipeline.yml index 0f908f1b..c2589453 100644 --- a/.github/workflows/pipeline.yml +++ b/.github/workflows/pipeline.yml @@ -142,4 +142,6 @@ jobs: # Run pytest outside sandbox (needs GPU access for MLX) export HOME="$RUNNER_TEMP" export EXO_TESTS=1 - EXO_RESOURCES_DIR="$PWD/resources" $TEST_ENV/bin/python -m pytest src -m "not slow" --import-mode=importlib + export EXO_DASHBOARD_DIR="$PWD/dashboard/" + export EXO_RESOURCES_DIR="$PWD/resources" + $TEST_ENV/bin/python -m pytest src -m "not slow" --import-mode=importlib diff --git a/.gitignore b/.gitignore index d0b8299e..139fc326 100644 --- a/.gitignore +++ b/.gitignore @@ -31,3 +31,4 @@ dashboard/.svelte-kit/ # host config snapshots hosts_*.json +.swp diff --git a/.mlx_typings/mlx_lm/tokenizer_utils.pyi b/.mlx_typings/mlx_lm/tokenizer_utils.pyi index 251e3d28..83eb4e33 100644 --- a/.mlx_typings/mlx_lm/tokenizer_utils.pyi +++ b/.mlx_typings/mlx_lm/tokenizer_utils.pyi @@ -108,6 +108,7 @@ class TokenizerWrapper: _tokenizer: PreTrainedTokenizerFast eos_token_id: int | None eos_token: str | None + eos_token_ids: list[int] | set[int] | None bos_token_id: int | None bos_token: str | None vocab_size: int @@ -117,7 +118,7 @@ class TokenizerWrapper: self, tokenizer: Any, detokenizer_class: Any = ..., - eos_token_ids: list[int] | None = ..., + eos_token_ids: list[int] | set[int] | None = ..., chat_template: Any = ..., tool_parser: Any = ..., tool_call_start: str | None = ..., diff --git a/app/EXO/EXO/EXOApp.swift b/app/EXO/EXO/EXOApp.swift index 7669c408..3aff58c3 100644 --- a/app/EXO/EXO/EXOApp.swift +++ b/app/EXO/EXO/EXOApp.swift @@ -14,7 +14,6 @@ import SwiftUI import UserNotifications import os.log -@main struct EXOApp: App { @StateObject private var controller: ExoProcessController @StateObject private var stateService: ClusterStateService diff --git a/app/EXO/EXO/Services/NetworkSetupHelper.swift b/app/EXO/EXO/Services/NetworkSetupHelper.swift index 82cb82d4..5428ee8e 100644 --- a/app/EXO/EXO/Services/NetworkSetupHelper.swift +++ b/app/EXO/EXO/Services/NetworkSetupHelper.swift @@ -288,6 +288,61 @@ enum NetworkSetupHelper { """ } + /// Direct install without GUI (requires root). + /// Returns true on success, false on failure. + static func installDirectly() -> Bool { + let script = makeInstallerScript() + return runShellDirectly(script) + } + + /// Direct uninstall without GUI (requires root). + /// Returns true on success, false on failure. + static func uninstallDirectly() -> Bool { + let script = makeUninstallScript() + return runShellDirectly(script) + } + + /// Run a shell script directly via Process (no AppleScript, requires root). + /// Returns true on success, false on failure. + private static func runShellDirectly(_ script: String) -> Bool { + let process = Process() + process.executableURL = URL(fileURLWithPath: "/bin/bash") + process.arguments = ["-c", script] + + let outputPipe = Pipe() + let errorPipe = Pipe() + process.standardOutput = outputPipe + process.standardError = errorPipe + + do { + try process.run() + process.waitUntilExit() + + let outputData = outputPipe.fileHandleForReading.readDataToEndOfFile() + let errorData = errorPipe.fileHandleForReading.readDataToEndOfFile() + + if let output = String(data: outputData, encoding: .utf8), !output.isEmpty { + print(output) + } + if let errorOutput = String(data: errorData, encoding: .utf8), !errorOutput.isEmpty { + fputs(errorOutput, stderr) + } + + if process.terminationStatus == 0 { + logger.info("Shell script completed successfully") + return true + } else { + logger.error("Shell script failed with exit code \(process.terminationStatus)") + return false + } + } catch { + logger.error( + "Failed to run shell script: \(error.localizedDescription, privacy: .public)") + fputs("Error: \(error.localizedDescription)\n", stderr) + return false + } + } + private static func runShellAsAdmin(_ script: String) throws { let escapedScript = script diff --git a/app/EXO/EXO/main.swift b/app/EXO/EXO/main.swift new file mode 100644 index 00000000..9383981f --- /dev/null +++ b/app/EXO/EXO/main.swift @@ -0,0 +1,85 @@ +// +// main.swift +// EXO +// +// Created by Jake Hillion on 2026-02-03. +// + +import Foundation + +/// Command line options for the EXO app +enum CLICommand { + case install + case uninstall + case help + case none +} + +/// Parse command line arguments to determine the CLI command +func parseArguments() -> CLICommand { + let args = CommandLine.arguments + if args.contains("--help") || args.contains("-h") { + return .help + } + if args.contains("--install") { + return .install + } + if args.contains("--uninstall") { + return .uninstall + } + return .none +} + +/// Print usage information +func printUsage() { + let programName = (CommandLine.arguments.first as NSString?)?.lastPathComponent ?? "EXO" + print( + """ + Usage: \(programName) [OPTIONS] + + Options: + --install Install EXO network configuration (requires root) + --uninstall Uninstall EXO network configuration (requires root) + --help, -h Show this help message + + When run without options, starts the normal GUI application. + + Examples: + sudo \(programName) --install Install network components as root + sudo \(programName) --uninstall Remove network components as root + """) +} + +/// Check if running as root +func isRunningAsRoot() -> Bool { + return getuid() == 0 +} + +// Main entry point +let command = parseArguments() + +switch command { +case .help: + printUsage() + exit(0) + +case .install: + if !isRunningAsRoot() { + fputs("Error: --install requires root privileges. Run with sudo.\n", stderr) + exit(1) + } + let success = NetworkSetupHelper.installDirectly() + exit(success ? 0 : 1) + +case .uninstall: + if !isRunningAsRoot() { + fputs("Error: --uninstall requires root privileges. Run with sudo.\n", stderr) + exit(1) + } + let success = NetworkSetupHelper.uninstallDirectly() + exit(success ? 0 : 1) + +case .none: + // Start normal GUI application + EXOApp.main() +} diff --git a/dashboard/src/lib/components/ChatMessages.svelte b/dashboard/src/lib/components/ChatMessages.svelte index 15ea088d..44b9ec0d 100644 --- a/dashboard/src/lib/components/ChatMessages.svelte +++ b/dashboard/src/lib/components/ChatMessages.svelte @@ -6,11 +6,13 @@ deleteMessage, editAndRegenerate, regenerateLastResponse, + regenerateFromToken, setEditingImage, } from "$lib/stores/app.svelte"; import type { Message } from "$lib/stores/app.svelte"; import type { MessageAttachment } from "$lib/stores/app.svelte"; import MarkdownContent from "./MarkdownContent.svelte"; + import TokenHeatmap from "./TokenHeatmap.svelte"; interface Props { class?: string; @@ -99,6 +101,23 @@ let copiedMessageId = $state(null); let expandedThinkingMessageIds = $state>(new Set()); + // Uncertainty heatmap toggle + let heatmapMessageIds = $state>(new Set()); + + function toggleHeatmap(messageId: string) { + const next = new Set(heatmapMessageIds); + if (next.has(messageId)) { + next.delete(messageId); + } else { + next.add(messageId); + } + heatmapMessageIds = next; + } + + function isHeatmapVisible(messageId: string): boolean { + return heatmapMessageIds.has(messageId); + } + function formatTimestamp(timestamp: number): string { return new Date(timestamp).toLocaleTimeString("en-US", { hour12: false, @@ -548,13 +567,23 @@ > {:else if message.content || (loading && !message.attachments?.some((a) => a.type === "generated-image"))} - - {#if loading && !message.content} - + {#if isHeatmapVisible(message.id) && message.tokens && message.tokens.length > 0} + + regenerateFromToken(message.id, tokenIndex)} + /> + {:else} + + {#if loading && !message.content} + + {/if} {/if} {/if} @@ -629,6 +658,35 @@ {/if} + + {#if message.role === "assistant" && message.tokens && message.tokens.length > 0} + + {/if} + {#if message.role === "assistant" && isLastAssistantMessage(message.id) && !loading} + + + {#if hasFavorites} + + {/if} + + + + +
+ + + {#each families as family} + + {/each} + diff --git a/dashboard/src/lib/components/HuggingFaceResultItem.svelte b/dashboard/src/lib/components/HuggingFaceResultItem.svelte new file mode 100644 index 00000000..566d8e17 --- /dev/null +++ b/dashboard/src/lib/components/HuggingFaceResultItem.svelte @@ -0,0 +1,127 @@ + + +
+
+
+ {modelName} + {#if isAdded} + Added + {/if} +
+
+ {model.author} + + + + + {formatNumber(model.downloads)} + + + + + + {formatNumber(model.likes)} + +
+
+ +
+ {#if isAdded} + + {:else} + + {/if} +
+
diff --git a/dashboard/src/lib/components/ModelFilterPopover.svelte b/dashboard/src/lib/components/ModelFilterPopover.svelte new file mode 100644 index 00000000..5406618a --- /dev/null +++ b/dashboard/src/lib/components/ModelFilterPopover.svelte @@ -0,0 +1,182 @@ + + + + + +
e.stopPropagation()} + role="dialog" + aria-label="Filter options" +> +
+ +
+

Capabilities

+
+ {#each availableCapabilities as cap} + {@const isSelected = filters.capabilities.includes(cap.id)} + + {/each} +
+
+ + +
+

Model Size

+
+ {#each sizeRanges as range} + {@const isSelected = + filters.sizeRange && + filters.sizeRange.min === range.min && + filters.sizeRange.max === range.max} + + {/each} +
+
+ + + +
+
diff --git a/dashboard/src/lib/components/ModelPickerGroup.svelte b/dashboard/src/lib/components/ModelPickerGroup.svelte new file mode 100644 index 00000000..b3ad425b --- /dev/null +++ b/dashboard/src/lib/components/ModelPickerGroup.svelte @@ -0,0 +1,324 @@ + + +
+ +
{ + if (group.hasMultipleVariants) { + onToggleExpand(); + } else { + const modelId = group.variants[0]?.id; + if (modelId && canModelFit(modelId)) { + onSelectModel(modelId); + } + } + }} + role="button" + tabindex="0" + onkeydown={(e) => { + if (e.key === "Enter" || e.key === " ") { + e.preventDefault(); + if (group.hasMultipleVariants) { + onToggleExpand(); + } else { + const modelId = group.variants[0]?.id; + if (modelId && canModelFit(modelId)) { + onSelectModel(modelId); + } + } + } + }} + > + + {#if group.hasMultipleVariants} + + + + {:else} +
+ {/if} + + +
+
+ + {group.name} + + + {#each group.capabilities.filter((c) => c !== "text") as cap} + {#if cap === "thinking"} + + + + {:else if cap === "code"} + + + + {:else if cap === "vision"} + + + + + {:else if cap === "image_gen"} + + + + + + {/if} + {/each} +
+
+ + + {#if !group.hasMultipleVariants && group.smallestVariant?.storage_size_megabytes} + + {formatSize(group.smallestVariant.storage_size_megabytes)} + + {/if} + + + {#if group.hasMultipleVariants} + + {group.variants.length} variants + + {/if} + + + {#if isMainSelected} + + + + {/if} + + + + + + +
+ + + {#if isExpanded && group.hasMultipleVariants} +
+ {#each group.variants as variant} + {@const modelCanFit = canModelFit(variant.id)} + {@const isSelected = selectedModelId === variant.id} + + {/each} +
+ {/if} +
diff --git a/dashboard/src/lib/components/ModelPickerModal.svelte b/dashboard/src/lib/components/ModelPickerModal.svelte new file mode 100644 index 00000000..40827a0e --- /dev/null +++ b/dashboard/src/lib/components/ModelPickerModal.svelte @@ -0,0 +1,748 @@ + + + + +{#if isOpen} + + + + + + + + {#if infoGroup} +
(infoGroup = null)} + role="presentation" + >
+ + {/if} +{/if} diff --git a/dashboard/src/lib/components/TokenHeatmap.svelte b/dashboard/src/lib/components/TokenHeatmap.svelte new file mode 100644 index 00000000..d6de667c --- /dev/null +++ b/dashboard/src/lib/components/TokenHeatmap.svelte @@ -0,0 +1,236 @@ + + +
+ {#each tokens as tokenData, i (i)} + handleMouseEnter(e, tokenData, i)} + onmouseleave={handleMouseLeave}>{tokenData.token} + {/each} +
+ + +{#if hoveredToken} +
+
+ +
+ Token: + "{hoveredToken.token.token}" + {formatProbability(hoveredToken.token.probability)} +
+ +
+ logprob: {formatLogprob(hoveredToken.token.logprob)} +
+ + + {#if hoveredToken.token.topLogprobs.length > 0} +
+
Alternatives:
+ {#each hoveredToken.token.topLogprobs.slice(0, 5) as alt, idx (idx)} + {@const altProb = Math.exp(alt.logprob)} +
+ "{alt.token}" + {formatProbability(altProb)} +
+ {/each} +
+ {/if} + + + {#if onRegenerateFrom} + + {/if} +
+ +
+
+
+
+{/if} + + diff --git a/dashboard/src/lib/components/index.ts b/dashboard/src/lib/components/index.ts index dc8a7d76..b0424bf4 100644 --- a/dashboard/src/lib/components/index.ts +++ b/dashboard/src/lib/components/index.ts @@ -6,3 +6,9 @@ export { default as ChatSidebar } from "./ChatSidebar.svelte"; export { default as ModelCard } from "./ModelCard.svelte"; export { default as MarkdownContent } from "./MarkdownContent.svelte"; export { default as ImageParamsPanel } from "./ImageParamsPanel.svelte"; +export { default as FamilyLogos } from "./FamilyLogos.svelte"; +export { default as FamilySidebar } from "./FamilySidebar.svelte"; +export { default as HuggingFaceResultItem } from "./HuggingFaceResultItem.svelte"; +export { default as ModelFilterPopover } from "./ModelFilterPopover.svelte"; +export { default as ModelPickerGroup } from "./ModelPickerGroup.svelte"; +export { default as ModelPickerModal } from "./ModelPickerModal.svelte"; diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts index 51de6c66..6fdb0c7c 100644 --- a/dashboard/src/lib/stores/app.svelte.ts +++ b/dashboard/src/lib/stores/app.svelte.ts @@ -242,6 +242,19 @@ export interface MessageAttachment { mimeType?: string; } +export interface TopLogprob { + token: string; + logprob: number; + bytes: number[] | null; +} + +export interface TokenData { + token: string; + logprob: number; + probability: number; + topLogprobs: TopLogprob[]; +} + export interface Message { id: string; role: "user" | "assistant" | "system"; @@ -253,6 +266,7 @@ export interface Message { tps?: number; // Tokens per second (for assistant messages) requestType?: "chat" | "image-generation" | "image-editing"; sourceImageDataUrl?: string; // For image editing regeneration + tokens?: TokenData[]; } export interface Conversation { @@ -540,7 +554,18 @@ class AppStore { */ private saveConversationsToStorage() { try { - localStorage.setItem(STORAGE_KEY, JSON.stringify(this.conversations)); + // Strip tokens from messages before saving to avoid bloating localStorage + const stripped = this.conversations.map((conv) => ({ + ...conv, + messages: conv.messages.map((msg) => { + if (msg.tokens) { + const { tokens: _, ...rest } = msg; + return rest; + } + return msg; + }), + })); + localStorage.setItem(STORAGE_KEY, JSON.stringify(stripped)); } catch (error) { console.error("Failed to save conversations:", error); } @@ -1445,6 +1470,213 @@ class AppStore { } } + /** + * Regenerate response from a specific token index. + * Truncates the assistant message at the given token and re-generates from there. + */ + async regenerateFromToken( + messageId: string, + tokenIndex: number, + ): Promise { + if (this.isLoading) return; + + const targetConversationId = this.activeConversationId; + if (!targetConversationId) return; + + const msgIndex = this.messages.findIndex((m) => m.id === messageId); + if (msgIndex === -1) return; + + const msg = this.messages[msgIndex]; + if ( + msg.role !== "assistant" || + !msg.tokens || + tokenIndex >= msg.tokens.length + ) + return; + + // Keep tokens up to (not including) the specified index + const tokensToKeep = msg.tokens.slice(0, tokenIndex); + const prefixText = tokensToKeep.map((t) => t.token).join(""); + + // Remove all messages after this assistant message + this.messages = this.messages.slice(0, msgIndex + 1); + + // Update the message to show the prefix + this.messages[msgIndex].content = prefixText; + this.messages[msgIndex].tokens = tokensToKeep; + this.updateActiveConversation(); + + // Set up for continuation - modify the existing message in place + this.isLoading = true; + this.currentResponse = prefixText; + this.ttftMs = null; + this.tps = null; + this.totalTokens = tokensToKeep.length; + + try { + // Build messages for API - include the partial assistant message + const systemPrompt = { + role: "system" as const, + content: + "You are a helpful AI assistant. Respond directly and concisely. Do not show your reasoning or thought process.", + }; + + const apiMessages = [ + systemPrompt, + ...this.messages.map((m) => { + let msgContent = m.content; + if (m.attachments) { + for (const attachment of m.attachments) { + if (attachment.type === "text" && attachment.content) { + msgContent += `\n\n[File: ${attachment.name}]\n\`\`\`\n${attachment.content}\n\`\`\``; + } + } + } + return { role: m.role, content: msgContent }; + }), + ]; + + const modelToUse = this.getModelForRequest(); + if (!modelToUse) { + throw new Error("No model available"); + } + + const requestStartTime = performance.now(); + let firstTokenTime: number | null = null; + let tokenCount = tokensToKeep.length; + + const response = await fetch("/v1/chat/completions", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + model: modelToUse, + messages: apiMessages, + stream: true, + logprobs: true, + top_logprobs: 5, + }), + }); + + if (!response.ok) { + const errorText = await response.text(); + throw new Error(`API error: ${response.status} - ${errorText}`); + } + + const reader = response.body?.getReader(); + if (!reader) throw new Error("No response body"); + + let fullContent = prefixText; + const collectedTokens: TokenData[] = [...tokensToKeep]; + + interface ChatCompletionChunk { + choices?: Array<{ + delta?: { content?: string }; + logprobs?: { + content?: Array<{ + token: string; + logprob: number; + top_logprobs?: Array<{ + token: string; + logprob: number; + bytes: number[] | null; + }>; + }>; + }; + }>; + } + + await this.parseSSEStream( + reader, + targetConversationId, + (parsed) => { + const choice = parsed.choices?.[0]; + const delta = choice?.delta?.content; + + // Collect logprobs data + const logprobsContent = choice?.logprobs?.content; + if (logprobsContent) { + for (const item of logprobsContent) { + collectedTokens.push({ + token: item.token, + logprob: item.logprob, + probability: Math.exp(item.logprob), + topLogprobs: (item.top_logprobs || []).map((t) => ({ + token: t.token, + logprob: t.logprob, + bytes: t.bytes, + })), + }); + } + } + + if (delta) { + if (firstTokenTime === null) { + firstTokenTime = performance.now(); + this.ttftMs = firstTokenTime - requestStartTime; + } + + tokenCount += 1; + this.totalTokens = tokenCount; + + if (firstTokenTime !== null && tokenCount > tokensToKeep.length) { + const elapsed = performance.now() - firstTokenTime; + this.tps = ((tokenCount - tokensToKeep.length) / elapsed) * 1000; + } + + fullContent += delta; + const { displayContent, thinkingContent } = + this.stripThinkingTags(fullContent); + + if (this.activeConversationId === targetConversationId) { + this.currentResponse = displayContent; + } + + // Update existing message in place + this.updateConversationMessage( + targetConversationId, + messageId, + (m) => { + m.content = displayContent; + m.thinking = thinkingContent || undefined; + m.tokens = [...collectedTokens]; + }, + ); + this.syncActiveMessagesIfNeeded(targetConversationId); + this.persistConversation(targetConversationId); + } + }, + ); + + // Final update + if (this.conversationExists(targetConversationId)) { + const { displayContent, thinkingContent } = + this.stripThinkingTags(fullContent); + this.updateConversationMessage(targetConversationId, messageId, (m) => { + m.content = displayContent; + m.thinking = thinkingContent || undefined; + m.tokens = [...collectedTokens]; + if (this.ttftMs !== null) m.ttftMs = this.ttftMs; + if (this.tps !== null) m.tps = this.tps; + }); + this.syncActiveMessagesIfNeeded(targetConversationId); + this.persistConversation(targetConversationId); + } + } catch (error) { + console.error("Error regenerating from token:", error); + if (this.conversationExists(targetConversationId)) { + this.updateConversationMessage(targetConversationId, messageId, (m) => { + m.content = `${prefixText}\n\nError: ${error instanceof Error ? error.message : "Unknown error"}`; + }); + this.syncActiveMessagesIfNeeded(targetConversationId); + this.persistConversation(targetConversationId); + } + } finally { + this.isLoading = false; + this.currentResponse = ""; + this.saveConversationsToStorage(); + } + } + /** * Helper method to regenerate a chat completion response */ @@ -1513,6 +1745,8 @@ class AppStore { model: modelToUse, messages: apiMessages, stream: true, + logprobs: true, + top_logprobs: 5, }), }); @@ -1527,16 +1761,49 @@ class AppStore { } let streamedContent = ""; + const collectedTokens: TokenData[] = []; interface ChatCompletionChunk { - choices?: Array<{ delta?: { content?: string } }>; + choices?: Array<{ + delta?: { content?: string }; + logprobs?: { + content?: Array<{ + token: string; + logprob: number; + top_logprobs?: Array<{ + token: string; + logprob: number; + bytes: number[] | null; + }>; + }>; + }; + }>; } await this.parseSSEStream( reader, targetConversationId, (parsed) => { - const delta = parsed.choices?.[0]?.delta?.content; + const choice = parsed.choices?.[0]; + const delta = choice?.delta?.content; + + // Collect logprobs data + const logprobsContent = choice?.logprobs?.content; + if (logprobsContent) { + for (const item of logprobsContent) { + collectedTokens.push({ + token: item.token, + logprob: item.logprob, + probability: Math.exp(item.logprob), + topLogprobs: (item.top_logprobs || []).map((t) => ({ + token: t.token, + logprob: t.logprob, + bytes: t.bytes, + })), + }); + } + } + if (delta) { streamedContent += delta; const { displayContent, thinkingContent } = @@ -1554,6 +1821,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; }, ); this.syncActiveMessagesIfNeeded(targetConversationId); @@ -1572,6 +1840,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; }, ); this.syncActiveMessagesIfNeeded(targetConversationId); @@ -1914,6 +2183,8 @@ class AppStore { messages: apiMessages, temperature: 0.7, stream: true, + logprobs: true, + top_logprobs: 5, }), }); @@ -1930,14 +2201,48 @@ class AppStore { let streamedContent = ""; interface ChatCompletionChunk { - choices?: Array<{ delta?: { content?: string } }>; + choices?: Array<{ + delta?: { content?: string }; + logprobs?: { + content?: Array<{ + token: string; + logprob: number; + top_logprobs?: Array<{ + token: string; + logprob: number; + bytes: number[] | null; + }>; + }>; + }; + }>; } + const collectedTokens: TokenData[] = []; + await this.parseSSEStream( reader, targetConversationId, (parsed) => { - const tokenContent = parsed.choices?.[0]?.delta?.content; + const choice = parsed.choices?.[0]; + const tokenContent = choice?.delta?.content; + + // Collect logprobs data + const logprobsContent = choice?.logprobs?.content; + if (logprobsContent) { + for (const item of logprobsContent) { + collectedTokens.push({ + token: item.token, + logprob: item.logprob, + probability: Math.exp(item.logprob), + topLogprobs: (item.top_logprobs || []).map((t) => ({ + token: t.token, + logprob: t.logprob, + bytes: t.bytes, + })), + }); + } + } + if (tokenContent) { // Track first token for TTFT if (firstTokenTime === null) { @@ -1973,6 +2278,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; }, ); this.syncActiveMessagesIfNeeded(targetConversationId); @@ -1997,6 +2303,7 @@ class AppStore { (msg) => { msg.content = displayContent; msg.thinking = thinkingContent || undefined; + msg.tokens = [...collectedTokens]; // Store performance metrics on the message if (this.ttftMs !== null) { msg.ttftMs = this.ttftMs; @@ -2693,6 +3000,8 @@ export const editMessage = (messageId: string, newContent: string) => export const editAndRegenerate = (messageId: string, newContent: string) => appStore.editAndRegenerate(messageId, newContent); export const regenerateLastResponse = () => appStore.regenerateLastResponse(); +export const regenerateFromToken = (messageId: string, tokenIndex: number) => + appStore.regenerateFromToken(messageId, tokenIndex); // Conversation actions export const conversations = () => appStore.conversations; diff --git a/dashboard/src/lib/stores/favorites.svelte.ts b/dashboard/src/lib/stores/favorites.svelte.ts new file mode 100644 index 00000000..877b059c --- /dev/null +++ b/dashboard/src/lib/stores/favorites.svelte.ts @@ -0,0 +1,97 @@ +/** + * FavoritesStore - Manages favorite models with localStorage persistence + */ + +import { browser } from "$app/environment"; + +const FAVORITES_KEY = "exo-favorite-models"; + +class FavoritesStore { + favorites = $state>(new Set()); + + constructor() { + if (browser) { + this.loadFromStorage(); + } + } + + private loadFromStorage() { + try { + const stored = localStorage.getItem(FAVORITES_KEY); + if (stored) { + const parsed = JSON.parse(stored) as string[]; + this.favorites = new Set(parsed); + } + } catch (error) { + console.error("Failed to load favorites:", error); + } + } + + private saveToStorage() { + try { + const array = Array.from(this.favorites); + localStorage.setItem(FAVORITES_KEY, JSON.stringify(array)); + } catch (error) { + console.error("Failed to save favorites:", error); + } + } + + add(baseModelId: string) { + const next = new Set(this.favorites); + next.add(baseModelId); + this.favorites = next; + this.saveToStorage(); + } + + remove(baseModelId: string) { + const next = new Set(this.favorites); + next.delete(baseModelId); + this.favorites = next; + this.saveToStorage(); + } + + toggle(baseModelId: string) { + if (this.favorites.has(baseModelId)) { + this.remove(baseModelId); + } else { + this.add(baseModelId); + } + } + + isFavorite(baseModelId: string): boolean { + return this.favorites.has(baseModelId); + } + + getAll(): string[] { + return Array.from(this.favorites); + } + + getSet(): Set { + return new Set(this.favorites); + } + + hasAny(): boolean { + return this.favorites.size > 0; + } + + clearAll() { + this.favorites = new Set(); + this.saveToStorage(); + } +} + +export const favoritesStore = new FavoritesStore(); + +export const favorites = () => favoritesStore.favorites; +export const hasFavorites = () => favoritesStore.hasAny(); +export const isFavorite = (baseModelId: string) => + favoritesStore.isFavorite(baseModelId); +export const toggleFavorite = (baseModelId: string) => + favoritesStore.toggle(baseModelId); +export const addFavorite = (baseModelId: string) => + favoritesStore.add(baseModelId); +export const removeFavorite = (baseModelId: string) => + favoritesStore.remove(baseModelId); +export const getFavorites = () => favoritesStore.getAll(); +export const getFavoritesSet = () => favoritesStore.getSet(); +export const clearFavorites = () => favoritesStore.clearAll(); diff --git a/dashboard/src/routes/+page.svelte b/dashboard/src/routes/+page.svelte index 183b6547..288c3991 100644 --- a/dashboard/src/routes/+page.svelte +++ b/dashboard/src/routes/+page.svelte @@ -5,7 +5,13 @@ ChatMessages, ChatSidebar, ModelCard, + ModelPickerModal, } from "$lib/components"; + import { + favorites, + toggleFavorite, + getFavoritesSet, + } from "$lib/stores/favorites.svelte"; import { hasStartedChat, isTopologyMinimized, @@ -100,6 +106,11 @@ storage_size_megabytes?: number; tasks?: string[]; hugging_face_id?: string; + is_custom?: boolean; + family?: string; + quantization?: string; + base_model?: string; + capabilities?: string[]; }> >([]); @@ -211,9 +222,11 @@ let launchingModelId = $state(null); let instanceDownloadExpandedNodes = $state>(new Set()); - // Custom dropdown state - let isModelDropdownOpen = $state(false); - let modelDropdownSearch = $state(""); + // Model picker modal state + let isModelPickerOpen = $state(false); + + // Favorites state (reactive) + const favoritesSet = $derived(getFavoritesSet()); // Slider dragging state let isDraggingSlider = $state(false); @@ -530,6 +543,47 @@ } } + async function addModelFromPicker(modelId: string) { + const response = await fetch("/models/add", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ model_id: modelId }), + }); + + if (!response.ok) { + let message = `Failed to add model (${response.status}: ${response.statusText})`; + try { + const err = await response.json(); + if (err.detail) message = err.detail; + } catch { + // use default message + } + throw new Error(message); + } + + await fetchModels(); + } + + async function deleteCustomModel(modelId: string) { + try { + const response = await fetch( + `/models/custom/${encodeURIComponent(modelId)}`, + { method: "DELETE" }, + ); + if (response.ok) { + await fetchModels(); + } + } catch { + console.error("Failed to delete custom model"); + } + } + + function handleModelPickerSelect(modelId: string) { + selectPreviewModel(modelId); + saveLaunchDefaults(); + isModelPickerOpen = false; + } + async function launchInstance( modelId: string, specificPreview?: PlacementPreview | null, @@ -2360,14 +2414,12 @@ > - -
+ +
-
- - - -
- - {#if isModelDropdownOpen} - - - -
- -
- -
- - -
- {#each sortedModels().filter((m) => !modelDropdownSearch || (m.name || m.id) - .toLowerCase() - .includes(modelDropdownSearch.toLowerCase())) as model} - {@const sizeGB = getModelSizeGB(model)} - {@const modelCanFit = hasEnoughMemory(model)} - {@const isImageModel = modelSupportsImageGeneration( - model.id, - )} - {@const isImageEditModel = modelSupportsImageEditing( - model.id, - )} - - {:else} -
- No models found -
- {/each} -
+
- {/if} +
@@ -3354,3 +3246,22 @@ {/if}
+ + m.id))} + canModelFit={(modelId) => { + const model = models.find((m) => m.id === modelId); + return model ? hasEnoughMemory(model) : false; + }} + onSelect={handleModelPickerSelect} + onClose={() => (isModelPickerOpen = false)} + onToggleFavorite={toggleFavorite} + onAddModel={addModelFromPicker} + onDeleteModel={deleteCustomModel} + totalMemoryGB={clusterMemory().total / (1024 * 1024 * 1024)} + usedMemoryGB={clusterMemory().used / (1024 * 1024 * 1024)} +/> diff --git a/python/parts.nix b/python/parts.nix index 7423ed31..9d5580ae 100644 --- a/python/parts.nix +++ b/python/parts.nix @@ -69,7 +69,8 @@ # Create wrapper scripts for script in exo exo-master exo-worker; do makeWrapper ${exoVenv}/bin/$script $out/bin/$script \ - --set DASHBOARD_DIR ${self'.packages.dashboard} \ + --set EXO_DASHBOARD_DIR ${self'.packages.dashboard} \ + --set EXO_RESOURCES_DIR ${inputs.self + "/resources"} \ ${lib.optionalString pkgs.stdenv.isDarwin "--prefix PATH : ${pkgs.macmon}/bin"} done ''; diff --git a/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml index 89cd0f6f..8d3a637e 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-4bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["TextToImage"] +uses_cfg = true [storage_size] in_bytes = 26799533856 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml index 43951dab..ddf78c4a 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-8bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["TextToImage"] +uses_cfg = true [storage_size] in_bytes = 37014734400 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml index 99a60af2..db2f5e54 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-4bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["ImageToImage"] +uses_cfg = true [storage_size] in_bytes = 26799533856 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml index 0f326b39..2db63265 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509-8bit.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["ImageToImage"] +uses_cfg = true [storage_size] in_bytes = 37014734400 diff --git a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml index 65044e6c..3b615da1 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image-Edit-2509.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["ImageToImage"] +uses_cfg = true [storage_size] in_bytes = 57445135488 diff --git a/resources/image_model_cards/exolabs--Qwen-Image.toml b/resources/image_model_cards/exolabs--Qwen-Image.toml index a39235ea..d012af50 100644 --- a/resources/image_model_cards/exolabs--Qwen-Image.toml +++ b/resources/image_model_cards/exolabs--Qwen-Image.toml @@ -3,6 +3,7 @@ n_layers = 60 hidden_size = 1 supports_tensor = false tasks = ["TextToImage"] +uses_cfg = true [storage_size] in_bytes = 57445135488 diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml index 26de8de8..41784cf6 100644 --- a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml +++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "deepseek" +quantization = "4bit" +base_model = "DeepSeek V3.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 405874409472 diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml index 13cf367b..a5d77bcd 100644 --- a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml +++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "deepseek" +quantization = "8bit" +base_model = "DeepSeek V3.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 765577920512 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml index 288392f6..a7acea44 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml @@ -3,6 +3,10 @@ n_layers = 46 hidden_size = 4096 supports_tensor = false tasks = ["TextGeneration"] +family = "glm" +quantization = "8bit" +base_model = "GLM 4.5 Air" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 122406567936 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml index 00a19df2..4258c225 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml @@ -3,6 +3,10 @@ n_layers = 46 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "bf16" +base_model = "GLM 4.5 Air" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 229780750336 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml index 816c9657..0672d664 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml @@ -3,6 +3,10 @@ n_layers = 91 hidden_size = 5120 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "4bit" +base_model = "GLM 4.7" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 198556925568 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml index b087164b..bcf1cae4 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml @@ -3,6 +3,10 @@ n_layers = 91 hidden_size = 5120 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "6bit" +base_model = "GLM 4.7" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 286737579648 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml index 6f221cef..0f56c2f7 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml @@ -3,6 +3,10 @@ n_layers = 91 hidden_size = 5120 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "8bit" +base_model = "GLM 4.7" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 396963397248 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml index 43eb0dcd..8637cef0 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "4bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 19327352832 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml index 6a512c0a..b9a9da4d 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "5bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 22548578304 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml index 86c65489..e3cb1fa8 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "6bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 26843545600 diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml index eb69183f..bd6df312 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml @@ -3,6 +3,10 @@ n_layers = 47 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "glm" +quantization = "8bit" +base_model = "GLM 4.7 Flash" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 34359738368 diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml index d7acabec..3f21d4c0 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "kimi" +quantization = "4bit" +base_model = "Kimi K2" +capabilities = ["text"] [storage_size] in_bytes = 620622774272 diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml b/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml index 8d21727b..0a955b04 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "kimi" +quantization = "" +base_model = "Kimi K2" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 706522120192 diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml b/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml index c44cf9b1..806c6b30 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 7168 supports_tensor = true tasks = ["TextGeneration"] +family = "kimi" +quantization = "" +base_model = "Kimi K2.5" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 662498705408 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml index db334221..b38ec20f 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 16 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.2 1B" +capabilities = ["text"] [storage_size] in_bytes = 729808896 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml index 001b4a1d..81ce4567 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.2 3B" +capabilities = ["text"] [storage_size] in_bytes = 1863319552 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml index 358db81b..ac9a203b 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "8bit" +base_model = "Llama 3.2 3B" +capabilities = ["text"] [storage_size] in_bytes = 3501195264 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml index cf6eece0..24c7cbaa 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.3 70B" +capabilities = ["text"] [storage_size] in_bytes = 40652242944 diff --git a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml index 15f0c551..3bfc97dc 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "8bit" +base_model = "Llama 3.3 70B" +capabilities = ["text"] [storage_size] in_bytes = 76799803392 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml index b766164d..27d0b724 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.1 70B" +capabilities = ["text"] [storage_size] in_bytes = 40652242944 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml index b6d10c40..1fe34ba8 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 32 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "4bit" +base_model = "Llama 3.1 8B" +capabilities = ["text"] [storage_size] in_bytes = 4637851648 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml index 4cfe47cf..5310a2a0 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 32 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "8bit" +base_model = "Llama 3.1 8B" +capabilities = ["text"] [storage_size] in_bytes = 8954839040 diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml index 9e04f27b..eb6405e0 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml @@ -3,6 +3,10 @@ n_layers = 32 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "bf16" +base_model = "Llama 3.1 8B" +capabilities = ["text"] [storage_size] in_bytes = 16882073600 diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml index 4bf81136..92ec6746 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "minimax" +quantization = "3bit" +base_model = "MiniMax M2.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 100086644736 diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml index 54a49f97..c1388d2f 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml @@ -3,6 +3,10 @@ n_layers = 61 hidden_size = 3072 supports_tensor = true tasks = ["TextGeneration"] +family = "minimax" +quantization = "8bit" +base_model = "MiniMax M2.1" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 242986745856 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml index 212cdef6..7929aaba 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 1024 supports_tensor = false tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 0.6B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 342884352 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml index ac591d6c..d9fcc368 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml @@ -3,6 +3,10 @@ n_layers = 28 hidden_size = 1024 supports_tensor = false tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 0.6B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 698351616 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml index 020c11be..ef835c6a 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml @@ -3,6 +3,10 @@ n_layers = 94 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 235B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 141733920768 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml index 64afc366..f6e079ab 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml @@ -3,6 +3,10 @@ n_layers = 94 hidden_size = 4096 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 235B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 268435456000 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml index 1b9f92a6..48a6666f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 30B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 17612931072 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml index f8e59ac1..c283396f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 30B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 33279705088 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml index 25a0cf2b..b390bd21 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 62 hidden_size = 6144 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 Coder 480B" +capabilities = ["text", "code"] [storage_size] in_bytes = 289910292480 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml index cdb1d6ec..1c21307c 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 62 hidden_size = 6144 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 Coder 480B" +capabilities = ["text", "code"] [storage_size] in_bytes = 579820584960 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml index db55b7f9..386a3fa1 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text"] [storage_size] in_bytes = 46976204800 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml index e36e24b8..0e2bf2a5 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text"] [storage_size] in_bytes = 88814387200 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml index bc3bdf50..2a3e3c19 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "4bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 47080074240 diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml index dd5512a7..65d33253 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml @@ -3,6 +3,10 @@ n_layers = 48 hidden_size = 2048 supports_tensor = true tasks = ["TextGeneration"] +family = "qwen" +quantization = "8bit" +base_model = "Qwen3 Next 80B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 88814387200 diff --git a/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml b/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml index c725e728..f579c618 100644 --- a/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml +++ b/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml @@ -3,6 +3,10 @@ n_layers = 36 hidden_size = 2880 supports_tensor = true tasks = ["TextGeneration"] +family = "gpt-oss" +quantization = "MXFP4-Q8" +base_model = "GPT-OSS 120B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 70652212224 diff --git a/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml b/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml index bf8f1a60..af1e04ad 100644 --- a/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml +++ b/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml @@ -3,6 +3,10 @@ n_layers = 24 hidden_size = 2880 supports_tensor = true tasks = ["TextGeneration"] +family = "gpt-oss" +quantization = "MXFP4-Q8" +base_model = "GPT-OSS 20B" +capabilities = ["text", "thinking"] [storage_size] in_bytes = 12025908224 diff --git a/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml b/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml index dd451015..e61660c2 100644 --- a/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml +++ b/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml @@ -3,6 +3,10 @@ n_layers = 80 hidden_size = 8192 supports_tensor = true tasks = ["TextGeneration"] +family = "llama" +quantization = "fp16" +base_model = "Llama 3.3 70B" +capabilities = ["text"] [storage_size] in_bytes = 144383672320 diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py index c2f7b9e9..f5798ad3 100644 --- a/src/exo/download/coordinator.py +++ b/src/exo/download/coordinator.py @@ -1,4 +1,5 @@ import asyncio +import socket from dataclasses import dataclass, field from typing import Iterator @@ -60,10 +61,37 @@ class DownloadCoordinator: async def run(self) -> None: logger.info("Starting DownloadCoordinator") + self._test_internet_connection() async with self._tg as tg: tg.start_soon(self._command_processor) tg.start_soon(self._forward_events) tg.start_soon(self._emit_existing_download_progress) + tg.start_soon(self._check_internet_connection) + + def _test_internet_connection(self) -> None: + try: + socket.create_connection(("1.1.1.1", 443), timeout=3).close() + self.shard_downloader.set_internet_connection(True) + except OSError: + self.shard_downloader.set_internet_connection(False) + logger.debug( + f"Internet connectivity: {self.shard_downloader.internet_connection}" + ) + + async def _check_internet_connection(self) -> None: + first_connection = True + while True: + await asyncio.sleep(10) + + # Assume that internet connection is set to False on 443 errors. + if self.shard_downloader.internet_connection: + continue + + self._test_internet_connection() + + if first_connection and self.shard_downloader.internet_connection: + first_connection = False + self._tg.start_soon(self._emit_existing_download_progress) def shutdown(self) -> None: self._tg.cancel_scope.cancel() @@ -241,7 +269,7 @@ class DownloadCoordinator: async def _emit_existing_download_progress(self) -> None: try: while True: - logger.info( + logger.debug( "DownloadCoordinator: Fetching and emitting existing download progress..." ) async for ( @@ -274,10 +302,10 @@ class DownloadCoordinator: await self.event_sender.send( NodeDownloadProgress(download_progress=status) ) - logger.info( + logger.debug( "DownloadCoordinator: Done emitting existing download progress." ) - await anyio.sleep(5 * 60) # 5 minutes + await anyio.sleep(60) except Exception as e: logger.error( f"DownloadCoordinator: Error emitting existing download progress: {e}" diff --git a/src/exo/download/download_utils.py b/src/exo/download/download_utils.py index 6dec4718..618e4f38 100644 --- a/src/exo/download/download_utils.py +++ b/src/exo/download/download_utils.py @@ -49,6 +49,10 @@ class HuggingFaceAuthenticationError(Exception): """Raised when HuggingFace returns 401/403 for a model download.""" +class HuggingFaceRateLimitError(Exception): + """429 Huggingface code""" + + async def _build_auth_error_message(status_code: int, model_id: ModelId) -> str: token = await get_hf_token() if status_code == 401 and token is None: @@ -154,49 +158,76 @@ async def seed_models(seed_dir: str | Path): logger.error(traceback.format_exc()) +_fetched_file_lists_this_session: set[str] = set() + + async def fetch_file_list_with_cache( - model_id: ModelId, revision: str = "main", recursive: bool = False + model_id: ModelId, + revision: str = "main", + recursive: bool = False, + skip_internet: bool = False, + on_connection_lost: Callable[[], None] = lambda: None, ) -> list[FileListEntry]: target_dir = (await ensure_models_dir()) / "caches" / model_id.normalize() await aios.makedirs(target_dir, exist_ok=True) cache_file = target_dir / f"{model_id.normalize()}--{revision}--file_list.json" + cache_key = f"{model_id.normalize()}--{revision}" + + if cache_key in _fetched_file_lists_this_session and await aios.path.exists( + cache_file + ): + async with aiofiles.open(cache_file, "r") as f: + return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) + + if skip_internet: + if await aios.path.exists(cache_file): + async with aiofiles.open(cache_file, "r") as f: + return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) + raise FileNotFoundError( + f"No internet connection and no cached file list for {model_id}" + ) - # Always try fresh first try: file_list = await fetch_file_list_with_retry( - model_id, revision, recursive=recursive + model_id, + revision, + recursive=recursive, + on_connection_lost=on_connection_lost, ) - # Update cache with fresh data async with aiofiles.open(cache_file, "w") as f: await f.write( TypeAdapter(list[FileListEntry]).dump_json(file_list).decode() ) + _fetched_file_lists_this_session.add(cache_key) return file_list except Exception as e: - # Fetch failed - try cache fallback if await aios.path.exists(cache_file): logger.warning( f"Failed to fetch file list for {model_id}, using cached data: {e}" ) async with aiofiles.open(cache_file, "r") as f: return TypeAdapter(list[FileListEntry]).validate_json(await f.read()) - # No cache available, propagate the error - raise + raise FileNotFoundError(f"Failed to fetch file list for {model_id}: {e}") from e async def fetch_file_list_with_retry( - model_id: ModelId, revision: str = "main", path: str = "", recursive: bool = False + model_id: ModelId, + revision: str = "main", + path: str = "", + recursive: bool = False, + on_connection_lost: Callable[[], None] = lambda: None, ) -> list[FileListEntry]: - n_attempts = 30 + n_attempts = 3 for attempt in range(n_attempts): try: return await _fetch_file_list(model_id, revision, path, recursive) except HuggingFaceAuthenticationError: raise except Exception as e: + on_connection_lost() if attempt == n_attempts - 1: raise e - await asyncio.sleep(min(8, 0.1 * float(2.0 ** int(attempt)))) + await asyncio.sleep(2.0**attempt) raise Exception( f"Failed to fetch file list for {model_id=} {revision=} {path=} {recursive=}" ) @@ -216,7 +247,11 @@ async def _fetch_file_list( if response.status in [401, 403]: msg = await _build_auth_error_message(response.status, model_id) raise HuggingFaceAuthenticationError(msg) - if response.status == 200: + elif response.status == 429: + raise HuggingFaceRateLimitError( + f"Couldn't download {model_id} because of HuggingFace rate limit." + ) + elif response.status == 200: data_json = await response.text() data = TypeAdapter(list[FileListEntry]).validate_json(data_json) files: list[FileListEntry] = [] @@ -249,7 +284,7 @@ def create_http_session( else: total_timeout = 1800 connect_timeout = 60 - sock_read_timeout = 1800 + sock_read_timeout = 60 sock_connect_timeout = 60 ssl_context = ssl.create_default_context( @@ -324,8 +359,9 @@ async def download_file_with_retry( path: str, target_dir: Path, on_progress: Callable[[int, int, bool], None] = lambda _, __, ___: None, + on_connection_lost: Callable[[], None] = lambda: None, ) -> Path: - n_attempts = 30 + n_attempts = 3 for attempt in range(n_attempts): try: return await _download_file( @@ -333,14 +369,19 @@ async def download_file_with_retry( ) except HuggingFaceAuthenticationError: raise - except Exception as e: - if isinstance(e, FileNotFoundError) or attempt == n_attempts - 1: + except HuggingFaceRateLimitError as e: + if attempt == n_attempts - 1: raise e logger.error( f"Download error on attempt {attempt}/{n_attempts} for {model_id=} {revision=} {path=} {target_dir=}" ) logger.error(traceback.format_exc()) - await asyncio.sleep(min(8, 0.1 * (2.0**attempt))) + await asyncio.sleep(2.0**attempt) + except Exception as e: + on_connection_lost() + if attempt == n_attempts - 1: + raise e + break raise Exception( f"Failed to download file {model_id=} {revision=} {path=} {target_dir=}" ) @@ -542,7 +583,9 @@ async def download_shard( on_progress: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], max_parallel_downloads: int = 8, skip_download: bool = False, + skip_internet: bool = False, allow_patterns: list[str] | None = None, + on_connection_lost: Callable[[], None] = lambda: None, ) -> tuple[Path, RepoDownloadProgress]: if not skip_download: logger.debug(f"Downloading {shard.model_card.model_id=}") @@ -562,7 +605,11 @@ async def download_shard( all_start_time = time.time() file_list = await fetch_file_list_with_cache( - shard.model_card.model_id, revision, recursive=True + shard.model_card.model_id, + revision, + recursive=True, + skip_internet=skip_internet, + on_connection_lost=on_connection_lost, ) filtered_file_list = list( filter_repo_objects( @@ -672,6 +719,7 @@ async def download_shard( lambda curr_bytes, total_bytes, is_renamed: schedule_progress( file, curr_bytes, total_bytes, is_renamed ), + on_connection_lost=on_connection_lost, ) if not skip_download: diff --git a/src/exo/download/impl_shard_downloader.py b/src/exo/download/impl_shard_downloader.py index 1b7f5eab..0e7aea1e 100644 --- a/src/exo/download/impl_shard_downloader.py +++ b/src/exo/download/impl_shard_downloader.py @@ -1,4 +1,5 @@ import asyncio +from asyncio import create_task from collections.abc import Awaitable from pathlib import Path from typing import AsyncIterator, Callable @@ -49,6 +50,10 @@ class SingletonShardDownloader(ShardDownloader): self.shard_downloader = shard_downloader self.active_downloads: dict[ShardMetadata, asyncio.Task[Path]] = {} + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + self.shard_downloader.set_internet_connection(value) + def on_progress( self, callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], @@ -85,6 +90,10 @@ class CachedShardDownloader(ShardDownloader): self.shard_downloader = shard_downloader self.cache: dict[tuple[str, ShardMetadata], Path] = {} + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + self.shard_downloader.set_internet_connection(value) + def on_progress( self, callback: Callable[[ShardMetadata, RepoDownloadProgress], Awaitable[None]], @@ -142,6 +151,8 @@ class ResumableShardDownloader(ShardDownloader): self.on_progress_wrapper, max_parallel_downloads=self.max_parallel_downloads, allow_patterns=allow_patterns, + skip_internet=not self.internet_connection, + on_connection_lost=lambda: self.set_internet_connection(False), ) return target_dir @@ -154,12 +165,23 @@ class ResumableShardDownloader(ShardDownloader): """Helper coroutine that builds the shard for a model and gets its download status.""" shard = await build_full_shard(model_id) return await download_shard( - shard, self.on_progress_wrapper, skip_download=True + shard, + self.on_progress_wrapper, + skip_download=True, + skip_internet=not self.internet_connection, + on_connection_lost=lambda: self.set_internet_connection(False), ) - # Kick off download status coroutines concurrently + semaphore = asyncio.Semaphore(self.max_parallel_downloads) + + async def download_with_semaphore( + model_card: ModelCard, + ) -> tuple[Path, RepoDownloadProgress]: + async with semaphore: + return await _status_for_model(model_card.model_id) + tasks = [ - asyncio.create_task(_status_for_model(model_card.model_id)) + create_task(download_with_semaphore(model_card)) for model_card in await get_model_cards() ] diff --git a/src/exo/download/shard_downloader.py b/src/exo/download/shard_downloader.py index 30c11d25..9dd8c324 100644 --- a/src/exo/download/shard_downloader.py +++ b/src/exo/download/shard_downloader.py @@ -16,6 +16,11 @@ from exo.shared.types.worker.shards import ( # TODO: the PipelineShardMetadata getting reinstantiated is a bit messy. Should this be a classmethod? class ShardDownloader(ABC): + internet_connection: bool = False + + def set_internet_connection(self, value: bool) -> None: + self.internet_connection = value + @abstractmethod async def ensure_shard( self, shard: ShardMetadata, config_only: bool = False diff --git a/src/exo/master/adapters/chat_completions.py b/src/exo/master/adapters/chat_completions.py index 5a27664d..3e013079 100644 --- a/src/exo/master/adapters/chat_completions.py +++ b/src/exo/master/adapters/chat_completions.py @@ -14,6 +14,8 @@ from exo.shared.types.api import ( ErrorInfo, ErrorResponse, FinishReason, + Logprobs, + LogprobsContentItem, StreamingChoiceResponse, ToolCall, ) @@ -66,7 +68,9 @@ def chat_request_to_text_generation( return TextGenerationTaskParams( model=request.model, - input=input_messages if input_messages else "", + input=input_messages + if input_messages + else [InputMessage(role="user", content="")], instructions=instructions, max_output_tokens=request.max_tokens, temperature=request.temperature, @@ -79,6 +83,8 @@ def chat_request_to_text_generation( chat_template_messages=chat_template_messages if chat_template_messages else None, + logprobs=request.logprobs or False, + top_logprobs=request.top_logprobs, ) @@ -86,6 +92,19 @@ def chunk_to_response( chunk: TokenChunk, command_id: CommandId ) -> ChatCompletionResponse: """Convert a TokenChunk to a streaming ChatCompletionResponse.""" + # Build logprobs if available + logprobs: Logprobs | None = None + if chunk.logprob is not None: + logprobs = Logprobs( + content=[ + LogprobsContentItem( + token=chunk.text, + logprob=chunk.logprob, + top_logprobs=chunk.top_logprobs or [], + ) + ] + ) + return ChatCompletionResponse( id=command_id, created=int(time.time()), @@ -94,6 +113,7 @@ def chunk_to_response( StreamingChoiceResponse( index=0, delta=ChatCompletionMessage(role="assistant", content=chunk.text), + logprobs=logprobs, finish_reason=chunk.finish_reason, ) ], @@ -160,6 +180,7 @@ async def collect_chat_response( """Collect all token chunks and return a single ChatCompletionResponse.""" text_parts: list[str] = [] tool_calls: list[ToolCall] = [] + logprobs_content: list[LogprobsContentItem] = [] model: str | None = None finish_reason: FinishReason | None = None error_message: str | None = None @@ -174,6 +195,14 @@ async def collect_chat_response( if isinstance(chunk, TokenChunk): text_parts.append(chunk.text) + if chunk.logprob is not None: + logprobs_content.append( + LogprobsContentItem( + token=chunk.text, + logprob=chunk.logprob, + top_logprobs=chunk.top_logprobs or [], + ) + ) if isinstance(chunk, ToolCallChunk): tool_calls.extend( @@ -206,6 +235,9 @@ async def collect_chat_response( content=combined_text, tool_calls=tool_calls if tool_calls else None, ), + logprobs=Logprobs(content=logprobs_content) + if logprobs_content + else None, finish_reason=finish_reason, ) ], diff --git a/src/exo/master/adapters/claude.py b/src/exo/master/adapters/claude.py index 13398012..6c17b49c 100644 --- a/src/exo/master/adapters/claude.py +++ b/src/exo/master/adapters/claude.py @@ -141,7 +141,9 @@ def claude_request_to_text_generation( return TextGenerationTaskParams( model=request.model, - input=input_messages if input_messages else "", + input=input_messages + if input_messages + else [InputMessage(role="user", content="")], instructions=instructions, max_output_tokens=request.max_tokens, temperature=request.temperature, diff --git a/src/exo/master/adapters/responses.py b/src/exo/master/adapters/responses.py index 27d845cd..c2a416ac 100644 --- a/src/exo/master/adapters/responses.py +++ b/src/exo/master/adapters/responses.py @@ -43,10 +43,10 @@ def _extract_content(content: str | list[ResponseContentPart]) -> str: def responses_request_to_text_generation( request: ResponsesRequest, ) -> TextGenerationTaskParams: - input_value: str | list[InputMessage] + input_value: list[InputMessage] built_chat_template: list[dict[str, Any]] | None = None if isinstance(request.input, str): - input_value = request.input + input_value = [InputMessage(role="user", content=request.input)] else: input_messages: list[InputMessage] = [] chat_template_messages: list[dict[str, Any]] = [] @@ -95,7 +95,11 @@ def responses_request_to_text_generation( } ) - input_value = input_messages if input_messages else "" + input_value = ( + input_messages + if input_messages + else [InputMessage(role="user", content="")] + ) built_chat_template = chat_template_messages if chat_template_messages else None return TextGenerationTaskParams( diff --git a/src/exo/master/api.py b/src/exo/master/api.py index 8cf86725..0ad5454c 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -1,6 +1,7 @@ import base64 import contextlib import json +import random import time from collections.abc import AsyncGenerator, Awaitable, Callable from datetime import datetime, timezone @@ -50,10 +51,13 @@ from exo.shared.logging import InterceptLogger from exo.shared.models.model_cards import ( ModelCard, ModelId, + delete_custom_card, get_model_cards, + is_custom_card, ) from exo.shared.tracing import TraceEvent, compute_stats, export_trace, load_trace_file from exo.shared.types.api import ( + AddCustomModelParams, AdvancedImageParams, BenchChatCompletionRequest, BenchChatCompletionResponse, @@ -71,6 +75,7 @@ from exo.shared.types.api import ( ErrorResponse, FinishReason, GenerationStats, + HuggingFaceSearchResult, ImageData, ImageEditsTaskParams, ImageGenerationResponse, @@ -146,6 +151,15 @@ def _format_to_content_type(image_format: Literal["png", "jpeg", "webp"] | None) return f"image/{image_format or 'png'}" +def _ensure_seed(params: AdvancedImageParams | None) -> AdvancedImageParams: + """Ensure advanced params has a seed set for distributed consistency.""" + if params is None: + return AdvancedImageParams(seed=random.randint(0, 2**32 - 1)) + if params.seed is None: + return params.model_copy(update={"seed": random.randint(0, 2**32 - 1)}) + return params + + class API: def __init__( self, @@ -257,6 +271,9 @@ class API: self.app.delete("/instance/{instance_id}")(self.delete_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) + self.app.delete("/models/custom/{model_id:path}")(self.delete_custom_model) + self.app.get("/models/search")(self.search_models) self.app.post("/v1/chat/completions", response_model=None)( self.chat_completions ) @@ -610,6 +627,11 @@ class API: self._token_chunk_stream(command.command_id), ), media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "close", + "X-Accel-Buffering": "no", + }, ) return await collect_chat_response( @@ -702,6 +724,9 @@ class API: with SSE-formatted events for partial and final images. """ payload.model = await self._validate_image_model(ModelId(payload.model)) + payload = payload.model_copy( + update={"advanced_params": _ensure_seed(payload.advanced_params)} + ) command = ImageGeneration( task_params=payload, @@ -950,6 +975,9 @@ class API: payload.stream = False payload.partial_images = 0 + payload = payload.model_copy( + update={"advanced_params": _ensure_seed(payload.advanced_params)} + ) command = ImageGeneration( task_params=payload, @@ -981,6 +1009,7 @@ class API: ) -> ImageEdits: """Prepare and send an image edits command with chunked image upload.""" resolved_model = await self._validate_image_model(model) + advanced_params = _ensure_seed(advanced_params) image_content = await image.read() image_data = base64.b64encode(image_content).decode("utf-8") @@ -1159,6 +1188,11 @@ class API: self._token_chunk_stream(command.command_id), ), media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "close", + "X-Accel-Buffering": "no", + }, ) return await collect_claude_response( @@ -1186,6 +1220,11 @@ class API: self._token_chunk_stream(command.command_id), ), media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "close", + "X-Accel-Buffering": "no", + }, ) return await collect_responses_response( @@ -1216,11 +1255,70 @@ class API: storage_size_megabytes=int(card.storage_size.in_mb), supports_tensor=card.supports_tensor, tasks=[task.value for task in card.tasks], + is_custom=is_custom_card(card.model_id), + family=card.family, + quantization=card.quantization, + base_model=card.base_model, + capabilities=card.capabilities, ) for card in await get_model_cards() ] ) + async def add_custom_model(self, payload: AddCustomModelParams) -> ModelListModel: + """Fetch a model from HuggingFace and save as a custom model card.""" + try: + card = await ModelCard.fetch_from_hf(payload.model_id) + except Exception as exc: + raise HTTPException( + status_code=400, detail=f"Failed to fetch model: {exc}" + ) from exc + + return ModelListModel( + id=card.model_id, + hugging_face_id=card.model_id, + name=card.model_id.short(), + description="", + tags=[], + storage_size_megabytes=int(card.storage_size.in_mb), + supports_tensor=card.supports_tensor, + tasks=[task.value for task in card.tasks], + is_custom=True, + ) + + async def delete_custom_model(self, model_id: ModelId) -> JSONResponse: + """Delete a user-added custom model card.""" + deleted = await delete_custom_card(model_id) + if not deleted: + raise HTTPException(status_code=404, detail="Custom model card not found") + return JSONResponse( + {"message": "Model card deleted", "model_id": str(model_id)} + ) + + async def search_models( + self, query: str = "", limit: int = 20 + ) -> list[HuggingFaceSearchResult]: + """Search HuggingFace Hub for mlx-community models.""" + from huggingface_hub import list_models + + results = list_models( + search=query or None, + author="mlx-community", + sort="downloads", + limit=limit, + ) + return [ + HuggingFaceSearchResult( + id=m.id, + author=m.author or "", + downloads=m.downloads or 0, + likes=m.likes or 0, + last_modified=str(m.last_modified or ""), + tags=list(m.tags or []), + ) + for m in results + ] + async def run(self): cfg = Config() cfg.bind = f"0.0.0.0:{self.port}" diff --git a/src/exo/master/placement_utils.py b/src/exo/master/placement_utils.py index 309abc25..b20a39cc 100644 --- a/src/exo/master/placement_utils.py +++ b/src/exo/master/placement_utils.py @@ -10,6 +10,7 @@ from exo.shared.types.profiling import MemoryUsage, NodeNetworkInfo from exo.shared.types.topology import Cycle, RDMAConnection, SocketConnection from exo.shared.types.worker.runners import RunnerId, ShardAssignments from exo.shared.types.worker.shards import ( + CfgShardMetadata, PipelineShardMetadata, Sharding, ShardMetadata, @@ -74,40 +75,43 @@ def allocate_layers_proportionally( return result -def get_shard_assignments_for_pipeline_parallel( - model_card: ModelCard, - cycle: Cycle, - node_memory: Mapping[NodeId, MemoryUsage], -): +def _validate_cycle(cycle: Cycle) -> None: if not cycle.node_ids: raise ValueError("Cannot create shard assignments for empty node cycle") - cycle_memory = sum( - (node_memory[node_id].ram_available for node_id in cycle.node_ids), + +def _compute_total_memory( + node_ids: list[NodeId], + node_memory: Mapping[NodeId, MemoryUsage], +) -> Memory: + total_memory = sum( + (node_memory[node_id].ram_available for node_id in node_ids), start=Memory(), ) - if cycle_memory.in_bytes == 0: + if total_memory.in_bytes == 0: raise ValueError("Cannot create shard assignments: total available memory is 0") + return total_memory - total_layers = model_card.n_layers - world_size = len(cycle) - runner_to_shard: dict[RunnerId, ShardMetadata] = {} - node_to_runner: dict[NodeId, RunnerId] = {} +def _allocate_and_validate_layers( + node_ids: list[NodeId], + node_memory: Mapping[NodeId, MemoryUsage], + total_memory: Memory, + model_card: ModelCard, +) -> list[int]: layer_allocations = allocate_layers_proportionally( - total_layers=total_layers, + total_layers=model_card.n_layers, memory_fractions=[ - node_memory[node_id].ram_available.in_bytes / cycle_memory.in_bytes - for node_id in cycle.node_ids + node_memory[node_id].ram_available.in_bytes / total_memory.in_bytes + for node_id in node_ids ], ) - # Validate each node has sufficient memory for its assigned layers - memory_per_layer = model_card.storage_size.in_bytes / total_layers - for i, (node_id, node_layers) in enumerate( - zip(cycle.node_ids, layer_allocations, strict=True) - ): - required_memory = node_layers * memory_per_layer + total_storage_bytes = model_card.storage_size.in_bytes + total_layers = model_card.n_layers + for i, node_id in enumerate(node_ids): + node_layers = layer_allocations[i] + required_memory = (total_storage_bytes * node_layers) // total_layers available_memory = node_memory[node_id].ram_available.in_bytes if required_memory > available_memory: raise ValueError( @@ -116,32 +120,125 @@ def get_shard_assignments_for_pipeline_parallel( f"but only has {available_memory / (1024**3):.2f} GB available" ) - layers_assigned = 0 - for i, (node_id, node_layers) in enumerate( - zip(cycle.node_ids, layer_allocations, strict=True) - ): - runner_id = RunnerId() + return layer_allocations - shard = PipelineShardMetadata( + +def get_shard_assignments_for_pipeline_parallel( + model_card: ModelCard, + cycle: Cycle, + node_memory: Mapping[NodeId, MemoryUsage], +) -> ShardAssignments: + """Create shard assignments for pipeline parallel execution.""" + world_size = len(cycle) + use_cfg_parallel = model_card.uses_cfg and world_size >= 2 and world_size % 2 == 0 + + if use_cfg_parallel: + return _get_shard_assignments_for_cfg_parallel(model_card, cycle, node_memory) + else: + return _get_shard_assignments_for_pure_pipeline(model_card, cycle, node_memory) + + +def _get_shard_assignments_for_cfg_parallel( + model_card: ModelCard, + cycle: Cycle, + node_memory: Mapping[NodeId, MemoryUsage], +) -> ShardAssignments: + """Create shard assignments for CFG parallel execution. + + CFG parallel runs two independent pipelines. Group 0 processes the positive + prompt, group 1 processes the negative prompt. The ring topology places + group 1's ranks in reverse order so both "last stages" are neighbors for + efficient CFG exchange. + """ + _validate_cycle(cycle) + + world_size = len(cycle) + cfg_world_size = 2 + pipeline_world_size = world_size // cfg_world_size + + # Allocate layers for one pipeline group (both groups run the same layers) + pipeline_node_ids = cycle.node_ids[:pipeline_world_size] + pipeline_memory = _compute_total_memory(pipeline_node_ids, node_memory) + layer_allocations = _allocate_and_validate_layers( + pipeline_node_ids, node_memory, pipeline_memory, model_card + ) + + # Ring topology: group 0 ascending [0,1,2,...], group 1 descending [...,2,1,0] + # This places both last stages as neighbors for CFG exchange. + position_to_cfg_pipeline = [(0, r) for r in range(pipeline_world_size)] + [ + (1, r) for r in reversed(range(pipeline_world_size)) + ] + + runner_to_shard: dict[RunnerId, ShardMetadata] = {} + node_to_runner: dict[NodeId, RunnerId] = {} + + for device_rank, node_id in enumerate(cycle.node_ids): + cfg_rank, pipeline_rank = position_to_cfg_pipeline[device_rank] + layers_before = sum(layer_allocations[:pipeline_rank]) + node_layers = layer_allocations[pipeline_rank] + + shard = CfgShardMetadata( model_card=model_card, - device_rank=i, + device_rank=device_rank, world_size=world_size, - start_layer=layers_assigned, - end_layer=layers_assigned + node_layers, - n_layers=total_layers, + start_layer=layers_before, + end_layer=layers_before + node_layers, + n_layers=model_card.n_layers, + cfg_rank=cfg_rank, + cfg_world_size=cfg_world_size, + pipeline_rank=pipeline_rank, + pipeline_world_size=pipeline_world_size, ) + runner_id = RunnerId() runner_to_shard[runner_id] = shard node_to_runner[node_id] = runner_id - layers_assigned += node_layers - shard_assignments = ShardAssignments( + return ShardAssignments( model_id=model_card.model_id, runner_to_shard=runner_to_shard, node_to_runner=node_to_runner, ) - return shard_assignments + +def _get_shard_assignments_for_pure_pipeline( + model_card: ModelCard, + cycle: Cycle, + node_memory: Mapping[NodeId, MemoryUsage], +) -> ShardAssignments: + """Create shard assignments for pure pipeline execution.""" + _validate_cycle(cycle) + total_memory = _compute_total_memory(cycle.node_ids, node_memory) + + layer_allocations = _allocate_and_validate_layers( + cycle.node_ids, node_memory, total_memory, model_card + ) + + runner_to_shard: dict[RunnerId, ShardMetadata] = {} + node_to_runner: dict[NodeId, RunnerId] = {} + + for pipeline_rank, node_id in enumerate(cycle.node_ids): + layers_before = sum(layer_allocations[:pipeline_rank]) + node_layers = layer_allocations[pipeline_rank] + + shard = PipelineShardMetadata( + model_card=model_card, + device_rank=pipeline_rank, + world_size=len(cycle), + start_layer=layers_before, + end_layer=layers_before + node_layers, + n_layers=model_card.n_layers, + ) + + runner_id = RunnerId() + runner_to_shard[runner_id] = shard + node_to_runner[node_id] = runner_id + + return ShardAssignments( + model_id=model_card.model_id, + runner_to_shard=runner_to_shard, + node_to_runner=node_to_runner, + ) def get_shard_assignments_for_tensor_parallel( diff --git a/src/exo/master/tests/test_master.py b/src/exo/master/tests/test_master.py index d9987727..ddf9aec8 100644 --- a/src/exo/master/tests/test_master.py +++ b/src/exo/master/tests/test_master.py @@ -28,7 +28,7 @@ from exo.shared.types.profiling import ( ) from exo.shared.types.tasks import TaskStatus from exo.shared.types.tasks import TextGeneration as TextGenerationTask -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.instances import ( InstanceMeta, MlxRingInstance, @@ -136,7 +136,9 @@ async def test_master(): command_id=CommandId(), task_params=TextGenerationTaskParams( model=ModelId("llama-3.2-1b"), - input="Hello, how are you?", + input=[ + InputMessage(role="user", content="Hello, how are you?") + ], ), ) ), @@ -189,7 +191,7 @@ async def test_master(): assert isinstance(events[2].event.task, TextGenerationTask) assert events[2].event.task.task_params == TextGenerationTaskParams( model=ModelId("llama-3.2-1b"), - input="Hello, how are you?", + input=[InputMessage(role="user", content="Hello, how are you?")], ) await master.shutdown() diff --git a/src/exo/master/tests/test_placement_utils.py b/src/exo/master/tests/test_placement_utils.py index f2cb1067..245c4fd7 100644 --- a/src/exo/master/tests/test_placement_utils.py +++ b/src/exo/master/tests/test_placement_utils.py @@ -5,6 +5,7 @@ from exo.master.placement_utils import ( filter_cycles_by_memory, get_mlx_jaccl_coordinators, get_shard_assignments, + get_shard_assignments_for_pipeline_parallel, get_smallest_cycles, ) from exo.master.tests.conftest import ( @@ -20,7 +21,11 @@ from exo.shared.types.profiling import ( NodeNetworkInfo, ) from exo.shared.types.topology import Connection, SocketConnection -from exo.shared.types.worker.shards import Sharding +from exo.shared.types.worker.shards import ( + CfgShardMetadata, + PipelineShardMetadata, + Sharding, +) def test_filter_cycles_by_memory(): @@ -487,3 +492,193 @@ def test_get_shard_assignments_insufficient_memory_raises(): get_shard_assignments( model_card, selected_cycle, Sharding.Pipeline, node_memory ) + + +class TestCfgParallelPlacement: + def _create_ring_topology(self, node_ids: list[NodeId]) -> Topology: + topology = Topology() + for node_id in node_ids: + topology.add_node(node_id) + + for i, node_id in enumerate(node_ids): + next_node = node_ids[(i + 1) % len(node_ids)] + conn = Connection( + source=node_id, + sink=next_node, + edge=create_socket_connection(i + 1), + ) + topology.add_connection(conn) + + return topology + + def test_two_nodes_cfg_model_uses_cfg_parallel(self): + """Two nodes with CFG model should use CFG parallel (no pipeline).""" + node_a = NodeId() + node_b = NodeId() + + topology = self._create_ring_topology([node_a, node_b]) + cycles = [c for c in topology.get_cycles() if len(c) == 2] + cycle = cycles[0] + + node_memory = { + node_a: create_node_memory(1000 * 1024), + node_b: create_node_memory(1000 * 1024), + } + + model_card = ModelCard( + model_id=ModelId("qwen-image-test"), + n_layers=60, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=True, + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 2 + + # CFG models should get CfgShardMetadata + for shard in shards: + assert isinstance(shard, CfgShardMetadata) + # Both nodes should have all layers (no pipeline split) + assert shard.start_layer == 0 + assert shard.end_layer == 60 + assert shard.cfg_world_size == 2 + # Each node is the only stage in its pipeline group + assert shard.pipeline_world_size == 1 + assert shard.pipeline_rank == 0 + + cfg_ranks = sorted( + s.cfg_rank for s in shards if isinstance(s, CfgShardMetadata) + ) + assert cfg_ranks == [0, 1] + + def test_four_nodes_cfg_model_uses_hybrid(self): + """Four nodes with CFG model should use 2 CFG groups x 2 pipeline stages.""" + nodes = [NodeId() for _ in range(4)] + + topology = self._create_ring_topology(nodes) + cycles = [c for c in topology.get_cycles() if len(c) == 4] + cycle = cycles[0] + + node_memory = {n: create_node_memory(1000 * 1024) for n in nodes} + + model_card = ModelCard( + model_id=ModelId("qwen-image-test"), + n_layers=60, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=True, + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 4 + + # CFG models should get CfgShardMetadata + for shard in shards: + assert isinstance(shard, CfgShardMetadata) + assert shard.cfg_world_size == 2 + assert shard.pipeline_world_size == 2 + assert shard.pipeline_rank in [0, 1] + + # Check we have 2 nodes in each CFG group + cfg_0_shards = [ + s for s in shards if isinstance(s, CfgShardMetadata) and s.cfg_rank == 0 + ] + cfg_1_shards = [ + s for s in shards if isinstance(s, CfgShardMetadata) and s.cfg_rank == 1 + ] + assert len(cfg_0_shards) == 2 + assert len(cfg_1_shards) == 2 + + # Both CFG groups should have the same layer assignments + cfg_0_layers = [(s.start_layer, s.end_layer) for s in cfg_0_shards] + cfg_1_layers = [(s.start_layer, s.end_layer) for s in cfg_1_shards] + assert sorted(cfg_0_layers) == sorted(cfg_1_layers) + + def test_three_nodes_cfg_model_uses_sequential_cfg(self): + """Three nodes (odd) with CFG model should use sequential CFG (PipelineShardMetadata).""" + nodes = [NodeId() for _ in range(3)] + + topology = self._create_ring_topology(nodes) + cycles = [c for c in topology.get_cycles() if len(c) == 3] + cycle = cycles[0] + + node_memory = {n: create_node_memory(1000 * 1024) for n in nodes} + + model_card = ModelCard( + model_id=ModelId("qwen-image-test"), + n_layers=60, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=True, + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 3 + + # Odd node count with CFG model falls back to PipelineShardMetadata (sequential CFG) + for shard in shards: + assert isinstance(shard, PipelineShardMetadata) + + def test_two_nodes_non_cfg_model_uses_pipeline(self): + """Two nodes with non-CFG model should use pure pipeline (PipelineShardMetadata).""" + node_a = NodeId() + node_b = NodeId() + + topology = self._create_ring_topology([node_a, node_b]) + cycles = [c for c in topology.get_cycles() if len(c) == 2] + cycle = cycles[0] + + node_memory = { + node_a: create_node_memory(1000 * 1024), + node_b: create_node_memory(1000 * 1024), + } + + model_card = ModelCard( + model_id=ModelId("flux-test"), + n_layers=57, + storage_size=Memory.from_kb(1000), + hidden_size=1, + supports_tensor=False, + uses_cfg=False, # Non-CFG model + tasks=[ModelTask.TextToImage], + ) + + assignments = get_shard_assignments_for_pipeline_parallel( + model_card, cycle, node_memory + ) + + shards = list(assignments.runner_to_shard.values()) + assert len(shards) == 2 + + # Non-CFG models should get PipelineShardMetadata + for shard in shards: + assert isinstance(shard, PipelineShardMetadata) + + # Should have actual layer sharding (pipeline) + layer_ranges = sorted( + (s.start_layer, s.end_layer) + for s in shards + if isinstance(s, PipelineShardMetadata) + ) + # First shard starts at 0, last shard ends at 57 + assert layer_ranges[0][0] == 0 + assert layer_ranges[-1][1] == 57 diff --git a/src/exo/shared/constants.py b/src/exo/shared/constants.py index 39438cbd..b385b5e8 100644 --- a/src/exo/shared/constants.py +++ b/src/exo/shared/constants.py @@ -39,7 +39,7 @@ RESOURCES_DIR = ( ) _DASHBOARD_DIR_ENV = os.environ.get("EXO_DASHBOARD_DIR", None) DASHBOARD_DIR = ( - find_dashboard() if _RESOURCES_DIR_ENV is None else Path.home() / _RESOURCES_DIR_ENV + find_dashboard() if _DASHBOARD_DIR_ENV is None else Path.home() / _DASHBOARD_DIR_ENV ) # Log files (data/logs or cache) @@ -58,6 +58,8 @@ LIBP2P_COMMANDS_TOPIC = "commands" EXO_MAX_CHUNK_SIZE = 512 * 1024 +EXO_CUSTOM_MODEL_CARDS_DIR = EXO_DATA_HOME / "custom_model_cards" + EXO_IMAGE_CACHE_DIR = EXO_CACHE_HOME / "images" EXO_TRACING_CACHE_DIR = EXO_CACHE_HOME / "traces" diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index bb041610..1291c2ec 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -18,14 +18,19 @@ from pydantic import ( ) from tomlkit.exceptions import TOMLKitError -from exo.shared.constants import EXO_ENABLE_IMAGE_MODELS, RESOURCES_DIR +from exo.shared.constants import ( + EXO_CUSTOM_MODEL_CARDS_DIR, + EXO_ENABLE_IMAGE_MODELS, + RESOURCES_DIR, +) from exo.shared.types.common import ModelId from exo.shared.types.memory import Memory from exo.utils.pydantic_ext import CamelCaseModel # kinda ugly... # TODO: load search path from config.toml -_csp = [Path(RESOURCES_DIR) / "inference_model_cards"] +_custom_cards_dir = Path(str(EXO_CUSTOM_MODEL_CARDS_DIR)) +_csp = [Path(RESOURCES_DIR) / "inference_model_cards", _custom_cards_dir] if EXO_ENABLE_IMAGE_MODELS: _csp.append(Path(RESOURCES_DIR) / "image_model_cards") @@ -60,9 +65,9 @@ class ComponentInfo(CamelCaseModel): component_name: str component_path: str storage_size: Memory - n_layers: PositiveInt | None + n_layers: PositiveInt | None = None can_shard: bool - safetensors_index_filename: str | None + safetensors_index_filename: str | None = None class ModelCard(CamelCaseModel): @@ -73,6 +78,11 @@ class ModelCard(CamelCaseModel): supports_tensor: bool tasks: list[ModelTask] components: list[ComponentInfo] | None = None + family: str = "" + quantization: str = "" + base_model: str = "" + capabilities: list[str] = [] + uses_cfg: bool = False @field_validator("tasks", mode="before") @classmethod @@ -85,8 +95,9 @@ class ModelCard(CamelCaseModel): data = tomlkit.dumps(py) # pyright: ignore[reportUnknownMemberType] await f.write(data) - async def save_to_default_path(self): - await self.save(Path(RESOURCES_DIR) / (self.model_id.normalize() + ".toml")) + async def save_to_custom_dir(self) -> None: + await aios.makedirs(str(_custom_cards_dir), exist_ok=True) + await self.save(_custom_cards_dir / (self.model_id.normalize() + ".toml")) @staticmethod async def load_from_path(path: Path) -> "ModelCard": @@ -108,9 +119,9 @@ class ModelCard(CamelCaseModel): async def fetch_from_hf(model_id: ModelId) -> "ModelCard": """Fetches storage size and number of layers for a Hugging Face model, returns Pydantic ModelMeta.""" # TODO: failure if files do not exist - config_data = await get_config_data(model_id) + config_data = await fetch_config_data(model_id) num_layers = config_data.layer_count - mem_size_bytes = await get_safetensors_size(model_id) + mem_size_bytes = await fetch_safetensors_size(model_id) mc = ModelCard( model_id=ModelId(model_id), @@ -120,90 +131,29 @@ class ModelCard(CamelCaseModel): supports_tensor=config_data.supports_tensor, tasks=[ModelTask.TextGeneration], ) - await mc.save_to_default_path() + await mc.save_to_custom_dir() _card_cache[model_id] = mc return mc -# TODO: quantizing and dynamically creating model cards -def _generate_image_model_quant_variants( # pyright: ignore[reportUnusedFunction] - base_name: str, - base_card: ModelCard, -) -> dict[str, ModelCard]: - """Create quantized variants of an image model 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 - Only the transformer component is quantized; text encoders stay at bf16. - Sizes are calculated exactly from the base card's component sizes. - """ - if base_card.components is None: - raise ValueError(f"Image model {base_name} must have components defined") - # quantizations = [8, 6, 5, 4, 3] - quantizations = [8, 4] +def is_custom_card(model_id: ModelId) -> bool: + """Check if a model card exists in the custom cards directory.""" + import os - num_transformer_bytes = next( - c.storage_size.in_bytes - for c in base_card.components - if c.component_name == "transformer" + card_path = Path(str(EXO_CUSTOM_MODEL_CARDS_DIR)) / ( + ModelId(model_id).normalize() + ".toml" ) - - transformer_bytes = Memory.from_bytes(num_transformer_bytes) - - remaining_bytes = Memory.from_bytes( - sum( - c.storage_size.in_bytes - for c in base_card.components - if c.component_name != "transformer" - ) - ) - - def with_transformer_size(new_size: Memory) -> list[ComponentInfo]: - assert base_card.components is not None - return [ - ComponentInfo( - component_name=c.component_name, - component_path=c.component_path, - storage_size=new_size - if c.component_name == "transformer" - else c.storage_size, - n_layers=c.n_layers, - can_shard=c.can_shard, - safetensors_index_filename=c.safetensors_index_filename, - ) - for c in base_card.components - ] - - variants = { - base_name: ModelCard( - model_id=base_card.model_id, - storage_size=transformer_bytes + remaining_bytes, - n_layers=base_card.n_layers, - hidden_size=base_card.hidden_size, - supports_tensor=base_card.supports_tensor, - tasks=base_card.tasks, - components=with_transformer_size(transformer_bytes), - ) - } - - for quant in quantizations: - quant_transformer_bytes = Memory.from_bytes( - (num_transformer_bytes * quant) // 16 - ) - total_bytes = remaining_bytes + quant_transformer_bytes - - model_id = ModelId(base_card.model_id + f"-{quant}bit") - - variants[f"{base_name}-{quant}bit"] = ModelCard( - model_id=model_id, - storage_size=total_bytes, - n_layers=base_card.n_layers, - hidden_size=base_card.hidden_size, - supports_tensor=base_card.supports_tensor, - tasks=base_card.tasks, - components=with_transformer_size(quant_transformer_bytes), - ) - - return variants + return os.path.isfile(str(card_path)) class ConfigData(BaseModel): @@ -259,7 +209,7 @@ class ConfigData(BaseModel): return data -async def get_config_data(model_id: ModelId) -> ConfigData: +async def fetch_config_data(model_id: ModelId) -> ConfigData: """Downloads and parses config.json for a model.""" from exo.download.download_utils import ( download_file_with_retry, @@ -281,7 +231,7 @@ async def get_config_data(model_id: ModelId) -> ConfigData: return ConfigData.model_validate_json(await f.read()) -async def get_safetensors_size(model_id: ModelId) -> Memory: +async def fetch_safetensors_size(model_id: ModelId) -> Memory: """Gets model size from safetensors index or falls back to HF API.""" from exo.download.download_utils import ( download_file_with_retry, diff --git a/src/exo/shared/types/api.py b/src/exo/shared/types/api.py index 40dbb288..5a7bae1e 100644 --- a/src/exo/shared/types/api.py +++ b/src/exo/shared/types/api.py @@ -42,6 +42,11 @@ class ModelListModel(BaseModel): storage_size_megabytes: int = Field(default=0) supports_tensor: bool = Field(default=False) tasks: list[str] = Field(default=[]) + is_custom: bool = Field(default=False) + family: str = Field(default="") + quantization: str = Field(default="") + base_model: str = Field(default="") + capabilities: list[str] = Field(default_factory=list) class ModelList(BaseModel): @@ -201,6 +206,19 @@ class BenchChatCompletionRequest(ChatCompletionRequest): pass +class AddCustomModelParams(BaseModel): + model_id: ModelId + + +class HuggingFaceSearchResult(BaseModel): + id: str + author: str = "" + downloads: int = 0 + likes: int = 0 + last_modified: str = "" + tags: list[str] = Field(default_factory=list) + + class PlaceInstanceParams(BaseModel): model_id: ModelId sharding: Sharding = Sharding.Pipeline diff --git a/src/exo/shared/types/chunks.py b/src/exo/shared/types/chunks.py index e96dbc9d..5fe9eb1c 100644 --- a/src/exo/shared/types/chunks.py +++ b/src/exo/shared/types/chunks.py @@ -2,7 +2,12 @@ from collections.abc import Generator from typing import Any, Literal from exo.shared.models.model_cards import ModelId -from exo.shared.types.api import GenerationStats, ImageGenerationStats, Usage +from exo.shared.types.api import ( + GenerationStats, + ImageGenerationStats, + TopLogprobItem, + Usage, +) from exo.utils.pydantic_ext import TaggedModel from .api import FinishReason @@ -20,6 +25,8 @@ class TokenChunk(BaseChunk): usage: Usage | None finish_reason: Literal["stop", "length", "content_filter"] | None = None stats: GenerationStats | None = None + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None class ErrorChunk(BaseChunk): diff --git a/src/exo/shared/types/text_generation.py b/src/exo/shared/types/text_generation.py index b9c5565c..3e7b89fd 100644 --- a/src/exo/shared/types/text_generation.py +++ b/src/exo/shared/types/text_generation.py @@ -28,7 +28,7 @@ class TextGenerationTaskParams(BaseModel, frozen=True): """ model: ModelId - input: str | list[InputMessage] + input: list[InputMessage] instructions: str | None = None max_output_tokens: int | None = None temperature: float | None = None @@ -40,3 +40,5 @@ class TextGenerationTaskParams(BaseModel, frozen=True): stop: str | list[str] | None = None seed: int | None = None chat_template_messages: list[dict[str, Any]] | None = None + logprobs: bool = False + top_logprobs: int | None = None diff --git a/src/exo/shared/types/worker/runner_response.py b/src/exo/shared/types/worker/runner_response.py index 5dfbe547..d1bea77e 100644 --- a/src/exo/shared/types/worker/runner_response.py +++ b/src/exo/shared/types/worker/runner_response.py @@ -6,6 +6,7 @@ from exo.shared.types.api import ( GenerationStats, ImageGenerationStats, ToolCallItem, + TopLogprobItem, Usage, ) from exo.utils.pydantic_ext import TaggedModel @@ -22,7 +23,8 @@ class TokenizedResponse(BaseRunnerResponse): class GenerationResponse(BaseRunnerResponse): text: str token: int - # logprobs: list[float] | None = None # too big. we can change to be top-k + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None finish_reason: FinishReason | None = None stats: GenerationStats | None = None usage: Usage | None diff --git a/src/exo/shared/types/worker/shards.py b/src/exo/shared/types/worker/shards.py index 8bb23a57..59a6c54e 100644 --- a/src/exo/shared/types/worker/shards.py +++ b/src/exo/shared/types/worker/shards.py @@ -1,4 +1,5 @@ from enum import Enum +from typing import TypeAlias, final from pydantic import Field @@ -51,6 +52,7 @@ class BaseShardMetadata(TaggedModel): ) +@final class PipelineShardMetadata(BaseShardMetadata): """ Pipeline parallelism shard meta. @@ -60,8 +62,23 @@ class PipelineShardMetadata(BaseShardMetadata): """ +@final +class CfgShardMetadata(BaseShardMetadata): + """Shard metadata for CFG-parallel image generation models.""" + + cfg_rank: int # 0 = positive branch, 1 = negative branch + cfg_world_size: int = 2 + + # Pipeline-relative coordinates (computed at placement time) + pipeline_rank: int # rank within the pipeline group (0, 1, 2, ...) + pipeline_world_size: int # number of nodes per pipeline group + + +@final class TensorShardMetadata(BaseShardMetadata): pass -ShardMetadata = PipelineShardMetadata | TensorShardMetadata +ShardMetadata: TypeAlias = ( + PipelineShardMetadata | CfgShardMetadata | TensorShardMetadata +) diff --git a/src/exo/worker/engines/image/distributed_model.py b/src/exo/worker/engines/image/distributed_model.py index bafa9319..8c9bd04c 100644 --- a/src/exo/worker/engines/image/distributed_model.py +++ b/src/exo/worker/engines/image/distributed_model.py @@ -9,7 +9,7 @@ from PIL import Image from exo.download.download_utils import build_model_path from exo.shared.types.api import AdvancedImageParams from exo.shared.types.worker.instances import BoundInstance -from exo.shared.types.worker.shards import PipelineShardMetadata +from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata from exo.worker.engines.image.config import ImageModelConfig from exo.worker.engines.image.models import ( create_adapter_for_model, @@ -30,14 +30,19 @@ class DistributedImageModel: self, model_id: str, local_path: Path, - shard_metadata: PipelineShardMetadata, + shard_metadata: PipelineShardMetadata | CfgShardMetadata, group: Optional[mx.distributed.Group] = None, quantize: int | None = None, ): config = get_config_for_model(model_id) adapter = create_adapter_for_model(config, model_id, local_path, quantize) - if group is not None: + has_layer_sharding = ( + shard_metadata.start_layer != 0 + or shard_metadata.end_layer != shard_metadata.n_layers + ) + + if group is not None and has_layer_sharding: adapter.slice_transformer_blocks( start_layer=shard_metadata.start_layer, end_layer=shard_metadata.end_layer, @@ -75,8 +80,10 @@ class DistributedImageModel: model_path = build_model_path(model_id) shard_metadata = bound_instance.bound_shard - if not isinstance(shard_metadata, PipelineShardMetadata): - raise ValueError("Expected PipelineShardMetadata for image generation") + if not isinstance(shard_metadata, (PipelineShardMetadata, CfgShardMetadata)): + raise ValueError( + "Expected PipelineShardMetadata or CfgShardMetadata for image generation" + ) is_distributed = ( len(bound_instance.instance.shard_assignments.node_to_runner) > 1 diff --git a/src/exo/worker/engines/image/models/base.py b/src/exo/worker/engines/image/models/base.py index 90439823..f77ea882 100644 --- a/src/exo/worker/engines/image/models/base.py +++ b/src/exo/worker/engines/image/models/base.py @@ -86,6 +86,27 @@ class PromptData(ABC): """ ... + @abstractmethod + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + """Get embeddings for a single CFG branch (positive or negative). + + Used for sequential CFG and CFG parallel modes where we process + one branch at a time instead of batching. + + Args: + positive: True for positive prompt, False for negative prompt + + Returns: + Tuple of: + - embeds: [1, seq, hidden] prompt embeddings + - mask: [1, seq] attention mask or None + - pooled: [1, hidden] pooled embeddings or None + - conditioning_latents: [1, latent_seq, latent_dim] or None + """ + ... + class ModelAdapter(ABC, Generic[ModelT, TransformerT]): _config: ImageModelConfig diff --git a/src/exo/worker/engines/image/models/flux/adapter.py b/src/exo/worker/engines/image/models/flux/adapter.py index be9b43f8..1aa510da 100644 --- a/src/exo/worker/engines/image/models/flux/adapter.py +++ b/src/exo/worker/engines/image/models/flux/adapter.py @@ -64,6 +64,12 @@ class FluxPromptData(PromptData): ) -> tuple[mx.array, mx.array, mx.array | None, mx.array | None] | None: return None + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + """Flux doesn't use CFG, but we return positive data for compatibility.""" + return (self._prompt_embeds, None, self._pooled_prompt_embeds, None) + class FluxModelAdapter(ModelAdapter[Flux1, Transformer]): def __init__( diff --git a/src/exo/worker/engines/image/models/qwen/adapter.py b/src/exo/worker/engines/image/models/qwen/adapter.py index d9f009ec..e88d2a75 100644 --- a/src/exo/worker/engines/image/models/qwen/adapter.py +++ b/src/exo/worker/engines/image/models/qwen/adapter.py @@ -133,6 +133,24 @@ class QwenPromptData(PromptData): return batched_embeds, batched_mask, None, cond_latents + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + if positive: + return ( + self._prompt_embeds, + self._prompt_mask, + None, + self.conditioning_latents, + ) + else: + return ( + self._negative_prompt_embeds, + self._negative_prompt_mask, + None, + self.conditioning_latents, + ) + class QwenModelAdapter(ModelAdapter[QwenImage, QwenTransformer]): """Adapter for Qwen-Image model. diff --git a/src/exo/worker/engines/image/models/qwen/config.py b/src/exo/worker/engines/image/models/qwen/config.py index d5da1bac..4ec2cb35 100644 --- a/src/exo/worker/engines/image/models/qwen/config.py +++ b/src/exo/worker/engines/image/models/qwen/config.py @@ -12,7 +12,7 @@ QWEN_IMAGE_CONFIG = ImageModelConfig( ), ), default_steps={"low": 10, "medium": 25, "high": 50}, - num_sync_steps_factor=0.125, # ~3 sync steps for medium (30 steps) + num_sync_steps_factor=0.25, guidance_scale=3.5, # Set to None or < 1.0 to disable CFG ) @@ -24,6 +24,6 @@ QWEN_IMAGE_EDIT_CONFIG = ImageModelConfig( ), ), default_steps={"low": 10, "medium": 25, "high": 50}, - num_sync_steps_factor=0.125, + num_sync_steps_factor=0.25, guidance_scale=3.5, ) diff --git a/src/exo/worker/engines/image/models/qwen/edit_adapter.py b/src/exo/worker/engines/image/models/qwen/edit_adapter.py index e327eb0c..4a88a4e3 100644 --- a/src/exo/worker/engines/image/models/qwen/edit_adapter.py +++ b/src/exo/worker/engines/image/models/qwen/edit_adapter.py @@ -153,6 +153,24 @@ class QwenEditPromptData(PromptData): return batched_embeds, batched_mask, None, batched_cond_latents + def get_cfg_branch_data( + self, positive: bool + ) -> tuple[mx.array, mx.array | None, mx.array | None, mx.array | None]: + if positive: + return ( + self._prompt_embeds, + self._prompt_mask, + None, + self._conditioning_latents, + ) + else: + return ( + self._negative_prompt_embeds, + self._negative_prompt_mask, + None, + self._conditioning_latents, + ) + class QwenEditModelAdapter(ModelAdapter[QwenImageEdit, QwenTransformer]): """Adapter for Qwen-Image-Edit model. diff --git a/src/exo/worker/engines/image/pipeline/runner.py b/src/exo/worker/engines/image/pipeline/runner.py index e1f65efd..f7054763 100644 --- a/src/exo/worker/engines/image/pipeline/runner.py +++ b/src/exo/worker/engines/image/pipeline/runner.py @@ -1,5 +1,7 @@ +from collections.abc import Iterator +from dataclasses import dataclass from math import ceil -from typing import Any, Optional +from typing import Any, Optional, final import mlx.core as mx from mflux.models.common.config.config import Config @@ -11,7 +13,7 @@ from exo.shared.tracing import ( clear_trace_buffer, trace, ) -from exo.shared.types.worker.shards import PipelineShardMetadata +from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata from exo.worker.engines.image.config import ImageModelConfig from exo.worker.engines.image.models.base import ( ModelAdapter, @@ -25,6 +27,16 @@ from exo.worker.engines.image.pipeline.block_wrapper import ( ) +@final +@dataclass(frozen=True) +class CfgBranch: + positive: bool + embeds: mx.array + mask: mx.array | None + pooled: mx.array | None + cond_latents: mx.array | None + + def calculate_patch_heights( latent_height: int, num_patches: int ) -> tuple[list[int], int]: @@ -70,29 +82,18 @@ class DiffusionRunner: config: ImageModelConfig, adapter: ModelAdapter[Any, Any], group: Optional[mx.distributed.Group], - shard_metadata: PipelineShardMetadata, + shard_metadata: PipelineShardMetadata | CfgShardMetadata, num_patches: Optional[int] = None, ): self.config = config self.adapter = adapter self.group = group - if group is None: - self.rank = 0 - self.world_size = 1 - self.next_rank = 0 - self.prev_rank = 0 - self.start_layer = 0 - self.end_layer = config.total_blocks - else: - self.rank = shard_metadata.device_rank - self.world_size = shard_metadata.world_size - self.next_rank = (self.rank + 1) % self.world_size - self.prev_rank = (self.rank - 1 + self.world_size) % self.world_size - self.start_layer = shard_metadata.start_layer - self.end_layer = shard_metadata.end_layer + self._init_cfg_topology(shard_metadata) - self.num_patches = num_patches if num_patches else max(1, self.world_size) + self.num_patches = ( + num_patches if num_patches else max(1, self.pipeline_world_size) + ) self.total_joint = config.joint_block_count self.total_single = config.single_block_count @@ -102,6 +103,97 @@ class DiffusionRunner: self._compute_assigned_blocks() + def _init_cfg_topology( + self, shard_metadata: PipelineShardMetadata | CfgShardMetadata + ) -> None: + """Initialize CFG and pipeline topology from shard metadata. + + Both CfgShardMetadata and PipelineShardMetadata represent pipeline parallel + execution. CFG adds a second parallel pipeline for negative prompt processing, + but within each pipeline group the communication pattern is identical. + """ + if self.group is None: + # Single node - no distributed communication + self.rank = 0 + self.world_size = 1 + self.start_layer = 0 + self.end_layer = self.config.total_blocks + self.cfg_rank = 0 + self.cfg_world_size = 1 + self.cfg_parallel = False + self.pipeline_rank = 0 + self.pipeline_world_size = 1 + self.next_pipeline_rank: int | None = None + self.prev_pipeline_rank: int | None = None + self.cfg_peer_rank: int | None = None + self.first_pipeline_rank: int = 0 + self.last_pipeline_rank: int = 0 + return + + # Common fields from base metadata + self.rank = shard_metadata.device_rank + self.world_size = shard_metadata.world_size + self.start_layer = shard_metadata.start_layer + self.end_layer = shard_metadata.end_layer + + if isinstance(shard_metadata, CfgShardMetadata): + # CFG parallel: two independent pipelines + self.cfg_rank = shard_metadata.cfg_rank + self.cfg_world_size = shard_metadata.cfg_world_size + self.cfg_parallel = True + self.pipeline_rank = shard_metadata.pipeline_rank + self.pipeline_world_size = shard_metadata.pipeline_world_size + else: + # Pure pipeline: single pipeline group, sequential CFG + self.cfg_rank = 0 + self.cfg_world_size = 1 + self.cfg_parallel = False + self.pipeline_rank = shard_metadata.device_rank + self.pipeline_world_size = shard_metadata.world_size + + # Pipeline neighbor computation (same logic for both types) + is_first = self.pipeline_rank == 0 + is_last = self.pipeline_rank == self.pipeline_world_size - 1 + + self.next_pipeline_rank = ( + None + if is_last + else self._device_rank_for(self.cfg_rank, self.pipeline_rank + 1) + ) + self.prev_pipeline_rank = ( + None + if is_first + else self._device_rank_for(self.cfg_rank, self.pipeline_rank - 1) + ) + + # CFG peer is the corresponding last stage in the other CFG group + if self.cfg_parallel and is_last: + other_cfg_rank = 1 - self.cfg_rank + self.cfg_peer_rank = self._device_rank_for( + other_cfg_rank, self.pipeline_rank + ) + else: + self.cfg_peer_rank = None + + # First/last pipeline ranks for ring communication (latent broadcast) + self.first_pipeline_rank = self._device_rank_for(self.cfg_rank, 0) + self.last_pipeline_rank = self._device_rank_for( + self.cfg_rank, self.pipeline_world_size - 1 + ) + + def _device_rank_for(self, cfg_rank: int, pipeline_rank: int) -> int: + """Convert (cfg_rank, pipeline_rank) to device_rank in the ring topology. + + Ring layout: [cfg0_pipe0, cfg0_pipe1, ..., cfg1_pipeN-1, cfg1_pipeN-2, ..., cfg1_pipe0] + Group 0 is in ascending order, group 1 is reversed so last stages are neighbors. + """ + if not self.cfg_parallel: + return pipeline_rank + if cfg_rank == 0: + return pipeline_rank + else: + return self.world_size - 1 - pipeline_rank + def _compute_assigned_blocks(self) -> None: """Determine which joint/single blocks this stage owns.""" start = self.start_layer @@ -138,11 +230,11 @@ class DiffusionRunner: @property def is_first_stage(self) -> bool: - return self.rank == 0 + return self.pipeline_rank == 0 @property def is_last_stage(self) -> bool: - return self.rank == self.world_size - 1 + return self.pipeline_rank == self.pipeline_world_size - 1 @property def is_distributed(self) -> bool: @@ -153,6 +245,97 @@ class DiffusionRunner: return self._guidance_override return self.config.guidance_scale + def _get_cfg_branches(self, prompt_data: PromptData) -> Iterator[CfgBranch]: + """Yield the CFG branches this node should process. + + - No CFG: yields one branch (positive) + - CFG parallel: yields one branch (our assigned branch) + - Sequential CFG: yields two branches (positive, then negative) + """ + if not self.adapter.needs_cfg: + embeds, mask, pooled, cond = prompt_data.get_cfg_branch_data(positive=True) + yield CfgBranch( + positive=True, + embeds=embeds, + mask=mask, + pooled=pooled, + cond_latents=cond, + ) + elif self.cfg_parallel: + positive = self.cfg_rank == 0 + embeds, mask, pooled, cond = prompt_data.get_cfg_branch_data(positive) + yield CfgBranch( + positive=positive, + embeds=embeds, + mask=mask, + pooled=pooled, + cond_latents=cond, + ) + else: + pos_embeds, pos_mask, pos_pooled, pos_cond = ( + prompt_data.get_cfg_branch_data(positive=True) + ) + yield CfgBranch( + positive=True, + embeds=pos_embeds, + mask=pos_mask, + pooled=pos_pooled, + cond_latents=pos_cond, + ) + neg_embeds, neg_mask, neg_pooled, neg_cond = ( + prompt_data.get_cfg_branch_data(positive=False) + ) + yield CfgBranch( + positive=False, + embeds=neg_embeds, + mask=neg_mask, + pooled=neg_pooled, + cond_latents=neg_cond, + ) + + def _combine_cfg_results(self, results: list[tuple[bool, mx.array]]) -> mx.array: + if len(results) == 1: + positive, noise = results[0] + if self.cfg_parallel and self.is_last_stage: + # TODO(ciaran): try to remove + mx.eval(noise) + return self._exchange_and_apply_guidance(noise, positive) + return noise + + noise_neg = next(n for p, n in results if not p) + noise_pos = next(n for p, n in results if p) + return self._apply_guidance(noise_pos, noise_neg) + + def _exchange_and_apply_guidance( + self, noise: mx.array, is_positive: bool + ) -> mx.array: + assert self.group is not None + assert self.cfg_peer_rank is not None + + if is_positive: + noise = mx.distributed.send(noise, self.cfg_peer_rank, group=self.group) + mx.async_eval(noise) + noise_neg = mx.distributed.recv_like( + noise, self.cfg_peer_rank, group=self.group + ) + mx.eval(noise_neg) + noise_pos = noise + else: + noise_pos = mx.distributed.recv_like( + noise, self.cfg_peer_rank, group=self.group + ) + mx.eval(noise_pos) + noise = mx.distributed.send(noise, self.cfg_peer_rank, group=self.group) + mx.async_eval(noise) + noise_neg = noise + + return self._apply_guidance(noise_pos, noise_neg) + + def _apply_guidance(self, noise_pos: mx.array, noise_neg: mx.array) -> mx.array: + scale = self._get_effective_guidance_scale() + assert scale is not None + return self.adapter.apply_guidance(noise_pos, noise_neg, scale) + def _ensure_wrappers( self, text_seq_len: int, @@ -470,7 +653,9 @@ class DiffusionRunner: ) -> mx.array: if self.group is None: return self._single_node_step(t, config, latents, prompt_data) - elif t < config.init_time_step + num_sync_steps: + elif ( + self.pipeline_world_size == 1 or t < config.init_time_step + num_sync_steps + ): with trace(name=f"sync {t}", rank=self.rank, category="sync"): return self._sync_pipeline_step( t, @@ -496,42 +681,29 @@ class DiffusionRunner: prompt_data: PromptData, ) -> mx.array: cond_image_grid = prompt_data.cond_image_grid - needs_cfg = self.adapter.needs_cfg + results: list[tuple[bool, mx.array]] = [] + + for branch in self._get_cfg_branches(prompt_data): + # Reset caches before each branch to ensure no state contamination + self._reset_all_caches() - if needs_cfg: - batched_data = prompt_data.get_batched_cfg_data() - assert batched_data is not None, "CFG model must provide batched data" - prompt_embeds, encoder_mask, batched_pooled, cond_latents = batched_data pooled_embeds = ( - batched_pooled if batched_pooled is not None else prompt_embeds - ) - step_latents = mx.concatenate([latents, latents], axis=0) - else: - prompt_embeds = prompt_data.prompt_embeds - pooled_embeds = prompt_data.pooled_prompt_embeds - encoder_mask = prompt_data.get_encoder_hidden_states_mask(positive=True) - cond_latents = prompt_data.conditioning_latents - step_latents = latents - - noise = self._forward_pass( - step_latents, - prompt_embeds, - pooled_embeds, - t=t, - config=config, - encoder_hidden_states_mask=encoder_mask, - cond_image_grid=cond_image_grid, - conditioning_latents=cond_latents, - ) - - if needs_cfg: - noise_pos, noise_neg = mx.split(noise, 2, axis=0) - guidance_scale = self._get_effective_guidance_scale() - assert guidance_scale is not None - noise = self.adapter.apply_guidance( - noise_pos, noise_neg, guidance_scale=guidance_scale + branch.pooled if branch.pooled is not None else branch.embeds ) + noise = self._forward_pass( + latents, + branch.embeds, + pooled_embeds, + t=t, + config=config, + encoder_hidden_states_mask=branch.mask, + cond_image_grid=cond_image_grid, + conditioning_latents=branch.cond_latents, + ) + results.append((branch.positive, noise)) + + noise = self._combine_cfg_results(results) return config.scheduler.step(noise=noise, timestep=t, latents=latents) # pyright: ignore[reportAny] def _create_patches( @@ -582,7 +754,7 @@ class DiffusionRunner: ) text_embeddings = self.adapter.compute_text_embeddings( - t, config, pooled_prompt_embeds + t, config, pooled_prompt_embeds, hidden_states=hidden_states ) image_rotary_embeddings = self.adapter.compute_rotary_embeddings( prompt_embeds, @@ -594,19 +766,22 @@ class DiffusionRunner: if self.has_joint_blocks: if not self.is_first_stage: + assert self.prev_pipeline_rank is not None with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): hidden_states = mx.distributed.recv( (batch_size, num_img_tokens, hidden_dim), dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) encoder_hidden_states = mx.distributed.recv( (batch_size, text_seq_len, hidden_dim), dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(hidden_states, encoder_hidden_states) @@ -639,34 +814,45 @@ class DiffusionRunner: if self.has_single_blocks or self.is_last_stage: hidden_states = concatenated else: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): concatenated = mx.distributed.send( - concatenated, self.next_rank, group=self.group + concatenated, self.next_pipeline_rank, group=self.group ) mx.async_eval(concatenated) elif self.has_joint_blocks and not self.is_last_stage: assert encoder_hidden_states is not None - with trace(name=f"send {self.next_rank}", rank=self.rank, category="comms"): + assert self.next_pipeline_rank is not None + with trace( + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", + ): hidden_states = mx.distributed.send( - hidden_states, self.next_rank, group=self.group + hidden_states, self.next_pipeline_rank, group=self.group ) encoder_hidden_states = mx.distributed.send( - encoder_hidden_states, self.next_rank, group=self.group + encoder_hidden_states, self.next_pipeline_rank, group=self.group ) mx.async_eval(hidden_states, encoder_hidden_states) if self.has_single_blocks: if not self.owns_concat_stage and not self.is_first_stage: + assert self.prev_pipeline_rank is not None with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): hidden_states = mx.distributed.recv( (batch_size, text_seq_len + num_img_tokens, hidden_dim), dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(hidden_states) @@ -689,11 +875,14 @@ class DiffusionRunner: mx.eval(hidden_states) if not self.is_last_stage: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): hidden_states = mx.distributed.send( - hidden_states, self.next_rank, group=self.group + hidden_states, self.next_pipeline_rank, group=self.group ) mx.async_eval(hidden_states) @@ -716,83 +905,67 @@ class DiffusionRunner: kontext_image_ids: mx.array | None = None, ) -> mx.array: prev_latents = hidden_states - needs_cfg = self.adapter.needs_cfg cond_image_grid = prompt_data.cond_image_grid scaled_hidden_states = config.scheduler.scale_model_input(hidden_states, t) # pyright: ignore[reportAny] original_latent_tokens: int = scaled_hidden_states.shape[1] # pyright: ignore[reportAny] - if needs_cfg: - batched_data = prompt_data.get_batched_cfg_data() - assert batched_data is not None, "CFG model must provide batched data" - prompt_embeds, encoder_mask, batched_pooled, cond_latents = batched_data + results: list[tuple[bool, mx.array]] = [] + + for branch in self._get_cfg_branches(prompt_data): pooled_embeds = ( - batched_pooled if batched_pooled is not None else prompt_embeds + branch.pooled if branch.pooled is not None else branch.embeds ) - step_latents = mx.concatenate( - [scaled_hidden_states, scaled_hidden_states], axis=0 + + cond_latents = branch.cond_latents + if cond_latents is not None: + num_img_tokens: int = original_latent_tokens + cond_latents.shape[1] + else: + num_img_tokens = original_latent_tokens + + step_latents: mx.array = scaled_hidden_states # pyright: ignore[reportAny] + if self.is_first_stage and cond_latents is not None: + step_latents = mx.concatenate([step_latents, cond_latents], axis=1) + + text_seq_len = branch.embeds.shape[1] + self._ensure_wrappers(text_seq_len, branch.mask) + + noise = self._run_sync_pass( + t, + config, + step_latents, + branch.embeds, + pooled_embeds, + branch.mask, + cond_image_grid, + kontext_image_ids, + num_img_tokens, + original_latent_tokens, + cond_latents, ) - else: - prompt_embeds = prompt_data.prompt_embeds - pooled_embeds = prompt_data.pooled_prompt_embeds - encoder_mask = prompt_data.get_encoder_hidden_states_mask(positive=True) - cond_latents = prompt_data.conditioning_latents - step_latents = scaled_hidden_states # pyright: ignore[reportAny] - if cond_latents is not None: - num_img_tokens: int = original_latent_tokens + cond_latents.shape[1] - else: - num_img_tokens = original_latent_tokens - - if self.is_first_stage and cond_latents is not None: - step_latents = mx.concatenate([step_latents, cond_latents], axis=1) - - text_seq_len = prompt_embeds.shape[1] - self._ensure_wrappers(text_seq_len, encoder_mask) - - noise = self._run_sync_pass( - t, - config, - step_latents, - prompt_embeds, - pooled_embeds, - encoder_mask, - cond_image_grid, - kontext_image_ids, - num_img_tokens, - original_latent_tokens, - cond_latents, - ) + if self.is_last_stage: + assert noise is not None + results.append((branch.positive, noise)) if self.is_last_stage: - assert noise is not None - if needs_cfg: - noise_pos, noise_neg = mx.split(noise, 2, axis=0) - guidance_scale = self._get_effective_guidance_scale() - assert guidance_scale is not None - noise = self.adapter.apply_guidance( - noise_pos, noise_neg, guidance_scale - ) + noise = self._combine_cfg_results(results) hidden_states = config.scheduler.step( # pyright: ignore[reportAny] noise=noise, timestep=t, latents=prev_latents ) if not self.is_first_stage: - with trace(name="send 0", rank=self.rank, category="comms"): - hidden_states = mx.distributed.send( - hidden_states, 0, group=self.group - ) - mx.async_eval(hidden_states) + hidden_states = mx.distributed.send( + hidden_states, self.first_pipeline_rank, group=self.group + ) + mx.async_eval(hidden_states) elif self.is_first_stage: - with trace( - name=f"recv {self.world_size - 1}", rank=self.rank, category="comms" - ): - hidden_states = mx.distributed.recv_like( - prev_latents, src=self.world_size - 1, group=self.group - ) - mx.eval(hidden_states) + hidden_states = mx.distributed.recv_like( + prev_latents, src=self.last_pipeline_rank, group=self.group + ) + mx.eval(hidden_states) else: hidden_states = prev_latents @@ -809,39 +982,10 @@ class DiffusionRunner: kontext_image_ids: mx.array | None = None, ) -> mx.array: patch_latents, token_indices = self._create_patches(latents, config) - needs_cfg = self.adapter.needs_cfg cond_image_grid = prompt_data.cond_image_grid - if needs_cfg: - batched_data = prompt_data.get_batched_cfg_data() - assert batched_data is not None, "CFG model must provide batched data" - prompt_embeds, encoder_mask, batched_pooled, _ = batched_data - pooled_embeds = ( - batched_pooled if batched_pooled is not None else prompt_embeds - ) - else: - prompt_embeds = prompt_data.prompt_embeds - pooled_embeds = prompt_data.pooled_prompt_embeds - encoder_mask = prompt_data.get_encoder_hidden_states_mask(positive=True) - - text_seq_len = prompt_embeds.shape[1] - self._ensure_wrappers(text_seq_len, encoder_mask) - self._set_text_seq_len(text_seq_len) - - if self.joint_block_wrappers: - for wrapper in self.joint_block_wrappers: - wrapper.set_encoder_mask(encoder_mask) - - text_embeddings = self.adapter.compute_text_embeddings(t, config, pooled_embeds) - image_rotary_embeddings = self.adapter.compute_rotary_embeddings( - prompt_embeds, - config, - encoder_hidden_states_mask=encoder_mask, - cond_image_grid=cond_image_grid, - kontext_image_ids=kontext_image_ids, - ) - prev_patch_latents = [p for p in patch_latents] + encoder_hidden_states: mx.array | None = None for patch_idx in range(len(patch_latents)): @@ -853,34 +997,57 @@ class DiffusionRunner: and not is_first_async_step ): with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.last_pipeline_rank}", + rank=self.rank, + category="comms", ): patch = mx.distributed.recv_like( - patch, src=self.prev_rank, group=self.group + patch, src=self.last_pipeline_rank, group=self.group ) mx.eval(patch) - step_patch = mx.concatenate([patch, patch], axis=0) if needs_cfg else patch + results: list[tuple[bool, mx.array]] = [] - noise, encoder_hidden_states = self._run_single_patch_pass( - patch=step_patch, - patch_idx=patch_idx, - token_indices=token_indices[patch_idx], - prompt_embeds=prompt_embeds, - text_embeddings=text_embeddings, - image_rotary_embeddings=image_rotary_embeddings, - encoder_hidden_states=encoder_hidden_states, - ) + for branch in self._get_cfg_branches(prompt_data): + pooled_embeds = ( + branch.pooled if branch.pooled is not None else branch.embeds + ) + + text_seq_len = branch.embeds.shape[1] + self._ensure_wrappers(text_seq_len, branch.mask) + self._set_text_seq_len(text_seq_len) + + if self.joint_block_wrappers: + for wrapper in self.joint_block_wrappers: + wrapper.set_encoder_mask(branch.mask) + + text_embeddings = self.adapter.compute_text_embeddings( + t, config, pooled_embeds + ) + image_rotary_embeddings = self.adapter.compute_rotary_embeddings( + branch.embeds, + config, + encoder_hidden_states_mask=branch.mask, + cond_image_grid=cond_image_grid, + kontext_image_ids=kontext_image_ids, + ) + + noise, encoder_hidden_states = self._run_single_patch_pass( + patch=patch, + patch_idx=patch_idx, + token_indices=token_indices[patch_idx], + prompt_embeds=branch.embeds, + text_embeddings=text_embeddings, + image_rotary_embeddings=image_rotary_embeddings, + encoder_hidden_states=encoder_hidden_states, + ) + + if self.is_last_stage: + assert noise is not None + results.append((branch.positive, noise)) if self.is_last_stage: - assert noise is not None - if needs_cfg: - noise_pos, noise_neg = mx.split(noise, 2, axis=0) - guidance_scale = self._get_effective_guidance_scale() - assert guidance_scale is not None - noise = self.adapter.apply_guidance( - noise_pos, noise_neg, guidance_scale - ) + noise = self._combine_cfg_results(results) patch_latents[patch_idx] = config.scheduler.step( # pyright: ignore[reportAny] noise=noise, @@ -890,10 +1057,14 @@ class DiffusionRunner: if not self.is_first_stage and t != config.num_inference_steps - 1: with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.first_pipeline_rank}", + rank=self.rank, + category="comms", ): patch_latents[patch_idx] = mx.distributed.send( - patch_latents[patch_idx], self.next_rank, group=self.group + patch_latents[patch_idx], + self.first_pipeline_rank, + group=self.group, ) mx.async_eval(patch_latents[patch_idx]) @@ -933,26 +1104,31 @@ class DiffusionRunner: if self.has_joint_blocks: if not self.is_first_stage: + assert self.prev_pipeline_rank is not None patch_len = patch.shape[1] with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): patch = mx.distributed.recv( (batch_size, patch_len, hidden_dim), patch.dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(patch) if patch_idx == 0: with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): encoder_hidden_states = mx.distributed.recv( (batch_size, text_seq_len, hidden_dim), patch.dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(encoder_hidden_states) @@ -988,39 +1164,54 @@ class DiffusionRunner: if self.has_single_blocks or self.is_last_stage: patch = patch_concat else: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): patch_concat = mx.distributed.send( - patch_concat, self.next_rank, group=self.group + patch_concat, self.next_pipeline_rank, group=self.group ) mx.async_eval(patch_concat) elif self.has_joint_blocks and not self.is_last_stage: - with trace(name=f"send {self.next_rank}", rank=self.rank, category="comms"): - patch = mx.distributed.send(patch, self.next_rank, group=self.group) + assert self.next_pipeline_rank is not None + with trace( + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", + ): + patch = mx.distributed.send( + patch, self.next_pipeline_rank, group=self.group + ) mx.async_eval(patch) if patch_idx == 0: assert encoder_hidden_states is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): encoder_hidden_states = mx.distributed.send( - encoder_hidden_states, self.next_rank, group=self.group + encoder_hidden_states, self.next_pipeline_rank, group=self.group ) mx.async_eval(encoder_hidden_states) if self.has_single_blocks: if not self.owns_concat_stage and not self.is_first_stage: + assert self.prev_pipeline_rank is not None patch_len = patch.shape[1] with trace( - name=f"recv {self.prev_rank}", rank=self.rank, category="comms" + name=f"recv {self.prev_pipeline_rank}", + rank=self.rank, + category="comms", ): patch = mx.distributed.recv( (batch_size, text_seq_len + patch_len, hidden_dim), patch.dtype, - self.prev_rank, + self.prev_pipeline_rank, group=self.group, ) mx.eval(patch) @@ -1043,15 +1234,20 @@ class DiffusionRunner: mx.eval(patch) if not self.is_last_stage: + assert self.next_pipeline_rank is not None with trace( - name=f"send {self.next_rank}", rank=self.rank, category="comms" + name=f"send {self.next_pipeline_rank}", + rank=self.rank, + category="comms", ): - patch = mx.distributed.send(patch, self.next_rank, group=self.group) + patch = mx.distributed.send( + patch, self.next_pipeline_rank, group=self.group + ) mx.async_eval(patch) noise: mx.array | None = None if self.is_last_stage: - patch = patch[:, text_seq_len:, :] - noise = self.adapter.final_projection(patch, text_embeddings) + patch_img_only = patch[:, text_seq_len:, :] + noise = self.adapter.final_projection(patch_img_only, text_embeddings) return noise, encoder_hidden_states diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index 81fb5e3c..29bc0183 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -166,6 +166,12 @@ def _inner_model(model: nn.Module) -> nn.Module: if isinstance(inner, nn.Module): return inner + inner = getattr(model, "language_model", None) + if isinstance(inner, nn.Module): + inner_inner = getattr(inner, "model", None) + if isinstance(inner_inner, nn.Module): + return inner_inner + raise ValueError("Model must either have a 'model' or 'transformer' attribute") diff --git a/src/exo/worker/engines/mlx/constants.py b/src/exo/worker/engines/mlx/constants.py index dbffdfa0..86a663e4 100644 --- a/src/exo/worker/engines/mlx/constants.py +++ b/src/exo/worker/engines/mlx/constants.py @@ -11,5 +11,7 @@ QUANTIZE_MODEL_MODE: str | None = "affine" CACHE_GROUP_SIZE: int = 64 KV_CACHE_BITS: int | None = None +DEFAULT_TOP_LOGPROBS: int = 5 + # TODO: We should really make this opt-in, but Kimi requires trust_remote_code=True TRUST_REMOTE_CODE: bool = True diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py index c9585f57..67a31ae0 100644 --- a/src/exo/worker/engines/mlx/generator/generate.py +++ b/src/exo/worker/engines/mlx/generator/generate.py @@ -12,18 +12,24 @@ from exo.shared.types.api import ( FinishReason, GenerationStats, PromptTokensDetails, + TopLogprobItem, Usage, ) from exo.shared.types.common import ModelId from exo.shared.types.memory import Memory from exo.shared.types.mlx import KVCacheType -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.runner_response import ( GenerationResponse, ) from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.cache import KVPrefixCache, encode_prompt, make_kv_cache -from exo.worker.engines.mlx.constants import KV_BITS, KV_GROUP_SIZE, MAX_TOKENS +from exo.worker.engines.mlx.constants import ( + DEFAULT_TOP_LOGPROBS, + KV_BITS, + KV_GROUP_SIZE, + MAX_TOKENS, +) from exo.worker.engines.mlx.utils_mlx import ( apply_chat_template, mx_barrier, @@ -100,7 +106,7 @@ def warmup_inference( tokenizer=tokenizer, task_params=TextGenerationTaskParams( model=ModelId(""), - input=content, + input=[InputMessage(role="user", content=content)], ), ) @@ -155,6 +161,60 @@ def eos_ids_from_tokenizer(tokenizer: TokenizerWrapper) -> list[int]: return eos +def extract_top_logprobs( + logprobs: mx.array, + tokenizer: TokenizerWrapper, + top_logprobs: int, + selected_token: int, +) -> tuple[float, list[TopLogprobItem]]: + """Extract the selected token's logprob and top alternative tokens. + + Args: + logprobs: Full vocabulary logprobs array from MLX + tokenizer: Tokenizer for decoding token IDs to strings + top_logprobs: Number of top alternatives to return + selected_token: The token ID that was actually sampled + + Returns: + Tuple of (selected_token_logprob, list of TopLogprobItem for top alternatives) + """ + # Get the logprob of the selected token + selected_logprob = float(logprobs[selected_token].item()) + + # Get top indices (most probable tokens) + # mx.argpartition gives indices that would partition the array + # We negate logprobs since argpartition finds smallest, and we want largest + top_logprobs = min(top_logprobs, logprobs.shape[0]) # Don't exceed vocab size + top_indices = mx.argpartition(-logprobs, top_logprobs)[:top_logprobs] + + # Get the actual logprob values for these indices + top_values = logprobs[top_indices] + + # Sort by logprob (descending) for consistent ordering + sort_order = mx.argsort(-top_values) + top_indices = top_indices[sort_order] + top_values = top_values[sort_order] + + # Convert to list of TopLogprobItem + top_logprob_items: list[TopLogprobItem] = [] + for i in range(top_logprobs): + token_id = int(top_indices[i].item()) + token_logprob = float(top_values[i].item()) + # Decode token ID to string + token_str = tokenizer.decode([token_id]) + # Get byte representation + token_bytes = list(token_str.encode("utf-8")) + top_logprob_items.append( + TopLogprobItem( + token=token_str, + logprob=token_logprob, + bytes=token_bytes, + ) + ) + + return selected_logprob, top_logprob_items + + def mlx_generate( model: Model, tokenizer: TokenizerWrapper, @@ -296,9 +356,22 @@ def mlx_generate( ), ) + # Extract logprobs from the full vocabulary logprobs array + logprob: float | None = None + top_logprobs: list[TopLogprobItem] | None = None + if task.logprobs: + logprob, top_logprobs = extract_top_logprobs( + logprobs=out.logprobs, + tokenizer=tokenizer, + top_logprobs=task.top_logprobs or DEFAULT_TOP_LOGPROBS, + selected_token=out.token, + ) + yield GenerationResponse( text=text, token=out.token, + logprob=logprob, + top_logprobs=top_logprobs, finish_reason=finish_reason, stats=stats, usage=usage, diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index de5cb190..4f1140fb 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -48,6 +48,7 @@ from exo.shared.types.worker.instances import ( MlxRingInstance, ) from exo.shared.types.worker.shards import ( + CfgShardMetadata, PipelineShardMetadata, ShardMetadata, TensorShardMetadata, @@ -274,6 +275,11 @@ def shard_and_load( logger.info(f"loading model from {model_path} with pipeline parallelism") model = pipeline_auto_parallel(model, group, shard_metadata) eval_with_timeout(model.parameters(), timeout_seconds, on_timeout) + case CfgShardMetadata(): + raise ValueError( + "CfgShardMetadata is not supported for text model loading - " + "this metadata type is only for image generation models" + ) # TODO: Do we need this? mx.eval(model) @@ -384,6 +390,17 @@ def load_tokenizer_for_model_id( eos_token_ids=eos_token_ids, ) + if "gemma-3" in model_id_lower: + gemma_3_eos_id = 1 + gemma_3_end_of_turn_id = 106 + if tokenizer.eos_token_ids is not None: + if gemma_3_end_of_turn_id not in tokenizer.eos_token_ids: + tokenizer.eos_token_ids = list(tokenizer.eos_token_ids) + [ + gemma_3_end_of_turn_id + ] + else: + tokenizer.eos_token_ids = [gemma_3_eos_id, gemma_3_end_of_turn_id] + return tokenizer @@ -436,16 +453,17 @@ def apply_chat_template( ) # Convert input to messages - if isinstance(task_params.input, str): - # Simple string input becomes a single user message - formatted_messages.append({"role": "user", "content": task_params.input}) - else: - # List of InputMessage - for msg in task_params.input: - if not msg.content: - logger.warning("Received message with empty content, skipping") - continue - formatted_messages.append({"role": msg.role, "content": msg.content}) + for msg in task_params.input: + if not msg.content: + logger.warning("Received message with empty content, skipping") + continue + formatted_messages.append({"role": msg.role, "content": msg.content}) + + # For assistant prefilling, append content after templating to avoid a closing turn token. + partial_assistant_content: str | None = None + if formatted_messages and formatted_messages[-1].get("role") == "assistant": + partial_assistant_content = cast(str, formatted_messages[-1].get("content", "")) + formatted_messages = formatted_messages[:-1] prompt: str = tokenizer.apply_chat_template( formatted_messages, @@ -454,6 +472,9 @@ def apply_chat_template( tools=task_params.tools, ) + if partial_assistant_content: + prompt += partial_assistant_content + logger.info(prompt) return prompt diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 5439b72b..b0e655cc 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -66,7 +66,11 @@ from exo.shared.types.worker.runners import ( RunnerStatus, RunnerWarmingUp, ) -from exo.shared.types.worker.shards import ShardMetadata +from exo.shared.types.worker.shards import ( + CfgShardMetadata, + PipelineShardMetadata, + ShardMetadata, +) from exo.utils.channels import MpReceiver, MpSender from exo.worker.engines.image import ( DistributedImageModel, @@ -87,6 +91,22 @@ from exo.worker.engines.mlx.utils_mlx import ( from exo.worker.runner.bootstrap import logger +def _is_primary_output_node(shard_metadata: ShardMetadata) -> bool: + """Check if this node is the primary output node for image generation. + + For CFG models: the last pipeline stage in CFG group 0 (positive prompt). + For non-CFG models: the last pipeline stage. + """ + if isinstance(shard_metadata, CfgShardMetadata): + is_pipeline_last = ( + shard_metadata.pipeline_rank == shard_metadata.pipeline_world_size - 1 + ) + return is_pipeline_last and shard_metadata.cfg_rank == 0 + elif isinstance(shard_metadata, PipelineShardMetadata): + return shard_metadata.device_rank == shard_metadata.world_size - 1 + return False + + def main( bound_instance: BoundInstance, event_sender: MpSender[Event], @@ -125,7 +145,6 @@ def main( event_sender.send( TaskStatusUpdated(task_id=task.task_id, task_status=TaskStatus.Running) ) - event_sender.send(TaskAcknowledged(task_id=task.task_id)) match task: case ConnectToGroup() if isinstance( current_status, (RunnerIdle, RunnerFailed) @@ -137,6 +156,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) group = initialize_mlx(bound_instance) logger.info("runner connected") @@ -153,6 +173,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) def on_model_load_timeout() -> None: event_sender.send( @@ -195,6 +216,7 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) logger.info(f"warming up inference for instance: {instance}") if ModelTask.TextGeneration in shard_metadata.model_card.tasks: @@ -234,6 +256,8 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) + assert model and not isinstance(model, DistributedImageModel) assert tokenizer @@ -320,6 +344,8 @@ def main( usage=response.usage, finish_reason=response.finish_reason, stats=response.stats, + logprob=response.logprob, + top_logprobs=response.top_logprobs, ), ) ) @@ -365,16 +391,14 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) try: - # Generate images using the image generation backend - # Track image_index for final images only image_index = 0 for response in generate_image(model=model, task=task_params): - if ( - shard_metadata.device_rank - == shard_metadata.world_size - 1 - ): + is_primary_output = _is_primary_output_node(shard_metadata) + + if is_primary_output: match response: case PartialImageResponse(): logger.info( @@ -399,7 +423,7 @@ def main( image_index += 1 # can we make this more explicit? except Exception as e: - if shard_metadata.device_rank == shard_metadata.world_size - 1: + if _is_primary_output_node(shard_metadata): event_sender.send( ChunkGenerated( command_id=command_id, @@ -430,14 +454,12 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) try: image_index = 0 for response in generate_image(model=model, task=task_params): - if ( - shard_metadata.device_rank - == shard_metadata.world_size - 1 - ): + if _is_primary_output_node(shard_metadata): match response: case PartialImageResponse(): logger.info( @@ -461,7 +483,7 @@ def main( ) image_index += 1 except Exception as e: - if shard_metadata.device_rank == shard_metadata.world_size - 1: + if _is_primary_output_node(shard_metadata): event_sender.send( ChunkGenerated( command_id=command_id, @@ -488,6 +510,8 @@ def main( runner_id=runner_id, runner_status=current_status ) ) + event_sender.send(TaskAcknowledged(task_id=task.task_id)) + current_status = RunnerShutdown() case _: raise ValueError( @@ -918,15 +942,10 @@ def _check_for_debug_prompts(task_params: TextGenerationTaskParams) -> None: Extracts the first user input text and checks for debug triggers. """ - prompt: str - if isinstance(task_params.input, str): - prompt = task_params.input - else: - # List of InputMessage - get first message content - if len(task_params.input) == 0: - logger.debug("Empty message list in debug prompt check") - return - prompt = task_params.input[0].content + if len(task_params.input) == 0: + logger.debug("Empty message list in debug prompt check") + return + prompt = task_params.input[0].content if not prompt: return diff --git a/src/exo/worker/tests/unittests/test_mlx/conftest.py b/src/exo/worker/tests/unittests/test_mlx/conftest.py index 7015d70b..9e897141 100644 --- a/src/exo/worker/tests/unittests/test_mlx/conftest.py +++ b/src/exo/worker/tests/unittests/test_mlx/conftest.py @@ -14,7 +14,7 @@ from exo.shared.constants import EXO_MODELS_DIR 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.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.shards import PipelineShardMetadata, TensorShardMetadata from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.generator.generate import mlx_generate @@ -114,7 +114,7 @@ def run_gpt_oss_pipeline_device( task = TextGenerationTaskParams( model=DEFAULT_GPT_OSS_MODEL_ID, - input=prompt_text, + input=[InputMessage(role="user", content=prompt_text)], max_output_tokens=max_tokens, ) @@ -182,7 +182,7 @@ def run_gpt_oss_tensor_parallel_device( task = TextGenerationTaskParams( model=DEFAULT_GPT_OSS_MODEL_ID, - input=prompt_text, + input=[InputMessage(role="user", content=prompt_text)], max_output_tokens=max_tokens, ) diff --git a/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py b/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py index 07376787..64be74e6 100644 --- a/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py +++ b/src/exo/worker/tests/unittests/test_plan/test_task_forwarding.py @@ -2,7 +2,7 @@ from typing import cast import exo.worker.plan as plan_mod from exo.shared.types.tasks import Task, TaskId, TaskStatus, TextGeneration -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.instances import BoundInstance, InstanceId from exo.shared.types.worker.runners import ( RunnerIdle, @@ -59,7 +59,9 @@ def test_plan_forwards_pending_chat_completion_when_runner_ready(): instance_id=INSTANCE_1_ID, task_status=TaskStatus.Pending, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) result = plan_mod.plan( @@ -106,7 +108,9 @@ def test_plan_does_not_forward_chat_completion_if_any_runner_not_ready(): instance_id=INSTANCE_1_ID, task_status=TaskStatus.Pending, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) result = plan_mod.plan( @@ -150,7 +154,9 @@ def test_plan_does_not_forward_tasks_for_other_instances(): instance_id=other_instance_id, task_status=TaskStatus.Pending, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) result = plan_mod.plan( @@ -198,7 +204,9 @@ def test_plan_ignores_non_pending_or_non_chat_tasks(): instance_id=INSTANCE_1_ID, task_status=TaskStatus.Complete, command_id=COMMAND_1_ID, - task_params=TextGenerationTaskParams(model=MODEL_A_ID, input=""), + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, input=[InputMessage(role="user", content="")] + ), ) other_task_id = TaskId("other-task") diff --git a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py index 9d7703d8..edf5ef3a 100644 --- a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py +++ b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py @@ -22,7 +22,7 @@ from exo.shared.types.tasks import ( TaskStatus, TextGeneration, ) -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.runner_response import GenerationResponse from exo.shared.types.worker.runners import ( RunnerConnected, @@ -86,7 +86,7 @@ SHUTDOWN_TASK = Shutdown( CHAT_PARAMS = TextGenerationTaskParams( model=MODEL_A_ID, - input="hello", + input=[InputMessage(role="user", content="hello")], stream=True, max_output_tokens=4, temperature=0.0, @@ -201,29 +201,29 @@ def test_events_processed_in_correct_order(patch_out_mlx: pytest.MonkeyPatch): TaskStatusUpdated( task_id=INITIALIZATION_TASK_ID, task_status=TaskStatus.Running ), - TaskAcknowledged(task_id=INITIALIZATION_TASK_ID), RunnerStatusUpdated( runner_id=RUNNER_1_ID, runner_status=RunnerConnecting() ), + TaskAcknowledged(task_id=INITIALIZATION_TASK_ID), TaskStatusUpdated( task_id=INITIALIZATION_TASK_ID, task_status=TaskStatus.Complete ), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerConnected()), TaskStatusUpdated(task_id=LOAD_TASK_ID, task_status=TaskStatus.Running), - TaskAcknowledged(task_id=LOAD_TASK_ID), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerLoading()), + TaskAcknowledged(task_id=LOAD_TASK_ID), TaskStatusUpdated(task_id=LOAD_TASK_ID, task_status=TaskStatus.Complete), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerLoaded()), TaskStatusUpdated(task_id=WARMUP_TASK_ID, task_status=TaskStatus.Running), - TaskAcknowledged(task_id=WARMUP_TASK_ID), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerWarmingUp()), + TaskAcknowledged(task_id=WARMUP_TASK_ID), TaskStatusUpdated(task_id=WARMUP_TASK_ID, task_status=TaskStatus.Complete), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerReady()), TaskStatusUpdated( task_id=CHAT_COMPLETION_TASK_ID, task_status=TaskStatus.Running ), - TaskAcknowledged(task_id=CHAT_COMPLETION_TASK_ID), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerRunning()), + TaskAcknowledged(task_id=CHAT_COMPLETION_TASK_ID), expected_chunk, TaskStatusUpdated( task_id=CHAT_COMPLETION_TASK_ID, task_status=TaskStatus.Complete @@ -231,10 +231,10 @@ def test_events_processed_in_correct_order(patch_out_mlx: pytest.MonkeyPatch): # CHAT COMPLETION TASK SHOULD COMPLETE BEFORE RUNNER READY RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerReady()), TaskStatusUpdated(task_id=SHUTDOWN_TASK_ID, task_status=TaskStatus.Running), - TaskAcknowledged(task_id=SHUTDOWN_TASK_ID), RunnerStatusUpdated( runner_id=RUNNER_1_ID, runner_status=RunnerShuttingDown() ), + TaskAcknowledged(task_id=SHUTDOWN_TASK_ID), TaskStatusUpdated( task_id=SHUTDOWN_TASK_ID, task_status=TaskStatus.Complete ), diff --git a/tests/headless_runner.py b/tests/headless_runner.py index ed57823b..56fb2632 100644 --- a/tests/headless_runner.py +++ b/tests/headless_runner.py @@ -23,7 +23,7 @@ from exo.shared.types.tasks import ( Task, TextGeneration, ) -from exo.shared.types.text_generation import TextGenerationTaskParams +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.instances import ( BoundInstance, Instance, @@ -196,7 +196,11 @@ async def execute_test(test: Tests, instance: Instance, hn: str) -> list[Event]: task_params=TextGenerationTaskParams( model=test.model_id, instructions="You are a helpful assistant", - input="What is the capital of France?", + input=[ + InputMessage( + role="user", content="What is the capital of France?" + ) + ], ), command_id=CommandId("yo"), instance_id=iid,