From cb9c9ee55cc3e628536923bed9cea80ddb69f451 Mon Sep 17 00:00:00 2001 From: ciaranbor <81697641+ciaranbor@users.noreply.github.com> Date: Fri, 23 Jan 2026 11:19:58 +0000 Subject: [PATCH 01/15] Enable generating multiple images. Optionally stream partial images (#1251) ## Motivation Support OpenAI API `n` setting ## Changes - Users can select `n` to generate more than one image with the same prompt - each image uses a different seed -> different results - `stream` and `partial_images` settings can be overwritten in UI --- .../lib/components/ImageParamsPanel.svelte | 83 ++++++++++ dashboard/src/lib/stores/app.svelte.ts | 99 +++++++++--- src/exo/master/api.py | 1 + .../shared/types/worker/runner_response.py | 2 + src/exo/worker/engines/image/generate.py | 150 ++++++++++-------- src/exo/worker/runner/runner.py | 2 +- 6 files changed, 244 insertions(+), 93 deletions(-) diff --git a/dashboard/src/lib/components/ImageParamsPanel.svelte b/dashboard/src/lib/components/ImageParamsPanel.svelte index a7d23798..ef17c278 100644 --- a/dashboard/src/lib/components/ImageParamsPanel.svelte +++ b/dashboard/src/lib/components/ImageParamsPanel.svelte @@ -110,6 +110,36 @@ setImageGenerationParams({ negativePrompt: value || null }); } + function handleNumImagesChange(event: Event) { + const input = event.target as HTMLInputElement; + const value = input.value.trim(); + if (value === "") { + setImageGenerationParams({ numImages: 1 }); + } else { + const num = parseInt(value, 10); + if (!isNaN(num) && num >= 1) { + setImageGenerationParams({ numImages: num }); + } + } + } + + function handleStreamChange(enabled: boolean) { + setImageGenerationParams({ stream: enabled }); + } + + function handlePartialImagesChange(event: Event) { + const input = event.target as HTMLInputElement; + const value = input.value.trim(); + if (value === "") { + setImageGenerationParams({ partialImages: 0 }); + } else { + const num = parseInt(value, 10); + if (!isNaN(num) && num >= 0) { + setImageGenerationParams({ partialImages: num }); + } + } + } + function clearSteps() { setImageGenerationParams({ numInferenceSteps: null }); } @@ -325,6 +355,59 @@ + + {#if !isEditMode} +
+ IMAGES: + +
+ {/if} + + +
+ STREAM: + +
+ + + {#if params.stream} +
+ PARTIALS: + +
+ {/if} + {#if isEditMode}
diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts index d1aff1a1..ebfe0cb8 100644 --- a/dashboard/src/lib/stores/app.svelte.ts +++ b/dashboard/src/lib/stores/app.svelte.ts @@ -238,6 +238,10 @@ export interface ImageGenerationParams { size: "512x512" | "768x768" | "1024x1024" | "1024x768" | "768x1024"; quality: "low" | "medium" | "high"; outputFormat: "png" | "jpeg"; + numImages: number; + // Streaming params + stream: boolean; + partialImages: number; // Advanced params seed: number | null; numInferenceSteps: number | null; @@ -257,6 +261,9 @@ const DEFAULT_IMAGE_PARAMS: ImageGenerationParams = { size: "1024x1024", quality: "medium", outputFormat: "png", + numImages: 1, + stream: true, + partialImages: 3, seed: null, numInferenceSteps: null, guidance: null, @@ -1809,12 +1816,13 @@ class AppStore { const requestBody: Record = { model, prompt, + n: params.numImages, quality: params.quality, size: params.size, output_format: params.outputFormat, response_format: "b64_json", - stream: true, - partial_images: 3, + stream: params.stream, + partial_images: params.partialImages, }; if (hasAdvancedParams) { @@ -1878,31 +1886,74 @@ class AppStore { if (imageData && idx !== -1) { const format = parsed.format || "png"; const mimeType = `image/${format}`; + const imageIndex = parsed.image_index ?? 0; + const numImages = params.numImages; + if (parsed.type === "partial") { // Update with partial image and progress const partialNum = (parsed.partial_index ?? 0) + 1; const totalPartials = parsed.total_partials ?? 3; - this.messages[idx].content = - `Generating... ${partialNum}/${totalPartials}`; - this.messages[idx].attachments = [ - { - type: "generated-image", - name: `generated-image.${format}`, - preview: `data:${mimeType};base64,${imageData}`, - mimeType, - }, - ]; + const progressText = + numImages > 1 + ? `Generating image ${imageIndex + 1}/${numImages}... ${partialNum}/${totalPartials}` + : `Generating... ${partialNum}/${totalPartials}`; + this.messages[idx].content = progressText; + + const partialAttachment: MessageAttachment = { + type: "generated-image", + name: `generated-image.${format}`, + preview: `data:${mimeType};base64,${imageData}`, + mimeType, + }; + + if (imageIndex === 0) { + // First image - safe to replace attachments with partial preview + this.messages[idx].attachments = [partialAttachment]; + } else { + // Subsequent images - keep existing finals, show partial at current position + const existingAttachments = + this.messages[idx].attachments || []; + // Keep only the completed final images (up to current imageIndex) + const finals = existingAttachments.slice(0, imageIndex); + this.messages[idx].attachments = [ + ...finals, + partialAttachment, + ]; + } } else if (parsed.type === "final") { - // Final image - this.messages[idx].content = ""; - this.messages[idx].attachments = [ - { - type: "generated-image", - name: `generated-image.${format}`, - preview: `data:${mimeType};base64,${imageData}`, - mimeType, - }, - ]; + // Final image - replace partial at this position + const newAttachment: MessageAttachment = { + type: "generated-image", + name: `generated-image-${imageIndex + 1}.${format}`, + preview: `data:${mimeType};base64,${imageData}`, + mimeType, + }; + + if (imageIndex === 0) { + // First final image - replace any partial preview + this.messages[idx].attachments = [newAttachment]; + } else { + // Subsequent images - keep previous finals, replace partial at current position + const existingAttachments = + this.messages[idx].attachments || []; + // Slice keeps indices 0 to imageIndex-1 (the previous final images) + const previousFinals = existingAttachments.slice( + 0, + imageIndex, + ); + this.messages[idx].attachments = [ + ...previousFinals, + newAttachment, + ]; + } + + // Update progress message for multiple images + if (numImages > 1 && imageIndex < numImages - 1) { + this.messages[idx].content = + `Generating image ${imageIndex + 2}/${numImages}...`; + } else { + this.messages[idx].content = ""; + } } } } catch { @@ -1983,8 +2034,8 @@ class AppStore { formData.append("size", params.size); formData.append("output_format", params.outputFormat); formData.append("response_format", "b64_json"); - formData.append("stream", "1"); // Use "1" instead of "true" for reliable FastAPI boolean parsing - formData.append("partial_images", "3"); + formData.append("stream", params.stream ? "1" : "0"); + formData.append("partial_images", params.partialImages.toString()); formData.append("input_fidelity", params.inputFidelity); // Advanced params diff --git a/src/exo/master/api.py b/src/exo/master/api.py index a436f53e..f893fedf 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -835,6 +835,7 @@ class API: # Yield partial image event (always use b64_json for partials) event_data = { "type": "partial", + "image_index": chunk.image_index, "partial_index": partial_idx, "total_partials": total_partials, "format": str(chunk.format), diff --git a/src/exo/shared/types/worker/runner_response.py b/src/exo/shared/types/worker/runner_response.py index 8d695ab0..9f867413 100644 --- a/src/exo/shared/types/worker/runner_response.py +++ b/src/exo/shared/types/worker/runner_response.py @@ -30,6 +30,7 @@ class ImageGenerationResponse(BaseRunnerResponse): image_data: bytes format: Literal["png", "jpeg", "webp"] = "png" stats: ImageGenerationStats | None = None + image_index: int = 0 def __repr_args__(self) -> Generator[tuple[str, Any], None, None]: for name, value in super().__repr_args__(): # pyright: ignore[reportAny] @@ -44,6 +45,7 @@ class PartialImageResponse(BaseRunnerResponse): format: Literal["png", "jpeg", "webp"] = "png" partial_index: int total_partials: int + image_index: int = 0 def __repr_args__(self) -> Generator[tuple[str, Any], None, None]: for name, value in super().__repr_args__(): # pyright: ignore[reportAny] diff --git a/src/exo/worker/engines/image/generate.py b/src/exo/worker/engines/image/generate.py index 8bd34749..66f21a1e 100644 --- a/src/exo/worker/engines/image/generate.py +++ b/src/exo/worker/engines/image/generate.py @@ -75,19 +75,20 @@ def generate_image( intermediate images, then ImageGenerationResponse for the final image. Yields: - PartialImageResponse for intermediate images (if partial_images > 0) - ImageGenerationResponse for the final complete image + PartialImageResponse for intermediate images (if partial_images > 0, first image only) + ImageGenerationResponse for final complete images """ width, height = parse_size(task.size) quality: Literal["low", "medium", "high"] = task.quality or "medium" advanced_params = task.advanced_params if advanced_params is not None and advanced_params.seed is not None: - seed = advanced_params.seed + base_seed = advanced_params.seed else: - seed = random.randint(0, 2**32 - 1) + base_seed = random.randint(0, 2**32 - 1) is_bench = getattr(task, "bench", False) + num_images = task.n or 1 generation_start_time: float = 0.0 @@ -95,7 +96,11 @@ def generate_image( mx.reset_peak_memory() generation_start_time = time.perf_counter() - partial_images = task.partial_images or (3 if task.stream else 0) + partial_images = ( + task.partial_images + if task.partial_images is not None + else (3 if task.stream else 0) + ) image_path: Path | None = None @@ -105,72 +110,81 @@ def generate_image( image_path = Path(tmpdir) / "input.png" image_path.write_bytes(base64.b64decode(task.image_data)) - # Iterate over generator results - for result in model.generate( - prompt=task.prompt, - height=height, - width=width, - quality=quality, - seed=seed, - image_path=image_path, - partial_images=partial_images, - advanced_params=advanced_params, - ): - if isinstance(result, tuple): - # Partial image: (Image, partial_index, total_partials) - image, partial_idx, total_partials = result - buffer = io.BytesIO() - image_format = task.output_format.upper() - if image_format == "JPG": - image_format = "JPEG" - if image_format == "JPEG" and image.mode == "RGBA": - image = image.convert("RGB") - image.save(buffer, format=image_format) + for image_num in range(num_images): + # Increment seed for each image to ensure unique results + current_seed = base_seed + image_num - yield PartialImageResponse( - image_data=buffer.getvalue(), - format=task.output_format, - partial_index=partial_idx, - total_partials=total_partials, - ) - else: - image = result + for result in model.generate( + prompt=task.prompt, + height=height, + width=width, + quality=quality, + seed=current_seed, + image_path=image_path, + partial_images=partial_images, + advanced_params=advanced_params, + ): + if isinstance(result, tuple): + # Partial image: (Image, partial_index, total_partials) + image, partial_idx, total_partials = result + buffer = io.BytesIO() + image_format = task.output_format.upper() + if image_format == "JPG": + image_format = "JPEG" + if image_format == "JPEG" and image.mode == "RGBA": + image = image.convert("RGB") + image.save(buffer, format=image_format) - stats: ImageGenerationStats | None = None - if is_bench: - generation_end_time = time.perf_counter() - total_generation_time = generation_end_time - generation_start_time - - num_inference_steps = model.get_steps_for_quality(quality) - - seconds_per_step = ( - total_generation_time / num_inference_steps - if num_inference_steps > 0 - else 0.0 + yield PartialImageResponse( + image_data=buffer.getvalue(), + format=task.output_format, + partial_index=partial_idx, + total_partials=total_partials, + image_index=image_num, ) + else: + image = result - peak_memory_gb = mx.get_peak_memory() / (1024**3) + # Only include stats on the final image + stats: ImageGenerationStats | None = None + if is_bench and image_num == num_images - 1: + generation_end_time = time.perf_counter() + total_generation_time = ( + generation_end_time - generation_start_time + ) - stats = ImageGenerationStats( - seconds_per_step=seconds_per_step, - total_generation_time=total_generation_time, - num_inference_steps=num_inference_steps, - num_images=task.n or 1, - image_width=width, - image_height=height, - peak_memory_usage=Memory.from_gb(peak_memory_gb), + num_inference_steps = model.get_steps_for_quality(quality) + total_steps = num_inference_steps * num_images + + seconds_per_step = ( + total_generation_time / total_steps + if total_steps > 0 + else 0.0 + ) + + peak_memory_gb = mx.get_peak_memory() / (1024**3) + + stats = ImageGenerationStats( + seconds_per_step=seconds_per_step, + total_generation_time=total_generation_time, + num_inference_steps=num_inference_steps, + num_images=num_images, + image_width=width, + image_height=height, + peak_memory_usage=Memory.from_gb(peak_memory_gb), + ) + + buffer = io.BytesIO() + image_format = task.output_format.upper() + if image_format == "JPG": + image_format = "JPEG" + if image_format == "JPEG" and image.mode == "RGBA": + image = image.convert("RGB") + image.save(buffer, format=image_format) + + yield ImageGenerationResponse( + image_data=buffer.getvalue(), + format=task.output_format, + stats=stats, + image_index=image_num, ) - - buffer = io.BytesIO() - image_format = task.output_format.upper() - if image_format == "JPG": - image_format = "JPEG" - if image_format == "JPEG" and image.mode == "RGBA": - image = image.convert("RGB") - image.save(buffer, format=image_format) - - yield ImageGenerationResponse( - image_data=buffer.getvalue(), - format=task.output_format, - stats=stats, - ) diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 24e42691..4f91deda 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -612,7 +612,7 @@ def _process_image_response( command_id=command_id, model_id=shard_metadata.model_card.model_id, event_sender=event_sender, - image_index=response.partial_index if is_partial else image_index, + image_index=response.image_index, is_partial=is_partial, partial_index=response.partial_index if is_partial else None, total_partials=response.total_partials if is_partial else None, From a1939c89f2d22bd7a53706f527801f938233d353 Mon Sep 17 00:00:00 2001 From: ciaranbor <81697641+ciaranbor@users.noreply.github.com> Date: Fri, 23 Jan 2026 13:37:25 +0000 Subject: [PATCH 02/15] Enable UI settings for image editing (#1258) ## Motivation Image editing was missing UI controls for quality, output format, and advanced parameters that text-to-image generation already supported. ## Changes - Added quality, output_format, and advanced_params to image edit API endpoints - Extended isImageModel check to include image editing models ## Why It Works The API now accepts and forwards these settings for image edits, and the UI displays the appropriate controls for image editing models. ## Test Plan ### Manual Testing Verified parameters can be set in UI and that they progagate through to model inference --- dashboard/src/lib/components/ChatForm.svelte | 22 ++- .../lib/components/ImageParamsPanel.svelte | 162 +++++++++--------- src/exo/master/api.py | 34 ++++ 3 files changed, 137 insertions(+), 81 deletions(-) diff --git a/dashboard/src/lib/components/ChatForm.svelte b/dashboard/src/lib/components/ChatForm.svelte index 42a9bb4c..6801287d 100644 --- a/dashboard/src/lib/components/ChatForm.svelte +++ b/dashboard/src/lib/components/ChatForm.svelte @@ -89,7 +89,10 @@ const isImageModel = $derived(() => { if (!currentModel) return false; - return modelSupportsTextToImage(currentModel); + return ( + modelSupportsTextToImage(currentModel) || + modelSupportsImageEditing(currentModel) + ); }); const isEditOnlyWithoutImage = $derived( @@ -646,6 +649,23 @@ EDIT + {:else if isEditOnlyWithoutImage} + + + + + EDIT + {:else if isImageModel()}
- -
- SIZE: -
- -
- + -
-
- - {#if isSizeDropdownOpen} - - - - -
-
- {#each sizeOptions as size} - - {/each} + {params.size} + +
+ + +
- {/if} -
+ + {#if isSizeDropdownOpen} + + + + +
+
+ {#each sizeOptions as size} + + {/each} +
+
+ {/if} +
+ {/if}
diff --git a/src/exo/master/api.py b/src/exo/master/api.py index f893fedf..989378b5 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -1,4 +1,5 @@ import base64 +import contextlib import json import time from collections.abc import AsyncGenerator @@ -33,6 +34,7 @@ from exo.shared.models.model_cards import ( ModelId, ) from exo.shared.types.api import ( + AdvancedImageParams, BenchChatCompletionResponse, BenchChatCompletionTaskParams, BenchImageGenerationResponse, @@ -1025,6 +1027,9 @@ class API: stream: bool, partial_images: int, bench: bool, + quality: Literal["high", "medium", "low"], + output_format: Literal["png", "jpeg", "webp"], + advanced_params: AdvancedImageParams | None, ) -> ImageEdits: """Prepare and send an image edits command with chunked image upload.""" resolved_model = await self._validate_image_model(model) @@ -1053,6 +1058,9 @@ class API: stream=stream, partial_images=partial_images, bench=bench, + quality=quality, + output_format=output_format, + advanced_params=advanced_params, ), ) @@ -1087,12 +1095,22 @@ class API: input_fidelity: Literal["low", "high"] = Form("low"), stream: str = Form("false"), partial_images: str = Form("0"), + quality: Literal["high", "medium", "low"] = Form("medium"), + output_format: Literal["png", "jpeg", "webp"] = Form("png"), + advanced_params: str | None = Form(None), ) -> ImageGenerationResponse | StreamingResponse: """Handle image editing requests (img2img).""" # Parse string form values to proper types stream_bool = stream.lower() in ("true", "1", "yes") partial_images_int = int(partial_images) if partial_images.isdigit() else 0 + parsed_advanced_params: AdvancedImageParams | None = None + if advanced_params: + with contextlib.suppress(Exception): + parsed_advanced_params = AdvancedImageParams.model_validate_json( + advanced_params + ) + command = await self._send_image_edits_command( image=image, prompt=prompt, @@ -1104,6 +1122,9 @@ class API: stream=stream_bool, partial_images=partial_images_int, bench=False, + quality=quality, + output_format=output_format, + advanced_params=parsed_advanced_params, ) if stream_bool and partial_images_int > 0: @@ -1134,8 +1155,18 @@ class API: size: str = Form("1024x1024"), response_format: Literal["url", "b64_json"] = Form("b64_json"), input_fidelity: Literal["low", "high"] = Form("low"), + quality: Literal["high", "medium", "low"] = Form("medium"), + output_format: Literal["png", "jpeg", "webp"] = Form("png"), + advanced_params: str | None = Form(None), ) -> BenchImageGenerationResponse: """Handle benchmark image editing requests with generation stats.""" + parsed_advanced_params: AdvancedImageParams | None = None + if advanced_params: + with contextlib.suppress(Exception): + parsed_advanced_params = AdvancedImageParams.model_validate_json( + advanced_params + ) + command = await self._send_image_edits_command( image=image, prompt=prompt, @@ -1147,6 +1178,9 @@ class API: stream=False, partial_images=0, bench=True, + quality=quality, + output_format=output_format, + advanced_params=parsed_advanced_params, ) return await self._collect_image_generation_with_stats( From f255345a1ad29e923b0f22d0f5d13c815d1f697d Mon Sep 17 00:00:00 2001 From: Jake Hillion Date: Fri, 23 Jan 2026 12:47:43 +0000 Subject: [PATCH 03/15] dashboard: decouple prettier-svelte from dashboard source The prettier-svelte formatter depended on the full dashboard build (dashboardFull), causing the devshell to rebuild whenever any dashboard source file changed. Created a deps-only dream2nix derivation (deps.nix) that uses a stub source containing only package.json, package-lock.json, and minimal files for vite to succeed. Updated prettier-svelte to use this derivation instead of dashboardFull. The stub source is constant unless lockfiles change, so prettier-svelte and the devshell no longer rebuild when dashboard source files are modified. Test plan: - nix flake check passed - nix fmt successfully formatted svelte files --- dashboard/parts.nix | 46 ++++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 43 insertions(+), 3 deletions(-) diff --git a/dashboard/parts.nix b/dashboard/parts.nix index 487078d5..8edcc6b7 100644 --- a/dashboard/parts.nix +++ b/dashboard/parts.nix @@ -3,6 +3,45 @@ perSystem = { pkgs, lib, ... }: let + # Stub source with lockfiles and minimal files for build to succeed + # This allows prettier-svelte to avoid rebuilding when dashboard source changes + dashboardStubSrc = pkgs.runCommand "dashboard-stub-src" { } '' + mkdir -p $out + cp ${inputs.self}/dashboard/package.json $out/ + cp ${inputs.self}/dashboard/package-lock.json $out/ + # Minimal files so vite build succeeds (produces empty output) + echo '' > $out/index.html + mkdir -p $out/src + touch $out/src/app.html + ''; + + # Deps-only build using stub source (for prettier-svelte) + # Only rebuilds when package.json or package-lock.json change + dashboardDeps = inputs.dream2nix.lib.evalModules { + packageSets.nixpkgs = pkgs; + modules = [ + ./dashboard.nix + { + paths.projectRoot = inputs.self; + paths.projectRootFile = "flake.nix"; + paths.package = inputs.self + "/dashboard"; + } + { + deps.dashboardSrc = lib.mkForce dashboardStubSrc; + } + # Override build phases to skip the actual build - just need node_modules + { + mkDerivation = { + buildPhase = lib.mkForce "true"; + installPhase = lib.mkForce '' + runHook preInstall + runHook postInstall + ''; + }; + } + ]; + }; + # Filter source to only include dashboard directory dashboardSrc = lib.cleanSourceWith { src = inputs.self; @@ -42,11 +81,12 @@ ''; # Prettier with svelte plugin for treefmt + # Uses dashboardDeps instead of dashboardFull to avoid rebuilding on source changes packages.prettier-svelte = pkgs.writeShellScriptBin "prettier-svelte" '' - export NODE_PATH="${dashboardFull}/lib/node_modules/exo-dashboard/node_modules" + export NODE_PATH="${dashboardDeps}/lib/node_modules/exo-dashboard/node_modules" exec ${pkgs.nodejs}/bin/node \ - ${dashboardFull}/lib/node_modules/exo-dashboard/node_modules/prettier/bin/prettier.cjs \ - --plugin "${dashboardFull}/lib/node_modules/exo-dashboard/node_modules/prettier-plugin-svelte/plugin.js" \ + ${dashboardDeps}/lib/node_modules/exo-dashboard/node_modules/prettier/bin/prettier.cjs \ + --plugin "${dashboardDeps}/lib/node_modules/exo-dashboard/node_modules/prettier-plugin-svelte/plugin.js" \ "$@" ''; }; From ba19940828e10f171fa5300441aab94b074802de Mon Sep 17 00:00:00 2001 From: ciaranbor <81697641+ciaranbor@users.noreply.github.com> Date: Fri, 23 Jan 2026 16:33:01 +0000 Subject: [PATCH 04/15] Fix regenerate for image models (#1263) ## Motivation The 'regenerate' button was hardcoded to chat completion. Clicking 'regenerate' for image request would result in an error after the model is loaded ## Changes Store request type and dispatch to appropriate request upon regeneration ## Why It Works We make sure to repeat the same request type as was performed originally ## Test Plan ### Manual Testing Checked 'regenerate' works for chat completion, image generation, image editing --- dashboard/src/lib/stores/app.svelte.ts | 47 ++++++++++++++++++++++++-- 1 file changed, 44 insertions(+), 3 deletions(-) diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts index ebfe0cb8..1b1a3847 100644 --- a/dashboard/src/lib/stores/app.svelte.ts +++ b/dashboard/src/lib/stores/app.svelte.ts @@ -216,6 +216,8 @@ export interface Message { attachments?: MessageAttachment[]; ttftMs?: number; // Time to first token in ms (for assistant messages) tps?: number; // Tokens per second (for assistant messages) + requestType?: "chat" | "image-generation" | "image-editing"; + sourceImageDataUrl?: string; // For image editing regeneration } export interface Conversation { @@ -1270,10 +1272,46 @@ class AppStore { if (lastUserIndex === -1) return; - // Remove any messages after the user message - this.messages = this.messages.slice(0, lastUserIndex + 1); + const lastUserMessage = this.messages[lastUserIndex]; + const requestType = lastUserMessage.requestType || "chat"; + const prompt = lastUserMessage.content; - // Resend the message to get a new response + // Remove messages after user message (including the user message for image requests + // since generateImage/editImage will re-add it) + this.messages = this.messages.slice(0, lastUserIndex); + + switch (requestType) { + case "image-generation": + await this.generateImage(prompt); + break; + case "image-editing": + if (lastUserMessage.sourceImageDataUrl) { + await this.editImage(prompt, lastUserMessage.sourceImageDataUrl); + } else { + // Can't regenerate edit without source image - restore user message and show error + this.messages.push(lastUserMessage); + const errorMessage = this.addMessage("assistant", ""); + const idx = this.messages.findIndex((m) => m.id === errorMessage.id); + if (idx !== -1) { + this.messages[idx].content = + "Error: Cannot regenerate image edit - source image not found"; + } + this.updateActiveConversation(); + } + break; + case "chat": + default: + // Restore the user message for chat regeneration + this.messages.push(lastUserMessage); + await this.regenerateChatCompletion(); + break; + } + } + + /** + * Helper method to regenerate a chat completion response + */ + private async regenerateChatCompletion(): Promise { this.isLoading = true; this.currentResponse = ""; @@ -1788,6 +1826,7 @@ class AppStore { role: "user", content: prompt, timestamp: Date.now(), + requestType: "image-generation", }; this.messages.push(userMessage); @@ -1998,6 +2037,8 @@ class AppStore { role: "user", content: prompt, timestamp: Date.now(), + requestType: "image-editing", + sourceImageDataUrl: imageDataUrl, }; this.messages.push(userMessage); From 9357503c6fcfb23bc2ca3cf57c945a5796a3319d Mon Sep 17 00:00:00 2001 From: Jake Hillion Date: Mon, 19 Jan 2026 18:59:32 +0000 Subject: [PATCH 05/15] downloads: refactor to run at node level The Worker previously owned the ShardDownloader directly via dependency injection, which prevented --no-worker nodes from downloading and made it impossible for multiple Workers to share a single downloader instance. Moved download functionality to a new DownloadCoordinator component at the Node level that communicates via the DOWNLOAD_COMMANDS pub/sub topic. Workers now send StartDownload commands instead of calling the downloader directly, and receive progress updates through the event-sourced state. This decouples downloads from the Worker lifecycle and enables future features like UI-triggered downloads to specific nodes and multi-worker download sharing. Test plan: - Mostly tested in the next PR that adds explicit downloads/deletions to the dashboard. - Started a model that isn't downloaded - it works. --- src/exo/download/coordinator.py | 284 ++++++++++++++++++ .../{worker => }/download/download_utils.py | 14 +- .../download/huggingface_utils.py | 0 .../download/impl_shard_downloader.py | 4 +- .../{worker => }/download/shard_downloader.py | 2 +- src/exo/main.py | 69 ++++- src/exo/routing/topics.py | 5 +- src/exo/shared/models/model_cards.py | 6 +- src/exo/shared/types/commands.py | 22 +- src/exo/utils/keyed_backoff.py | 32 ++ .../worker/engines/image/distributed_model.py | 2 +- src/exo/worker/engines/mlx/utils_mlx.py | 2 +- src/exo/worker/main.py | 222 +++----------- src/exo/worker/plan.py | 14 +- .../unittests/test_mlx/test_tokenizers.py | 4 +- .../test_plan/test_download_and_loading.py | 39 +-- .../test_plan/test_runner_lifecycle.py | 11 +- .../test_plan/test_task_forwarding.py | 5 - .../tests/unittests/test_plan/test_warmup.py | 8 - tests/headless_runner.py | 8 +- 20 files changed, 479 insertions(+), 274 deletions(-) create mode 100644 src/exo/download/coordinator.py rename src/exo/{worker => }/download/download_utils.py (99%) rename src/exo/{worker => }/download/huggingface_utils.py (100%) rename src/exo/{worker => }/download/impl_shard_downloader.py (97%) rename src/exo/{worker => }/download/shard_downloader.py (97%) create mode 100644 src/exo/utils/keyed_backoff.py diff --git a/src/exo/download/coordinator.py b/src/exo/download/coordinator.py new file mode 100644 index 00000000..c2f7b9e9 --- /dev/null +++ b/src/exo/download/coordinator.py @@ -0,0 +1,284 @@ +import asyncio +from dataclasses import dataclass, field +from typing import Iterator + +import anyio +from anyio import current_time +from anyio.abc import TaskGroup +from loguru import logger + +from exo.download.download_utils import ( + RepoDownloadProgress, + delete_model, + map_repo_download_progress_to_download_progress_data, +) +from exo.download.shard_downloader import ShardDownloader +from exo.shared.models.model_cards import ModelId +from exo.shared.types.commands import ( + DeleteDownload, + ForwarderDownloadCommand, + StartDownload, +) +from exo.shared.types.common import NodeId, SessionId +from exo.shared.types.events import ( + Event, + ForwarderEvent, + NodeDownloadProgress, +) +from exo.shared.types.worker.downloads import ( + DownloadCompleted, + DownloadFailed, + DownloadOngoing, + DownloadPending, + DownloadProgress, +) +from exo.shared.types.worker.shards import ShardMetadata +from exo.utils.channels import Receiver, Sender, channel + + +@dataclass +class DownloadCoordinator: + node_id: NodeId + session_id: SessionId + shard_downloader: ShardDownloader + download_command_receiver: Receiver[ForwarderDownloadCommand] + local_event_sender: Sender[ForwarderEvent] + event_index_counter: Iterator[int] + + # Local state + download_status: dict[ModelId, DownloadProgress] = field(default_factory=dict) + active_downloads: dict[ModelId, asyncio.Task[None]] = field(default_factory=dict) + + # Internal event channel for forwarding (initialized in __post_init__) + event_sender: Sender[Event] = field(init=False) + event_receiver: Receiver[Event] = field(init=False) + _tg: TaskGroup = field(init=False) + + def __post_init__(self) -> None: + self.event_sender, self.event_receiver = channel[Event]() + self._tg = anyio.create_task_group() + + async def run(self) -> None: + logger.info("Starting DownloadCoordinator") + 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) + + def shutdown(self) -> None: + self._tg.cancel_scope.cancel() + + async def _command_processor(self) -> None: + with self.download_command_receiver as commands: + async for cmd in commands: + # Only process commands targeting this node + if cmd.command.target_node_id != self.node_id: + continue + + match cmd.command: + case StartDownload(shard_metadata=shard): + await self._start_download(shard) + case DeleteDownload(model_id=model_id): + await self._delete_download(model_id) + + async def _start_download(self, shard: ShardMetadata) -> None: + model_id = shard.model_card.model_id + + # Check if already downloading or complete + if model_id in self.download_status: + status = self.download_status[model_id] + if isinstance(status, (DownloadOngoing, DownloadCompleted)): + logger.debug( + f"Download for {model_id} already in progress or complete, skipping" + ) + return + + # Emit pending status + progress = DownloadPending(shard_metadata=shard, node_id=self.node_id) + self.download_status[model_id] = progress + await self.event_sender.send(NodeDownloadProgress(download_progress=progress)) + + # Check initial status from downloader + initial_progress = ( + await self.shard_downloader.get_shard_download_status_for_shard(shard) + ) + + if initial_progress.status == "complete": + completed = DownloadCompleted( + shard_metadata=shard, + node_id=self.node_id, + total_bytes=initial_progress.total_bytes, + ) + self.download_status[model_id] = completed + await self.event_sender.send( + NodeDownloadProgress(download_progress=completed) + ) + return + + # Start actual download + self._start_download_task(shard, initial_progress) + + def _start_download_task( + self, shard: ShardMetadata, initial_progress: RepoDownloadProgress + ) -> None: + model_id = shard.model_card.model_id + + # Emit ongoing status + status = DownloadOngoing( + node_id=self.node_id, + shard_metadata=shard, + download_progress=map_repo_download_progress_to_download_progress_data( + initial_progress + ), + ) + self.download_status[model_id] = status + self.event_sender.send_nowait(NodeDownloadProgress(download_progress=status)) + + last_progress_time = 0.0 + throttle_interval_secs = 1.0 + + async def download_progress_callback( + callback_shard: ShardMetadata, progress: RepoDownloadProgress + ) -> None: + nonlocal last_progress_time + + if progress.status == "complete": + completed = DownloadCompleted( + shard_metadata=callback_shard, + node_id=self.node_id, + total_bytes=progress.total_bytes, + ) + self.download_status[callback_shard.model_card.model_id] = completed + await self.event_sender.send( + NodeDownloadProgress(download_progress=completed) + ) + # Clean up active download tracking + if callback_shard.model_card.model_id in self.active_downloads: + del self.active_downloads[callback_shard.model_card.model_id] + elif ( + progress.status == "in_progress" + and current_time() - last_progress_time > throttle_interval_secs + ): + ongoing = DownloadOngoing( + node_id=self.node_id, + shard_metadata=callback_shard, + download_progress=map_repo_download_progress_to_download_progress_data( + progress + ), + ) + self.download_status[callback_shard.model_card.model_id] = ongoing + await self.event_sender.send( + NodeDownloadProgress(download_progress=ongoing) + ) + last_progress_time = current_time() + + self.shard_downloader.on_progress(download_progress_callback) + + async def download_wrapper() -> None: + try: + await self.shard_downloader.ensure_shard(shard) + except Exception as e: + logger.error(f"Download failed for {model_id}: {e}") + failed = DownloadFailed( + shard_metadata=shard, + node_id=self.node_id, + error_message=str(e), + ) + self.download_status[model_id] = failed + await self.event_sender.send( + NodeDownloadProgress(download_progress=failed) + ) + finally: + if model_id in self.active_downloads: + del self.active_downloads[model_id] + + task = asyncio.create_task(download_wrapper()) + self.active_downloads[model_id] = task + + async def _delete_download(self, model_id: ModelId) -> None: + # Cancel if active + if model_id in self.active_downloads: + logger.info(f"Cancelling active download for {model_id} before deletion") + self.active_downloads[model_id].cancel() + del self.active_downloads[model_id] + + # Delete from disk + logger.info(f"Deleting model files for {model_id}") + deleted = await delete_model(model_id) + + if deleted: + logger.info(f"Successfully deleted model {model_id}") + else: + logger.warning(f"Model {model_id} was not found on disk") + + # Emit pending status to reset UI state, then remove from local tracking + if model_id in self.download_status: + current_status = self.download_status[model_id] + pending = DownloadPending( + shard_metadata=current_status.shard_metadata, + node_id=self.node_id, + ) + await self.event_sender.send( + NodeDownloadProgress(download_progress=pending) + ) + del self.download_status[model_id] + + async def _forward_events(self) -> None: + with self.event_receiver as events: + async for event in events: + idx = next(self.event_index_counter) + fe = ForwarderEvent( + origin_idx=idx, + origin=self.node_id, + session=self.session_id, + event=event, + ) + logger.debug( + f"DownloadCoordinator published event {idx}: {str(event)[:100]}" + ) + await self.local_event_sender.send(fe) + + async def _emit_existing_download_progress(self) -> None: + try: + while True: + logger.info( + "DownloadCoordinator: Fetching and emitting existing download progress..." + ) + async for ( + _, + progress, + ) in self.shard_downloader.get_shard_download_status(): + if progress.status == "complete": + status: DownloadProgress = DownloadCompleted( + node_id=self.node_id, + shard_metadata=progress.shard, + total_bytes=progress.total_bytes, + ) + elif progress.status in ["in_progress", "not_started"]: + if progress.downloaded_bytes_this_session.in_bytes == 0: + status = DownloadPending( + node_id=self.node_id, shard_metadata=progress.shard + ) + else: + status = DownloadOngoing( + node_id=self.node_id, + shard_metadata=progress.shard, + download_progress=map_repo_download_progress_to_download_progress_data( + progress + ), + ) + else: + continue + + self.download_status[progress.shard.model_card.model_id] = status + await self.event_sender.send( + NodeDownloadProgress(download_progress=status) + ) + logger.info( + "DownloadCoordinator: Done emitting existing download progress." + ) + await anyio.sleep(5 * 60) # 5 minutes + except Exception as e: + logger.error( + f"DownloadCoordinator: Error emitting existing download progress: {e}" + ) diff --git a/src/exo/worker/download/download_utils.py b/src/exo/download/download_utils.py similarity index 99% rename from src/exo/worker/download/download_utils.py rename to src/exo/download/download_utils.py index 8816d701..f08aa0ed 100644 --- a/src/exo/worker/download/download_utils.py +++ b/src/exo/download/download_utils.py @@ -24,6 +24,13 @@ from pydantic import ( TypeAdapter, ) +from exo.download.huggingface_utils import ( + filter_repo_objects, + get_allow_patterns, + get_auth_headers, + get_hf_endpoint, + get_hf_token, +) from exo.shared.constants import EXO_MODELS_DIR from exo.shared.types.common import ModelId from exo.shared.types.memory import Memory @@ -35,13 +42,6 @@ from exo.shared.types.worker.downloads import ( RepoFileDownloadProgress, ) from exo.shared.types.worker.shards import ShardMetadata -from exo.worker.download.huggingface_utils import ( - filter_repo_objects, - get_allow_patterns, - get_auth_headers, - get_hf_endpoint, - get_hf_token, -) class HuggingFaceAuthenticationError(Exception): diff --git a/src/exo/worker/download/huggingface_utils.py b/src/exo/download/huggingface_utils.py similarity index 100% rename from src/exo/worker/download/huggingface_utils.py rename to src/exo/download/huggingface_utils.py diff --git a/src/exo/worker/download/impl_shard_downloader.py b/src/exo/download/impl_shard_downloader.py similarity index 97% rename from src/exo/worker/download/impl_shard_downloader.py rename to src/exo/download/impl_shard_downloader.py index b253ad9d..24f13afe 100644 --- a/src/exo/worker/download/impl_shard_downloader.py +++ b/src/exo/download/impl_shard_downloader.py @@ -5,13 +5,13 @@ from typing import AsyncIterator, Callable from loguru import logger +from exo.download.download_utils import RepoDownloadProgress, download_shard +from exo.download.shard_downloader import ShardDownloader from exo.shared.models.model_cards import MODEL_CARDS, ModelCard, ModelId from exo.shared.types.worker.shards import ( PipelineShardMetadata, ShardMetadata, ) -from exo.worker.download.download_utils import RepoDownloadProgress, download_shard -from exo.worker.download.shard_downloader import ShardDownloader def exo_shard_downloader(max_parallel_downloads: int = 8) -> ShardDownloader: diff --git a/src/exo/worker/download/shard_downloader.py b/src/exo/download/shard_downloader.py similarity index 97% rename from src/exo/worker/download/shard_downloader.py rename to src/exo/download/shard_downloader.py index deb0aaee..30c11d25 100644 --- a/src/exo/worker/download/shard_downloader.py +++ b/src/exo/download/shard_downloader.py @@ -5,13 +5,13 @@ from datetime import timedelta from pathlib import Path from typing import AsyncIterator, Callable +from exo.download.download_utils import RepoDownloadProgress from exo.shared.models.model_cards import ModelCard, ModelId, ModelTask from exo.shared.types.memory import Memory from exo.shared.types.worker.shards import ( PipelineShardMetadata, ShardMetadata, ) -from exo.worker.download.download_utils import RepoDownloadProgress # TODO: the PipelineShardMetadata getting reinstantiated is a bit messy. Should this be a classmethod? diff --git a/src/exo/main.py b/src/exo/main.py index 6425595d..8223383e 100644 --- a/src/exo/main.py +++ b/src/exo/main.py @@ -1,10 +1,11 @@ import argparse +import itertools import multiprocessing as mp import os import resource import signal from dataclasses import dataclass, field -from typing import Self +from typing import Iterator, Self import anyio from anyio.abc import TaskGroup @@ -12,6 +13,8 @@ from loguru import logger from pydantic import PositiveInt import exo.routing.topics as topics +from exo.download.coordinator import DownloadCoordinator +from exo.download.impl_shard_downloader import exo_shard_downloader from exo.master.api import API # TODO: should API be in master? from exo.master.main import Master from exo.routing.router import Router, get_node_id_keypair @@ -21,7 +24,6 @@ from exo.shared.logging import logger_cleanup, logger_setup from exo.shared.types.common import NodeId, SessionId from exo.utils.channels import Receiver, channel from exo.utils.pydantic_ext import CamelCaseModel -from exo.worker.download.impl_shard_downloader import exo_shard_downloader from exo.worker.main import Worker @@ -29,6 +31,7 @@ from exo.worker.main import Worker @dataclass class Node: router: Router + download_coordinator: DownloadCoordinator | None worker: Worker | None election: Election # Every node participates in election, as we do want a node to become master even if it isn't a master candidate if no master candidates are present. election_result_receiver: Receiver[ElectionResult] @@ -36,6 +39,7 @@ class Node: api: API | None node_id: NodeId + event_index_counter: Iterator[int] _tg: TaskGroup = field(init=False, default_factory=anyio.create_task_group) @classmethod @@ -49,8 +53,26 @@ class Node: await router.register_topic(topics.COMMANDS) await router.register_topic(topics.ELECTION_MESSAGES) await router.register_topic(topics.CONNECTION_MESSAGES) + await router.register_topic(topics.DOWNLOAD_COMMANDS) logger.info(f"Starting node {node_id}") + + # Create shared event index counter for Worker and DownloadCoordinator + event_index_counter = itertools.count() + + # Create DownloadCoordinator (unless --no-downloads) + if not args.no_downloads: + download_coordinator = DownloadCoordinator( + node_id, + session_id, + exo_shard_downloader(), + download_command_receiver=router.receiver(topics.DOWNLOAD_COMMANDS), + local_event_sender=router.sender(topics.LOCAL_EVENTS), + event_index_counter=event_index_counter, + ) + else: + download_coordinator = None + if args.spawn_api: api = API( node_id, @@ -67,11 +89,12 @@ class Node: worker = Worker( node_id, session_id, - exo_shard_downloader(), connection_message_receiver=router.receiver(topics.CONNECTION_MESSAGES), global_event_receiver=router.receiver(topics.GLOBAL_EVENTS), local_event_sender=router.sender(topics.LOCAL_EVENTS), command_sender=router.sender(topics.COMMANDS), + download_command_sender=router.sender(topics.DOWNLOAD_COMMANDS), + event_index_counter=event_index_counter, ) else: worker = None @@ -99,13 +122,25 @@ class Node: election_result_sender=er_send, ) - return cls(router, worker, election, er_recv, master, api, node_id) + return cls( + router, + download_coordinator, + worker, + election, + er_recv, + master, + api, + node_id, + event_index_counter, + ) async def run(self): async with self._tg as tg: signal.signal(signal.SIGINT, lambda _, __: self.shutdown()) tg.start_soon(self.router.run) tg.start_soon(self.election.run) + if self.download_coordinator: + tg.start_soon(self.download_coordinator.run) if self.worker: tg.start_soon(self.worker.run) if self.master: @@ -170,13 +205,27 @@ class Node: ) if result.is_new_master: await anyio.sleep(0) + # Fresh counter for new session (buffer expects indices from 0) + self.event_index_counter = itertools.count() + if self.download_coordinator: + self.download_coordinator.shutdown() + self.download_coordinator = DownloadCoordinator( + self.node_id, + result.session_id, + exo_shard_downloader(), + download_command_receiver=self.router.receiver( + topics.DOWNLOAD_COMMANDS + ), + local_event_sender=self.router.sender(topics.LOCAL_EVENTS), + event_index_counter=self.event_index_counter, + ) + self._tg.start_soon(self.download_coordinator.run) if self.worker: self.worker.shutdown() # TODO: add profiling etc to resource monitor self.worker = Worker( self.node_id, result.session_id, - exo_shard_downloader(), connection_message_receiver=self.router.receiver( topics.CONNECTION_MESSAGES ), @@ -185,6 +234,10 @@ class Node: ), local_event_sender=self.router.sender(topics.LOCAL_EVENTS), command_sender=self.router.sender(topics.COMMANDS), + download_command_sender=self.router.sender( + topics.DOWNLOAD_COMMANDS + ), + event_index_counter=self.event_index_counter, ) self._tg.start_soon(self.worker.run) if self.api: @@ -226,6 +279,7 @@ class Args(CamelCaseModel): api_port: PositiveInt = 52415 tb_only: bool = False no_worker: bool = False + no_downloads: bool = False fast_synch: bool | None = None # None = auto, True = force on, False = force off @classmethod @@ -268,6 +322,11 @@ class Args(CamelCaseModel): "--no-worker", action="store_true", ) + parser.add_argument( + "--no-downloads", + action="store_true", + help="Disable the download coordinator (node won't download models)", + ) fast_synch_group = parser.add_mutually_exclusive_group() fast_synch_group.add_argument( "--fast-synch", diff --git a/src/exo/routing/topics.py b/src/exo/routing/topics.py index 50f1c9af..5a122a95 100644 --- a/src/exo/routing/topics.py +++ b/src/exo/routing/topics.py @@ -3,7 +3,7 @@ from enum import Enum from exo.routing.connection_message import ConnectionMessage from exo.shared.election import ElectionMessage -from exo.shared.types.commands import ForwarderCommand +from exo.shared.types.commands import ForwarderCommand, ForwarderDownloadCommand from exo.shared.types.events import ( ForwarderEvent, ) @@ -45,3 +45,6 @@ ELECTION_MESSAGES = TypedTopic( CONNECTION_MESSAGES = TypedTopic( "connection_messages", PublishPolicy.Never, ConnectionMessage ) +DOWNLOAD_COMMANDS = TypedTopic( + "download_commands", PublishPolicy.Always, ForwarderDownloadCommand +) diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index eaad0154..35e077ba 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -621,7 +621,7 @@ class ConfigData(BaseModel): async def get_config_data(model_id: ModelId) -> ConfigData: """Downloads and parses config.json for a model.""" - from exo.worker.download.download_utils import ( + from exo.download.download_utils import ( download_file_with_retry, ensure_models_dir, ) @@ -643,11 +643,11 @@ async def get_config_data(model_id: ModelId) -> ConfigData: async def get_safetensors_size(model_id: ModelId) -> Memory: """Gets model size from safetensors index or falls back to HF API.""" - from exo.shared.types.worker.downloads import ModelSafetensorsIndex - from exo.worker.download.download_utils import ( + from exo.download.download_utils import ( download_file_with_retry, ensure_models_dir, ) + from exo.shared.types.worker.downloads import ModelSafetensorsIndex target_dir = (await ensure_models_dir()) / model_id.normalize() await aios.makedirs(target_dir, exist_ok=True) diff --git a/src/exo/shared/types/commands.py b/src/exo/shared/types/commands.py index a0ddb968..ed086f35 100644 --- a/src/exo/shared/types/commands.py +++ b/src/exo/shared/types/commands.py @@ -1,6 +1,6 @@ from pydantic import Field -from exo.shared.models.model_cards import ModelCard +from exo.shared.models.model_cards import ModelCard, ModelId from exo.shared.types.api import ( ChatCompletionTaskParams, ImageEditsInternalParams, @@ -9,7 +9,7 @@ from exo.shared.types.api import ( from exo.shared.types.chunks import InputImageChunk from exo.shared.types.common import CommandId, NodeId from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta -from exo.shared.types.worker.shards import Sharding +from exo.shared.types.worker.shards import Sharding, ShardMetadata from exo.utils.pydantic_ext import CamelCaseModel, TaggedModel @@ -62,6 +62,19 @@ class RequestEventLog(BaseCommand): since_idx: int +class StartDownload(BaseCommand): + target_node_id: NodeId + shard_metadata: ShardMetadata + + +class DeleteDownload(BaseCommand): + target_node_id: NodeId + model_id: ModelId + + +DownloadCommand = StartDownload | DeleteDownload + + Command = ( TestCommand | RequestEventLog @@ -79,3 +92,8 @@ Command = ( class ForwarderCommand(CamelCaseModel): origin: NodeId command: Command + + +class ForwarderDownloadCommand(CamelCaseModel): + origin: NodeId + command: DownloadCommand diff --git a/src/exo/utils/keyed_backoff.py b/src/exo/utils/keyed_backoff.py new file mode 100644 index 00000000..b481d43e --- /dev/null +++ b/src/exo/utils/keyed_backoff.py @@ -0,0 +1,32 @@ +import time +from typing import Generic, TypeVar + +K = TypeVar("K") + + +class KeyedBackoff(Generic[K]): + """Tracks exponential backoff state per key.""" + + def __init__(self, base: float = 0.5, cap: float = 10.0): + self._base = base + self._cap = cap + self._attempts: dict[K, int] = {} + self._last_time: dict[K, float] = {} + + def should_proceed(self, key: K) -> bool: + """Returns True if enough time has elapsed since last attempt.""" + now = time.monotonic() + last = self._last_time.get(key, 0.0) + attempts = self._attempts.get(key, 0) + delay = min(self._cap, self._base * (2.0**attempts)) + return now - last >= delay + + def record_attempt(self, key: K) -> None: + """Record that an attempt was made for this key.""" + self._last_time[key] = time.monotonic() + self._attempts[key] = self._attempts.get(key, 0) + 1 + + def reset(self, key: K) -> None: + """Reset backoff state for a key (e.g., on success).""" + self._attempts.pop(key, None) + self._last_time.pop(key, None) diff --git a/src/exo/worker/engines/image/distributed_model.py b/src/exo/worker/engines/image/distributed_model.py index 9c127f36..45fd0e29 100644 --- a/src/exo/worker/engines/image/distributed_model.py +++ b/src/exo/worker/engines/image/distributed_model.py @@ -6,10 +6,10 @@ import mlx.core as mx from mflux.models.common.config.config import Config 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.worker.download.download_utils import build_model_path from exo.worker.engines.image.config import ImageModelConfig from exo.worker.engines.image.models import ( create_adapter_for_model, diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index 0703d6b5..cb0e72b3 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -41,6 +41,7 @@ import mlx.nn as nn from mlx_lm.utils import load_model from pydantic import RootModel +from exo.download.download_utils import build_model_path from exo.shared.types.api import ChatCompletionMessageText from exo.shared.types.common import Host from exo.shared.types.memory import Memory @@ -55,7 +56,6 @@ from exo.shared.types.worker.shards import ( ShardMetadata, TensorShardMetadata, ) -from exo.worker.download.download_utils import build_model_path from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.auto_parallel import ( TimeoutCallback, diff --git a/src/exo/worker/main.py b/src/exo/worker/main.py index e861884a..3473be41 100644 --- a/src/exo/worker/main.py +++ b/src/exo/worker/main.py @@ -1,8 +1,9 @@ from datetime import datetime, timezone from random import random +from typing import Iterator import anyio -from anyio import CancelScope, create_task_group, current_time, fail_after +from anyio import CancelScope, create_task_group, fail_after from anyio.abc import TaskGroup from loguru import logger @@ -10,7 +11,12 @@ from exo.routing.connection_message import ConnectionMessage, ConnectionMessageT from exo.shared.apply import apply from exo.shared.models.model_cards import ModelId from exo.shared.types.api import ImageEditsInternalParams -from exo.shared.types.commands import ForwarderCommand, RequestEventLog +from exo.shared.types.commands import ( + ForwarderCommand, + ForwarderDownloadCommand, + RequestEventLog, + StartDownload, +) from exo.shared.types.common import CommandId, NodeId, SessionId from exo.shared.types.events import ( Event, @@ -18,7 +24,6 @@ from exo.shared.types.events import ( ForwarderEvent, IndexedEvent, InputChunkReceived, - NodeDownloadProgress, NodeGatheredInfo, TaskCreated, TaskStatusUpdated, @@ -36,23 +41,12 @@ from exo.shared.types.tasks import ( TaskStatus, ) from exo.shared.types.topology import Connection, SocketConnection -from exo.shared.types.worker.downloads import ( - DownloadCompleted, - DownloadFailed, - DownloadOngoing, - DownloadPending, - DownloadProgress, -) from exo.shared.types.worker.runners import RunnerId -from exo.shared.types.worker.shards import ShardMetadata from exo.utils.channels import Receiver, Sender, channel from exo.utils.event_buffer import OrderedBuffer from exo.utils.info_gatherer.info_gatherer import GatheredInfo, InfoGatherer from exo.utils.info_gatherer.net_profile import check_reachable -from exo.worker.download.download_utils import ( - map_repo_download_progress_to_download_progress_data, -) -from exo.worker.download.shard_downloader import RepoDownloadProgress, ShardDownloader +from exo.utils.keyed_backoff import KeyedBackoff from exo.worker.plan import plan from exo.worker.runner.runner_supervisor import RunnerSupervisor @@ -62,7 +56,6 @@ class Worker: self, node_id: NodeId, session_id: SessionId, - shard_downloader: ShardDownloader, *, connection_message_receiver: Receiver[ConnectionMessage], global_event_receiver: Receiver[ForwarderEvent], @@ -70,23 +63,22 @@ class Worker: # This is for requesting updates. It doesn't need to be a general command sender right now, # but I think it's the correct way to be thinking about commands command_sender: Sender[ForwarderCommand], + download_command_sender: Sender[ForwarderDownloadCommand], + event_index_counter: Iterator[int], ): self.node_id: NodeId = node_id self.session_id: SessionId = session_id - self.shard_downloader: ShardDownloader = shard_downloader - self._pending_downloads: dict[RunnerId, ShardMetadata] = {} - self.global_event_receiver = global_event_receiver self.local_event_sender = local_event_sender - self.local_event_index = 0 + self.event_index_counter = event_index_counter self.command_sender = command_sender + self.download_command_sender = download_command_sender self.connection_message_receiver = connection_message_receiver self.event_buffer = OrderedBuffer[Event]() self.out_for_delivery: dict[EventId, ForwarderEvent] = {} self.state: State = State() - self.download_status: dict[ModelId, DownloadProgress] = {} self.runners: dict[RunnerId, RunnerSupervisor] = {} self._tg: TaskGroup = create_task_group() @@ -101,6 +93,8 @@ class Worker: self.input_chunk_buffer: dict[CommandId, dict[int, str]] = {} self.input_chunk_counts: dict[CommandId, int] = {} + self._download_backoff: KeyedBackoff[ModelId] = KeyedBackoff(base=0.5, cap=10.0) + async def run(self): logger.info("Starting Worker") @@ -111,7 +105,6 @@ class Worker: tg.start_soon(info_gatherer.run) tg.start_soon(self._forward_info, info_recv) tg.start_soon(self.plan_step) - tg.start_soon(self._emit_existing_download_progress) tg.start_soon(self._connection_message_event_writer) tg.start_soon(self._resend_out_for_delivery) tg.start_soon(self._event_applier) @@ -121,6 +114,7 @@ class Worker: # Actual shutdown code - waits for all tasks to complete before executing. self.local_event_sender.close() self.command_sender.close() + self.download_command_sender.close() for runner in self.runners.values(): runner.shutdown() @@ -179,11 +173,9 @@ class Worker: async def plan_step(self): while True: await anyio.sleep(0.1) - # 3. based on the updated state, we plan & execute an operation. task: Task | None = plan( self.node_id, self.runners, - self.download_status, self.state.downloads, self.state.instances, self.state.runners, @@ -207,42 +199,26 @@ class Worker: ) ) case DownloadModel(shard_metadata=shard): - if shard.model_card.model_id not in self.download_status: - progress = DownloadPending( - shard_metadata=shard, node_id=self.node_id - ) - self.download_status[shard.model_card.model_id] = progress - await self.event_sender.send( - NodeDownloadProgress(download_progress=progress) - ) - initial_progress = ( - await self.shard_downloader.get_shard_download_status_for_shard( - shard + model_id = shard.model_card.model_id + if not self._download_backoff.should_proceed(model_id): + continue + + self._download_backoff.record_attempt(model_id) + + await self.download_command_sender.send( + ForwarderDownloadCommand( + origin=self.node_id, + command=StartDownload( + target_node_id=self.node_id, + shard_metadata=shard, + ), ) ) - if initial_progress.status == "complete": - progress = DownloadCompleted( - shard_metadata=shard, - node_id=self.node_id, - total_bytes=initial_progress.total_bytes, + await self.event_sender.send( + TaskStatusUpdated( + task_id=task.task_id, task_status=TaskStatus.Running ) - self.download_status[shard.model_card.model_id] = progress - await self.event_sender.send( - NodeDownloadProgress(download_progress=progress) - ) - await self.event_sender.send( - TaskStatusUpdated( - task_id=task.task_id, - task_status=TaskStatus.Complete, - ) - ) - else: - await self.event_sender.send( - TaskStatusUpdated( - task_id=task.task_id, task_status=TaskStatus.Running - ) - ) - self._handle_shard_download_process(task, initial_progress) + ) case Shutdown(runner_id=runner_id): try: with fail_after(3): @@ -387,104 +363,17 @@ class Worker: self._tg.start_soon(runner.run) return runner - def _handle_shard_download_process( - self, - task: DownloadModel, - initial_progress: RepoDownloadProgress, - ): - """Manages the shard download process with progress tracking.""" - status = DownloadOngoing( - node_id=self.node_id, - shard_metadata=task.shard_metadata, - download_progress=map_repo_download_progress_to_download_progress_data( - initial_progress - ), - ) - self.download_status[task.shard_metadata.model_card.model_id] = status - self.event_sender.send_nowait(NodeDownloadProgress(download_progress=status)) - - last_progress_time = 0.0 - throttle_interval_secs = 1.0 - - async def download_progress_callback( - shard: ShardMetadata, progress: RepoDownloadProgress - ) -> None: - nonlocal self - nonlocal last_progress_time - if progress.status == "complete": - status = DownloadCompleted( - shard_metadata=shard, - node_id=self.node_id, - total_bytes=progress.total_bytes, - ) - self.download_status[shard.model_card.model_id] = status - await self.event_sender.send( - NodeDownloadProgress(download_progress=status) - ) - await self.event_sender.send( - TaskStatusUpdated( - task_id=task.task_id, task_status=TaskStatus.Complete - ) - ) - elif ( - progress.status == "in_progress" - and current_time() - last_progress_time > throttle_interval_secs - ): - status = DownloadOngoing( - node_id=self.node_id, - shard_metadata=shard, - download_progress=map_repo_download_progress_to_download_progress_data( - progress - ), - ) - self.download_status[shard.model_card.model_id] = status - await self.event_sender.send( - NodeDownloadProgress(download_progress=status) - ) - last_progress_time = current_time() - - self.shard_downloader.on_progress(download_progress_callback) - - async def download_with_error_handling() -> None: - try: - await self.shard_downloader.ensure_shard(task.shard_metadata) - except Exception as e: - error_message = str(e) - logger.error( - f"Download failed for {task.shard_metadata.model_card.model_id}: {error_message}" - ) - failed_status = DownloadFailed( - node_id=self.node_id, - shard_metadata=task.shard_metadata, - error_message=error_message, - ) - self.download_status[task.shard_metadata.model_card.model_id] = ( - failed_status - ) - await self.event_sender.send( - NodeDownloadProgress(download_progress=failed_status) - ) - await self.event_sender.send( - TaskStatusUpdated( - task_id=task.task_id, task_status=TaskStatus.Failed - ) - ) - - self._tg.start_soon(download_with_error_handling) - async def _forward_events(self) -> None: with self.event_receiver as events: async for event in events: + idx = next(self.event_index_counter) fe = ForwarderEvent( - origin_idx=self.local_event_index, + origin_idx=idx, origin=self.node_id, session=self.session_id, event=event, ) - logger.debug( - f"Worker published event {self.local_event_index}: {str(event)[:100]}" - ) - self.local_event_index += 1 + logger.debug(f"Worker published event {idx}: {str(event)[:100]}") await self.local_event_sender.send(fe) self.out_for_delivery[event.event_id] = fe @@ -532,42 +421,3 @@ class Worker: await self.event_sender.send(TopologyEdgeDeleted(conn=conn)) await anyio.sleep(10) - - async def _emit_existing_download_progress(self) -> None: - try: - while True: - logger.debug("Fetching and emitting existing download progress...") - async for ( - _, - progress, - ) in self.shard_downloader.get_shard_download_status(): - if progress.status == "complete": - status = DownloadCompleted( - node_id=self.node_id, - shard_metadata=progress.shard, - total_bytes=progress.total_bytes, - ) - elif progress.status in ["in_progress", "not_started"]: - if progress.downloaded_bytes_this_session.in_bytes == 0: - status = DownloadPending( - node_id=self.node_id, shard_metadata=progress.shard - ) - else: - status = DownloadOngoing( - node_id=self.node_id, - shard_metadata=progress.shard, - download_progress=map_repo_download_progress_to_download_progress_data( - progress - ), - ) - else: - continue - - self.download_status[progress.shard.model_card.model_id] = status - await self.event_sender.send( - NodeDownloadProgress(download_progress=status) - ) - logger.debug("Done emitting existing download progress.") - await anyio.sleep(5 * 60) # 5 minutes - except Exception as e: - logger.error(f"Error emitting existing download progress: {e}") diff --git a/src/exo/worker/plan.py b/src/exo/worker/plan.py index 7877b96c..cb6fbbf3 100644 --- a/src/exo/worker/plan.py +++ b/src/exo/worker/plan.py @@ -2,7 +2,6 @@ from collections.abc import Mapping, Sequence -from exo.shared.models.model_cards import ModelId from exo.shared.types.common import CommandId, NodeId from exo.shared.types.tasks import ( ChatCompletion, @@ -45,9 +44,6 @@ def plan( node_id: NodeId, # Runners is expected to be FRESH and so should not come from state runners: Mapping[RunnerId, RunnerSupervisor], - # DL_status is expected to be FRESH and so should not come from state - download_status: Mapping[ModelId, DownloadProgress], - # gdls is not expected to be fresh global_download_status: Mapping[NodeId, Sequence[DownloadProgress]], instances: Mapping[InstanceId, Instance], all_runners: Mapping[RunnerId, RunnerStatus], # all global @@ -59,7 +55,7 @@ def plan( return ( _kill_runner(runners, all_runners, instances) or _create_runner(node_id, runners, instances) - or _model_needs_download(runners, download_status) + or _model_needs_download(node_id, runners, global_download_status) or _init_distributed_backend(runners, all_runners) or _load_model(runners, all_runners, global_download_status) or _ready_to_warmup(runners, all_runners) @@ -115,9 +111,15 @@ def _create_runner( def _model_needs_download( + node_id: NodeId, runners: Mapping[RunnerId, RunnerSupervisor], - download_status: Mapping[ModelId, DownloadProgress], + global_download_status: Mapping[NodeId, Sequence[DownloadProgress]], ) -> DownloadModel | None: + local_downloads = global_download_status.get(node_id, []) + download_status = { + dp.shard_metadata.model_card.model_id: dp for dp in local_downloads + } + for runner in runners.values(): model_id = runner.bound_instance.bound_shard.model_card.model_id if isinstance(runner.status, RunnerIdle) and ( diff --git a/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py b/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py index d0b6707b..1f51b703 100644 --- a/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py +++ b/src/exo/worker/tests/unittests/test_mlx/test_tokenizers.py @@ -11,12 +11,12 @@ from pathlib import Path import pytest -from exo.shared.models.model_cards import MODEL_CARDS, ModelCard, ModelId -from exo.worker.download.download_utils import ( +from exo.download.download_utils import ( download_file_with_retry, ensure_models_dir, fetch_file_list_with_cache, ) +from exo.shared.models.model_cards import MODEL_CARDS, ModelCard, ModelId from exo.worker.engines.mlx.utils_mlx import ( get_eos_token_ids_for_model, load_tokenizer_for_model_id, diff --git a/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py b/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py index 20cf72df..28ecd8e4 100644 --- a/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py +++ b/src/exo/worker/tests/unittests/test_plan/test_download_and_loading.py @@ -1,5 +1,5 @@ import exo.worker.plan as plan_mod -from exo.shared.types.common import ModelId, NodeId +from exo.shared.types.common import NodeId from exo.shared.types.memory import Memory from exo.shared.types.tasks import LoadModel from exo.shared.types.worker.downloads import DownloadCompleted, DownloadProgress @@ -45,13 +45,9 @@ def test_plan_requests_download_when_waiting_and_shard_not_downloaded(): instances = {INSTANCE_1_ID: instance} all_runners = {RUNNER_1_ID: RunnerIdle()} - # No entry for this shard -> should trigger DownloadModel - download_status: dict[ModelId, DownloadProgress] = {} - result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status=download_status, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -92,14 +88,6 @@ def test_plan_loads_model_when_all_shards_downloaded_and_waiting(): RUNNER_2_ID: RunnerConnected(), } - # Local node has already marked its shard as downloaded (not actually used by _load_model) - local_download_status = { - MODEL_A_ID: DownloadCompleted( - shard_metadata=shard1, node_id=NODE_A, total_bytes=Memory() - ) - } - - # Global view has completed downloads for both nodes global_download_status = { NODE_A: [ DownloadCompleted( @@ -116,7 +104,6 @@ def test_plan_loads_model_when_all_shards_downloaded_and_waiting(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status=local_download_status, global_download_status=global_download_status, instances=instances, all_runners=all_runners, @@ -148,23 +135,19 @@ def test_plan_does_not_request_download_when_shard_already_downloaded(): instances = {INSTANCE_1_ID: instance} all_runners = {RUNNER_1_ID: RunnerIdle()} - # Local status claims the shard is downloaded already - local_download_status = { - MODEL_A_ID: DownloadCompleted( - shard_metadata=shard, node_id=NODE_A, total_bytes=Memory() - ) - } - - # Global view hasn't caught up yet (no completed shards recorded for NODE_A) + # Global state shows shard is downloaded for NODE_A global_download_status: dict[NodeId, list[DownloadProgress]] = { - NODE_A: [], + NODE_A: [ + DownloadCompleted( + shard_metadata=shard, node_id=NODE_A, total_bytes=Memory() + ) + ], NODE_B: [], } result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status=local_download_status, global_download_status=global_download_status, instances=instances, all_runners=all_runners, @@ -202,12 +185,6 @@ def test_plan_does_not_load_model_until_all_shards_downloaded_globally(): RUNNER_2_ID: RunnerConnected(), } - # Only NODE_A's shard is recorded as downloaded globally - local_download_status = { - MODEL_A_ID: DownloadCompleted( - shard_metadata=shard1, node_id=NODE_A, total_bytes=Memory() - ) - } global_download_status = { NODE_A: [ DownloadCompleted( @@ -220,7 +197,6 @@ def test_plan_does_not_load_model_until_all_shards_downloaded_globally(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status=local_download_status, global_download_status=global_download_status, instances=instances, all_runners=all_runners, @@ -245,7 +221,6 @@ def test_plan_does_not_load_model_until_all_shards_downloaded_globally(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status=local_download_status, global_download_status=global_download_status, instances=instances, all_runners=all_runners, diff --git a/src/exo/worker/tests/unittests/test_plan/test_runner_lifecycle.py b/src/exo/worker/tests/unittests/test_plan/test_runner_lifecycle.py index 944cb6db..a1b87ca1 100644 --- a/src/exo/worker/tests/unittests/test_plan/test_runner_lifecycle.py +++ b/src/exo/worker/tests/unittests/test_plan/test_runner_lifecycle.py @@ -47,8 +47,7 @@ def test_plan_kills_runner_when_instance_missing(): result = plan_mod.plan( node_id=NODE_A, - runners=runners, # type: ignore - download_status={}, + runners=runners, # type: ignore[arg-type] global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -87,8 +86,7 @@ def test_plan_kills_runner_when_sibling_failed(): result = plan_mod.plan( node_id=NODE_A, - runners=runners, # type: ignore - download_status={}, + runners=runners, # type: ignore[arg-type] global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -120,7 +118,6 @@ def test_plan_creates_runner_when_missing_for_node(): result = plan_mod.plan( node_id=NODE_A, runners=runners, - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -158,8 +155,7 @@ def test_plan_does_not_create_runner_when_supervisor_already_present(): result = plan_mod.plan( node_id=NODE_A, - runners=runners, # type: ignore - download_status={}, + runners=runners, # type: ignore[arg-type] global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -189,7 +185,6 @@ def test_plan_does_not_create_runner_for_unassigned_node(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, 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 5b3aa5ec..6c8fea8d 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 @@ -65,7 +65,6 @@ def test_plan_forwards_pending_chat_completion_when_runner_ready(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -113,7 +112,6 @@ def test_plan_does_not_forward_chat_completion_if_any_runner_not_ready(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: [], NODE_B: []}, instances=instances, all_runners=all_runners, @@ -158,7 +156,6 @@ def test_plan_does_not_forward_tasks_for_other_instances(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -221,7 +218,6 @@ def test_plan_ignores_non_pending_or_non_chat_tasks(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: [], NODE_B: []}, instances=instances, all_runners=all_runners, @@ -261,7 +257,6 @@ def test_plan_returns_none_when_nothing_to_do(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: [], NODE_B: []}, instances=instances, all_runners=all_runners, diff --git a/src/exo/worker/tests/unittests/test_plan/test_warmup.py b/src/exo/worker/tests/unittests/test_plan/test_warmup.py index b42a5afd..52848446 100644 --- a/src/exo/worker/tests/unittests/test_plan/test_warmup.py +++ b/src/exo/worker/tests/unittests/test_plan/test_warmup.py @@ -57,7 +57,6 @@ def test_plan_starts_warmup_for_accepting_rank_when_all_loaded_or_warming(): result = plan_mod.plan( node_id=NODE_B, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -99,7 +98,6 @@ def test_plan_starts_warmup_for_rank_zero_after_others_warming(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -140,7 +138,6 @@ def test_plan_does_not_start_warmup_for_non_zero_rank_until_all_loaded_or_warmin result = plan_mod.plan( node_id=NODE_B, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: [], NODE_B: []}, instances=instances, all_runners=all_runners, @@ -185,7 +182,6 @@ def test_plan_does_not_start_warmup_for_rank_zero_until_others_warming(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -202,7 +198,6 @@ def test_plan_does_not_start_warmup_for_rank_zero_until_others_warming(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: []}, instances=instances, all_runners=all_runners, @@ -246,7 +241,6 @@ def test_plan_starts_warmup_for_connecting_rank_after_others_warming(): result = plan_mod.plan( node_id=NODE_B, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_B: []}, instances=instances, all_runners=all_runners, @@ -289,7 +283,6 @@ def test_plan_does_not_start_warmup_for_accepting_rank_until_all_loaded_or_warmi result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: [], NODE_B: []}, instances=instances, all_runners=all_runners, @@ -331,7 +324,6 @@ def test_plan_does_not_start_warmup_for_connecting_rank_until_others_warming(): result = plan_mod.plan( node_id=NODE_A, runners=runners, # type: ignore - download_status={}, global_download_status={NODE_A: [], NODE_B: []}, instances=instances, all_runners=all_runners, diff --git a/tests/headless_runner.py b/tests/headless_runner.py index 354186a5..30fd19d4 100644 --- a/tests/headless_runner.py +++ b/tests/headless_runner.py @@ -11,6 +11,10 @@ from hypercorn.asyncio import serve # pyright: ignore[reportUnknownVariableType from loguru import logger from pydantic import BaseModel +from exo.download.impl_shard_downloader import ( + build_full_shard, + exo_shard_downloader, +) from exo.shared.logging import InterceptLogger, logger_setup from exo.shared.models.model_cards import MODEL_CARDS, ModelId from exo.shared.types.api import ChatCompletionMessage, ChatCompletionTaskParams @@ -36,10 +40,6 @@ from exo.shared.types.worker.runners import RunnerId, ShardAssignments from exo.shared.types.worker.shards import PipelineShardMetadata, TensorShardMetadata from exo.utils.channels import MpReceiver, MpSender, channel, mp_channel from exo.utils.info_gatherer.info_gatherer import GatheredInfo, InfoGatherer -from exo.worker.download.impl_shard_downloader import ( - build_full_shard, - exo_shard_downloader, -) from exo.worker.runner.bootstrap import entrypoint From 6dbbe7797b5b2596e1b796b9682725287a399a4c Mon Sep 17 00:00:00 2001 From: Jake Hillion Date: Mon, 19 Jan 2026 19:13:55 +0000 Subject: [PATCH 06/15] downloads: add download and delete buttons to downloads UI The downloads page showed model download progress but provided no way for users to trigger downloads or remove completed models from disk. Added API endpoints (POST /download/start, DELETE /download/{node_id}/{model_id}) that send StartDownload and DeleteDownload commands via the download_command_sender. Updated the dashboard downloads page with per-model buttons: a download button for incomplete downloads and a delete button for completed ones. This allows users to manage downloads directly from the UI without needing to trigger downloads through other means. Test plan: - Deployed on a 3 machine cluster. Did several downloads/deletions - all work and the dashboard updates relatively fluently. It takes roughly 5 seconds to render a 131GB model deletion which isn't too bad. --- dashboard/src/lib/stores/app.svelte.ts | 54 ++++++++++++++++++++ dashboard/src/routes/downloads/+page.svelte | 56 +++++++++++++++++++++ src/exo/main.py | 1 + src/exo/master/api.py | 36 +++++++++++++ src/exo/shared/types/api.py | 18 ++++++- 5 files changed, 163 insertions(+), 2 deletions(-) diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts index 1b1a3847..1cbedbc2 100644 --- a/dashboard/src/lib/stores/app.svelte.ts +++ b/dashboard/src/lib/stores/app.svelte.ts @@ -2228,6 +2228,54 @@ class AppStore { this.conversations.find((c) => c.id === this.activeConversationId) || null ); } + + /** + * Start a download on a specific node + */ + async startDownload(nodeId: string, shardMetadata: object): Promise { + try { + const response = await fetch("/download/start", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + targetNodeId: nodeId, + shardMetadata: shardMetadata, + }), + }); + if (!response.ok) { + const errorText = await response.text(); + throw new Error( + `Failed to start download: ${response.status} - ${errorText}`, + ); + } + } catch (error) { + console.error("Error starting download:", error); + throw error; + } + } + + /** + * Delete a downloaded model from a specific node + */ + async deleteDownload(nodeId: string, modelId: string): Promise { + try { + const response = await fetch( + `/download/${encodeURIComponent(nodeId)}/${encodeURIComponent(modelId)}`, + { + method: "DELETE", + }, + ); + if (!response.ok) { + const errorText = await response.text(); + throw new Error( + `Failed to delete download: ${response.status} - ${errorText}`, + ); + } + } catch (error) { + console.error("Error deleting download:", error); + throw error; + } + } } export const appStore = new AppStore(); @@ -2333,3 +2381,9 @@ export const setImageGenerationParams = ( ) => appStore.setImageGenerationParams(params); export const resetImageGenerationParams = () => appStore.resetImageGenerationParams(); + +// Download actions +export const startDownload = (nodeId: string, shardMetadata: object) => + appStore.startDownload(nodeId, shardMetadata); +export const deleteDownload = (nodeId: string, modelId: string) => + appStore.deleteDownload(nodeId, modelId); diff --git a/dashboard/src/routes/downloads/+page.svelte b/dashboard/src/routes/downloads/+page.svelte index a7ee2003..72fe149c 100644 --- a/dashboard/src/routes/downloads/+page.svelte +++ b/dashboard/src/routes/downloads/+page.svelte @@ -6,6 +6,8 @@ type DownloadProgress, refreshState, lastUpdate as lastUpdateStore, + startDownload, + deleteDownload, } from "$lib/stores/app.svelte"; import HeaderNav from "$lib/components/HeaderNav.svelte"; @@ -28,6 +30,7 @@ etaMs: number; status: "completed" | "downloading"; files: FileProgress[]; + shardMetadata?: Record; }; type NodeEntry = { @@ -269,6 +272,12 @@ } } + // Extract shard_metadata for use with download actions + const shardMetadata = (downloadPayload.shard_metadata ?? + downloadPayload.shardMetadata) as + | Record + | undefined; + const entry: ModelEntry = { modelId, prettyName, @@ -285,6 +294,7 @@ ? "completed" : "downloading", files, + shardMetadata, }; const existing = modelMap.get(modelId); @@ -469,6 +479,52 @@ > {pct.toFixed(1)}% + {#if model.status !== "completed" && model.shardMetadata} + + {/if} + {#if model.status === "completed"} + + {/if}