diff --git a/dashboard/src/routes/+page.svelte b/dashboard/src/routes/+page.svelte index 4c68ac6c..4c35d9ef 100644 --- a/dashboard/src/routes/+page.svelte +++ b/dashboard/src/routes/+page.svelte @@ -701,6 +701,22 @@ return null; }); + // Helper to get onboarding model loading progress (layers loaded) + const onboardingLoadProgress = $derived.by(() => { + if (instanceCount === 0 || !onboardingModelId) return null; + let layersLoaded = 0, totalLayers = 0; + for (const [, inst] of Object.entries(instanceData)) { + if (getInstanceModelId(inst) !== onboardingModelId) continue; + const status = deriveInstanceStatus(inst); + if (status.statusText === "LOADING" && status.totalLayers && status.totalLayers > 0) { + layersLoaded += status.layersLoaded ?? 0; + totalLayers += status.totalLayers; + } + } + if (totalLayers === 0) return null; + return { layersLoaded, totalLayers, percentage: (layersLoaded / totalLayers) * 100 }; + }); + // Instance launch state let models = $state< Array<{ @@ -1720,6 +1736,8 @@ function deriveInstanceStatus(instanceWrapped: unknown): { statusText: string; statusClass: string; + layersLoaded?: number; + totalLayers?: number; } { const [, instance] = getTagged(instanceWrapped); if (!instance || typeof instance !== "object") { @@ -1759,8 +1777,20 @@ if (has("Failed")) return { statusText: "FAILED", statusClass: "failed" }; if (has("Shutdown")) return { statusText: "SHUTDOWN", statusClass: "inactive" }; - if (has("Loading")) - return { statusText: "LOADING", statusClass: "starting" }; + if (has("Loading")) { + let layersLoaded = 0, totalLayers = 0; + for (const rid of runnerIds) { + const r = runnersData[rid]; + if (!r) continue; + const [kind, payload] = getTagged(r); + if (kind === "RunnerLoading" && payload && typeof payload === "object") { + const p = payload as { layersLoaded?: number; totalLayers?: number }; + layersLoaded += p.layersLoaded ?? 0; + totalLayers += p.totalLayers ?? 0; + } + } + return { statusText: "LOADING", statusClass: "starting", layersLoaded, totalLayers }; + } if (has("WarmingUp")) return { statusText: "WARMING UP", statusClass: "starting" }; if (has("Running")) @@ -3502,13 +3532,26 @@ -
Almost ready...
+ {#if onboardingLoadProgress} ++ {onboardingLoadProgress.layersLoaded} / {onboardingLoadProgress.totalLayers} layers loaded +
+Loading...
+ {/if} {:else if onboardingStep === 9} @@ -4397,12 +4440,27 @@ {downloadInfo.statusText} {#if isLoading} -- Loading model into memory for fast - inference... -
+ {@const loadStatus = deriveInstanceStatus(instance)} + {#if loadStatus.totalLayers && loadStatus.totalLayers > 0} ++ Loading model into memory... +
+ {/if} {:else if isReady || isRunning}{#if isLoading} -
- Loading model into memory for fast - inference... -
+ {@const loadStatus = deriveInstanceStatus(instance)} + {#if loadStatus.totalLayers && loadStatus.totalLayers > 0} ++ Loading model into memory... +
+ {/if} {:else if isReady || isRunning}nn.Module: """ Automatically parallelize a model across multiple devices. @@ -238,8 +240,11 @@ def pipeline_auto_parallel( device_rank, world_size = model_shard_meta.device_rank, model_shard_meta.world_size layers = layers[start_layer:end_layer] - for layer in layers: + total = len(layers) + for i, layer in enumerate(layers): mx.eval(layer) # type: ignore + if on_layer_loaded is not None: + on_layer_loaded(i + 1, total) layers[0] = PipelineFirstLayer(layers[0], device_rank, group=group) layers[-1] = PipelineLastLayer( diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index 360a018b..5cbd5745 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -55,6 +55,7 @@ from exo.shared.types.worker.shards import ( ) from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.auto_parallel import ( + LayerLoadedCallback, TimeoutCallback, eval_with_timeout, pipeline_auto_parallel, @@ -172,6 +173,7 @@ def load_mlx_items( bound_instance: BoundInstance, group: Group | None, on_timeout: TimeoutCallback | None = None, + on_layer_loaded: LayerLoadedCallback | None = None, ) -> tuple[Model, TokenizerWrapper]: if group is None: logger.info(f"Single device used for {bound_instance.instance}") @@ -186,7 +188,7 @@ def load_mlx_items( logger.info("Starting distributed init") start_time = time.perf_counter() model, tokenizer = shard_and_load( - bound_instance.bound_shard, group=group, on_timeout=on_timeout + bound_instance.bound_shard, group=group, on_timeout=on_timeout, on_layer_loaded=on_layer_loaded ) end_time = time.perf_counter() logger.info( @@ -202,6 +204,7 @@ def shard_and_load( shard_metadata: ShardMetadata, group: Group, on_timeout: TimeoutCallback | None = None, + on_layer_loaded: LayerLoadedCallback | None = None, ) -> tuple[nn.Module, TokenizerWrapper]: model_path = build_model_path(shard_metadata.model_card.model_id) @@ -245,7 +248,7 @@ def shard_and_load( model = tensor_auto_parallel(model, group, timeout_seconds, on_timeout) case PipelineShardMetadata(): logger.info(f"loading model from {model_path} with pipeline parallelism") - model = pipeline_auto_parallel(model, group, shard_metadata) + model = pipeline_auto_parallel(model, group, shard_metadata, on_layer_loaded=on_layer_loaded) eval_with_timeout(model.parameters(), timeout_seconds, on_timeout) case CfgShardMetadata(): raise ValueError( diff --git a/src/exo/worker/runner/llm_inference/runner.py b/src/exo/worker/runner/llm_inference/runner.py index c4059b9e..09fd676f 100644 --- a/src/exo/worker/runner/llm_inference/runner.py +++ b/src/exo/worker/runner/llm_inference/runner.py @@ -150,7 +150,8 @@ def main( case LoadModel() if ( isinstance(current_status, RunnerConnected) and group is not None ) or (isinstance(current_status, RunnerIdle) and group is None): - current_status = RunnerLoading() + total_layers = shard_metadata.end_layer - shard_metadata.start_layer + current_status = RunnerLoading(layers_loaded=0, total_layers=total_layers) logger.info("runner loading") event_sender.send( RunnerStatusUpdated( @@ -170,11 +171,18 @@ def main( ) time.sleep(0.5) + def on_layer_loaded(layers_loaded: int, total: int) -> None: + event_sender.send( + RunnerStatusUpdated( + runner_id=runner_id, runner_status=RunnerLoading(layers_loaded=layers_loaded, total_layers=total) + ) + ) + assert ( ModelTask.TextGeneration in shard_metadata.model_card.tasks ), f"Incorrect model task(s): {shard_metadata.model_card.tasks}" inference_model, tokenizer = load_mlx_items( - bound_instance, group, on_timeout=on_model_load_timeout + bound_instance, group, on_timeout=on_model_load_timeout, on_layer_loaded=on_layer_loaded ) logger.info( f"model has_tool_calling={tokenizer.has_tool_calling} using tokens {tokenizer.tool_call_start}, {tokenizer.tool_call_end}" 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 6d964583..b4d02453 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 @@ -224,7 +224,7 @@ def test_events_processed_in_correct_order(patch_out_mlx: pytest.MonkeyPatch): ), RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerConnected()), TaskStatusUpdated(task_id=LOAD_TASK_ID, task_status=TaskStatus.Running), - RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerLoading()), + RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerLoading(layers_loaded=0, total_layers=32)), 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()),