Compare commits

..
Author SHA1 Message Date
Ryuichi Leo Takashige 84db569167 Optimizations 6 2026-04-29 13:29:36 +01:00
Ryuichi Leo Takashige b83d5e6a6f Optimizations 5 2026-04-29 13:29:36 +01:00
Ryuichi Leo Takashige 9e67e89862 Optimizations 4 2026-04-29 13:29:36 +01:00
Evan 92ea4ed0a4 banner update 2026-04-29 13:29:28 +01:00
Evan 9f37340b89 cleanup 2026-04-29 08:01:53 +01:00
Evan 344381fd74 snailed it 2026-04-29 08:00:50 +01:00
Ryuichi Leo Takashige 701c9b1cf6 Optimizations 3 2026-04-29 00:53:48 +01:00
Ryuichi Leo Takashige dc709e933a Optimizations 2 2026-04-29 00:02:18 +01:00
Ryuichi Leo Takashige 0a736d7eaf Optimizations 2026-04-28 20:50:47 +01:00
Ryuichi Leo Takashige 94b1813f76 tmp 2 2026-04-28 20:40:56 +01:00
Ryuichi Leo Takashige 8774513367 tmp 2026-04-28 17:13:42 +01:00
Ryuichi Leo Takashige 35e3335d6d Select VLLM instances 2026-04-28 15:58:09 +01:00
Ryuichi Leo Takashige c2b35f4d9e Fix linux CI 2026-04-28 14:59:31 +01:00
Ryuichi Leo Takashige d96f8379ce Add Linux dashboard 2026-04-28 14:52:57 +01:00
Ryuichi Leo Takashige c1eca8d026 Fix pyproject for Macs 2026-04-28 14:36:21 +01:00
Evan dbc736c845 vllm support 2026-04-28 02:08:10 +01:00
667a3bb0e5 feat: keep-models option when uninstalling EXO (#1997)
## Summary

- Adds a **Keep downloaded models (~/.exo/models)** checkbox to the
macOS uninstall confirmation dialog (Settings → Advanced → Danger Zone).
The full `~/.exo` directory is now removed on uninstall by default; if
the checkbox is checked, `~/.exo/models` is preserved.
- The standalone `app/EXO/uninstall-exo.sh` gains a matching
`--keep-models` flag and the same `~/.exo` cleanup so GUI and CLI flows
stay in sync. Resolves the user home via `$SUDO_USER` since the script
runs under `sudo`.

Previously, "Uninstall EXO" only cleaned up system-level components
(LaunchDaemon, network location, logs, app bundle) and left the entire
`~/.exo` directory behind. Now uninstalling actually removes EXO's user
data, with a one-click opt-out for the (potentially many GB) of
downloaded models.

![Uninstall dialog with new
checkbox](https://raw.githubusercontent.com/exo-explore/exo/703b7fbbf13441217ad2903bb199f07e92af4490/uninstall-dialog.png)

> Note: the rendered icon in the screenshot above is the generic system
folder icon because it was captured from a small standalone Swift binary
(no app bundle / icon resource). When triggered from the actual EXO.app,
the EXO app icon is shown.

## Test plan

- [ ] Build EXO.app locally; open Settings → Advanced → Danger Zone →
Uninstall EXO; confirm the new "Keep downloaded models (~/.exo/models)"
checkbox is present and unchecked by default.
- [ ] Uninstall with the checkbox **checked** → `~/.exo/models/`
survives, all other entries under `~/.exo` are gone, system components
removed, app moved to Trash.
- [ ] Uninstall with the checkbox **unchecked** → `~/.exo` is fully
removed.
- [ ] `sudo app/EXO/uninstall-exo.sh --keep-models` → `~/.exo/models/`
is preserved, the rest of `~/.exo` is removed.
- [ ] `sudo app/EXO/uninstall-exo.sh` (no flag) → `~/.exo` is fully
removed.
- [ ] `app/EXO/uninstall-exo.sh --help` prints usage and exits 0;
unknown args exit 2 with a usage hint.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

---------

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
Co-authored-by: Evan <evanev7@gmail.com>
2026-04-28 01:06:02 +00:00
Evan Quiney c80b10c013 implement engine abstraction for mlx and mflux (#2000)
refactor for future versions.
2026-04-28 00:58:17 +00:00
Alex CheemaandClaude Opus 4.7 18ffe1df23 fix: uninstall-exo.sh removes both current and legacy bridge scripts (#1998)
## Summary

The standalone `app/EXO/uninstall-exo.sh` only knew about the legacy
filename `disable_bridge_enable_dhcp.sh`. On machines installed with
newer EXO versions, the current `/Library/Application
Support/EXO/disable_bridge.sh` was left behind, and the script then
reported `EXO support directory not empty, leaving in place`.

This PR makes the script try both filenames, removing whichever ones
exist. Tolerates **either**, **both**, or **neither** being present
without erroring.

The Swift `NetworkSetupHelper.makeUninstallScript()` already handles
both paths correctly, so the GUI uninstall flow is unaffected — this is
a script-only fix.

Caught while running an end-to-end uninstall on a real machine for
#1997.

## Test plan

Verified the new block in isolation against all four states:

- [x] both `disable_bridge.sh` and `disable_bridge_enable_dhcp.sh`
present → both removed
- [x] only `disable_bridge.sh` present → removed cleanly
- [x] only `disable_bridge_enable_dhcp.sh` present → removed cleanly
(legacy install)
- [x] neither present → prints the existing "already removed?" warning,
exits 0

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-04-28 00:28:12 +00:00
rltakashige f0d1371d89 MLX P/D (#1993)
## Motivation

MLX only prefill server for Apple Silicon
2026-04-28 00:12:42 +00:00
5d10188d3a fix: route by in-flight tasks only — completed tasks were skewing load balance (#1989)
The load balancer counted ALL tasks (Complete, Cancelled, TimedOut,
Failed) instead of only Pending/Running ones. With 138 accumulated tasks
and only 7 active, routing decisions were based on historical
distribution, causing one node to appear permanently 'busier' and
starving the other of work.

Co-authored-by: Adam Durham <adam@example.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-04-27 16:03:12 +00:00
ciaranbor f2a0db4e23 Extend bench/eval tooling (#1905)
## Motivation

Extend bench/eval tooling with robustness features, streaming support,
and align model configs with vllm eval for reproducible comparisons.

## Changes

- **exo_eval**: Checkpoint/resume (JSONL), instance health monitoring +
early abort, `top_k`/`min_p`/`enable_thinking` params, LCB
`--release-version`/`--offset`
- **exo_bench**: Streaming SSE (`--stream`), Kimi tokenizer fix for
transformers 5.x
- **Both tools**: Auto-detect running instances instead of requiring
`--skip-instance-setup`; `--fresh-instance` to override
- **harness**: SSE streaming client, `find_existing_instance()` shared
helper, removed download timeout, settle-timeout default 0→7200s
- **models.toml**: Added `enable_thinking`, aligned `max_tokens`/temps
with vllm, added new models
- **API**: Streaming SSE for `/bench/chat/completions`

## Why It Works

- Checkpoint/resume uses append-only JSONL + skip-on-load so interrupted
evals resume without re-running completed questions
- Health monitoring races an `asyncio.Event` against API calls for fast
abort when the instance dies
- Auto-detection queries `/state` for existing instances matching the
model ID before attempting placement
- Streaming reuses the existing `generate_chat_stream` infrastructure
from the regular chat endpoint
2026-04-27 16:53:43 +01:00
rltakashigeandEvan 37f6f4f6c2 Add DeepSeek V4 Flash/Pro (#1978)
Wait for upstream merge.

---------

Co-authored-by: Evan <evanev7@gmail.com>
2026-04-27 15:20:50 +01:00
Adam DurhamandAdam Durham 48a922fd5c fix: map presence_penalty and frequency_penalty from ChatCompletionRequest (#1991)
Upstream PR #1947 added `presence_penalty` and `frequency_penalty` to
`TextGenerationTaskParams` and the mlx-lm generator call sites, but
missed wiring them up in the API adapter so they were silently dropped
from incoming requests. This fixes the API mapping.

Co-authored-by: Adam Durham <adam@example.com>
2026-04-27 08:58:59 +00:00
137 changed files with 16726 additions and 2727 deletions

No files matched your search

-7
View File
@@ -1,8 +1 @@
use flake
# creates .venv if doesn't exist and loads its environment
export VIRTUAL_ENV=".venv"
if ! [ -d "./$VIRTUAL_ENV" ]; then
uv venv
fi
layout python
+3 -6
View File
@@ -191,13 +191,10 @@ class RotatingKVCache(_BaseCache):
def state(self, v): # -> None:
...
@property
def meta_state(self): # -> tuple[str, ...]:
...
def meta_state(self) -> tuple[str, ...]: ...
@meta_state.setter
def meta_state(self, v): # -> None:
...
def is_trimmable(self): # -> bool:
...
def meta_state(self, v: tuple[str, ...]) -> None: ...
def is_trimmable(self) -> bool: ...
def trim(self, n: int) -> int: ...
def to_quantized(
self, group_size: int = ..., bits: int = ...
+19 -6
View File
@@ -108,12 +108,10 @@ class Compressor(nn.Module):
def __call__(
self,
x: mx.array,
state: ArraysCache,
cache: "DeepseekV4Cache",
offset: Any,
slot_compressed: int,
slot_kv_state: int,
slot_score_state: int,
) -> Optional[mx.array]: ...
key: str = ...,
) -> mx.array: ...
class Indexer(nn.Module):
def __init__(
@@ -123,14 +121,29 @@ class Indexer(nn.Module):
rope: DeepseekV4RoPE,
) -> None: ...
class _CompressorBranch:
buffer_kv: Optional[mx.array]
buffer_gate: Optional[mx.array]
prev_kv: Optional[mx.array]
prev_gate: Optional[mx.array]
pool: Optional[mx.array]
buffer_lengths: Optional[List[int]]
pool_lengths: Optional[List[int]]
buffer_count: int
_new_pool_lengths: Optional[List[int]]
def __init__(self) -> None: ...
class DeepseekV4Cache:
local: Any
local: RotatingKVCache
offset: int
keys: Optional[mx.array]
values: Optional[mx.array]
state: Any
meta_state: Any
nbytes: int
_branches: Dict[str, _CompressorBranch]
_pending_lengths: Optional[List[int]]
def __init__(self, sliding_window: int) -> None: ...
def update_and_fetch(
File diff suppressed because it is too large. Load diff
+30 -3
View File
@@ -552,15 +552,24 @@ struct SettingsView: View {
let alert = NSAlert()
alert.messageText = "Uninstall EXO"
alert.informativeText = """
This will remove EXO and all its system components:
This will remove EXO and all its components:
• Network configuration daemon
• Launch at login registration
• EXO network location
• EXO data directory (~/.exo)
The app will be moved to Trash.
"""
alert.alertStyle = .warning
let checkbox = NSButton(
checkboxWithTitle: "Keep downloaded models (~/.exo/models)",
target: nil, action: nil)
checkbox.state = .off
checkbox.sizeToFit()
alert.accessoryView = checkbox
alert.addButton(withTitle: "Uninstall")
alert.addButton(withTitle: "Cancel")
@@ -570,11 +579,11 @@ struct SettingsView: View {
let response = alert.runModal()
if response == .alertFirstButtonReturn {
performUninstall()
performUninstall(keepModels: checkbox.state == .on)
}
}
private func performUninstall() {
private func performUninstall(keepModels: Bool) {
uninstallInProgress = true
controller.cancelPendingLaunch()
@@ -584,6 +593,7 @@ struct SettingsView: View {
DispatchQueue.global(qos: .utility).async {
do {
try NetworkSetupHelper.uninstall()
try Self.removeExoDirectory(keepModels: keepModels)
DispatchQueue.main.async {
LaunchAtLoginHelper.disable()
@@ -607,6 +617,23 @@ struct SettingsView: View {
}
}
private static func removeExoDirectory(keepModels: Bool) throws {
let fm = FileManager.default
let exoDir = ExoProcessController.exoDirectoryURL
guard fm.fileExists(atPath: exoDir.path) else { return }
if !keepModels {
try fm.removeItem(at: exoDir)
return
}
let contents = try fm.contentsOfDirectory(
at: exoDir, includingPropertiesForKeys: nil, options: [])
for entry in contents where entry.lastPathComponent != "models" {
try? fm.removeItem(at: entry)
}
}
private func moveAppToTrash() {
guard let appURL = Bundle.main.bundleURL as URL? else { return }
do {
+63 -7
View File
@@ -3,25 +3,55 @@
# EXO Uninstaller Script
#
# This script removes all EXO system components that persist after deleting the app.
# Run with: sudo ./uninstall-exo.sh
# Run with: sudo ./uninstall-exo.sh [--keep-models]
#
# Options:
# --keep-models Preserve ~/.exo/models when removing the EXO data directory.
#
# Components removed:
# - LaunchDaemon: /Library/LaunchDaemons/io.exo.networksetup.plist
# - Network script: /Library/Application Support/EXO/
# - Log files: /var/log/io.exo.networksetup.*
# - Network location: "exo"
# - EXO data directory: ~/.exo (or all of ~/.exo except models/ when --keep-models is set)
# - Launch at login registration
#
set -euo pipefail
KEEP_MODELS=0
for arg in "$@"; do
case "$arg" in
--keep-models)
KEEP_MODELS=1
;;
-h | --help)
echo "Usage: sudo ./uninstall-exo.sh [--keep-models]"
echo " --keep-models Preserve ~/.exo/models when removing the EXO data directory."
exit 0
;;
*)
echo "Unknown argument: $arg" >&2
echo "Usage: sudo ./uninstall-exo.sh [--keep-models]" >&2
exit 2
;;
esac
done
LABEL="io.exo.networksetup"
SCRIPT_DEST="/Library/Application Support/EXO/disable_bridge_enable_dhcp.sh"
# Current script path. Older installs used a different filename; keep the
# legacy path here so a fresh uninstall still cleans up upgraded machines.
CURRENT_SCRIPT_DEST="/Library/Application Support/EXO/disable_bridge.sh"
LEGACY_SCRIPT_DEST="/Library/Application Support/EXO/disable_bridge_enable_dhcp.sh"
PLIST_DEST="/Library/LaunchDaemons/io.exo.networksetup.plist"
LOG_OUT="/var/log/${LABEL}.log"
LOG_ERR="/var/log/${LABEL}.err.log"
APP_BUNDLE_ID="io.exo.EXO"
# Resolve the invoking user's home, even when run via sudo.
USER_HOME="$(eval echo "~${SUDO_USER:-$USER}")"
EXO_DIR="$USER_HOME/.exo"
# Colors for output
RED='\033[0;31m'
GREEN='\033[0;32m'
@@ -69,11 +99,17 @@ else
echo_warn "LaunchDaemon plist not found (already removed?)"
fi
# Remove the script and parent directory
if [[ -f $SCRIPT_DEST ]]; then
rm -f "$SCRIPT_DEST"
echo_info "Removed network setup script"
else
# Remove the script (current and legacy filenames) — backwards-compatible:
# tolerate either, both, or neither being present.
removed_any_script=0
for script in "$CURRENT_SCRIPT_DEST" "$LEGACY_SCRIPT_DEST"; do
if [[ -f $script ]]; then
rm -f "$script"
echo_info "Removed network setup script: $script"
removed_any_script=1
fi
done
if [[ $removed_any_script -eq 0 ]]; then
echo_warn "Network setup script not found (already removed?)"
fi
@@ -115,6 +151,22 @@ if networksetup -listnetworkservices 2>/dev/null | grep -q "Thunderbolt Bridge";
echo_info "Re-enabled Thunderbolt Bridge"
fi
# Remove EXO data directory (~/.exo)
EXO_DIR_REMOVED=""
if [[ -d $EXO_DIR ]]; then
if [[ $KEEP_MODELS == "1" && -d "$EXO_DIR/models" ]]; then
find "$EXO_DIR" -mindepth 1 -maxdepth 1 ! -name models -exec rm -rf {} +
EXO_DIR_REMOVED="kept_models"
echo_info "Removed ~/.exo (preserved models/)"
else
rm -rf "$EXO_DIR"
EXO_DIR_REMOVED="full"
echo_info "Removed ~/.exo"
fi
else
echo_warn "~/.exo not found (already removed?)"
fi
# Note about launch at login registration
# SMAppService-based login items cannot be removed from a shell script.
# They can only be unregistered from within the app itself or manually via System Settings.
@@ -144,6 +196,10 @@ echo " • Network setup LaunchDaemon"
echo " • Network configuration script"
echo " • Log files"
echo " • 'exo' network location"
case "$EXO_DIR_REMOVED" in
full) echo " • EXO data directory (~/.exo)" ;;
kept_models) echo " • EXO data directory (~/.exo, models preserved)" ;;
esac
echo ""
echo "Your network has been restored to use the 'Automatic' location."
echo "Thunderbolt Bridge has been re-enabled (if present)."
+77 -22
View File
@@ -7,7 +7,7 @@
# name, patterns, reasoning
#
# Optional per-model overrides (CLI flags take priority over these):
# temperature, top_p, max_tokens, reasoning_effort
# temperature, top_p, max_tokens, reasoning_effort, enable_thinking
#
# Fallback defaults (when no per-model config):
# reasoning: temperature=1.0, max_tokens=131072, reasoning_effort="high"
@@ -18,10 +18,9 @@
# ─── Qwen3.5 (Feb 2026) ─────────────────────────────────────────────
# Source: HuggingFace model cards (Qwen/Qwen3.5-*)
# 35B-A3B thinking general: temp=1.0, top_p=0.95, top_k=20
# 397B thinking: temp=0.6, top_p=0.95, top_k=20
# Non-thinking: temp=0.7, top_p=0.8, top_k=20
# max_tokens: 32768 general, 81920 for complex math/code
# Model card recommends: temp=0.6, top_p=0.95, top_k=20
# We omit top_k to match vllm eval (which doesn't set it).
# max_tokens=121072 to match vllm eval (131072 context - 10000 safety margin).
[[model]]
name = "Qwen3.5 2B"
@@ -29,7 +28,8 @@ patterns = ["Qwen3.5-2B"]
reasoning = true
temperature = 0.6
top_p = 0.95
max_tokens = 81920
enable_thinking = true
max_tokens = 121072
[[model]]
name = "Qwen3.5 9B"
@@ -37,7 +37,8 @@ patterns = ["Qwen3.5-9B"]
reasoning = true
temperature = 0.6
top_p = 0.95
max_tokens = 81920
enable_thinking = true
max_tokens = 121072
[[model]]
name = "Qwen3.5 27B"
@@ -45,15 +46,17 @@ patterns = ["Qwen3.5-27B"]
reasoning = true
temperature = 0.6
top_p = 0.95
max_tokens = 81920
enable_thinking = true
max_tokens = 121072
[[model]]
name = "Qwen3.5 35B A3B"
patterns = ["Qwen3.5-35B-A3B"]
reasoning = true
temperature = 1.0
temperature = 0.6
top_p = 0.95
max_tokens = 81920
enable_thinking = true
max_tokens = 121072
[[model]]
name = "Qwen3.5 122B A10B"
@@ -61,7 +64,8 @@ patterns = ["Qwen3.5-122B-A10B"]
reasoning = true
temperature = 0.6
top_p = 0.95
max_tokens = 81920
enable_thinking = true
max_tokens = 121072
[[model]]
name = "Qwen3.5 397B A17B"
@@ -69,12 +73,14 @@ patterns = ["Qwen3.5-397B-A17B"]
reasoning = true
temperature = 0.6
top_p = 0.95
max_tokens = 81920
enable_thinking = true
max_tokens = 121072
# ─── Qwen3 (Apr 2025) ───────────────────────────────────────────────
# Source: HuggingFace model cards (Qwen/Qwen3-*)
# Thinking: temp=0.6, top_p=0.95, top_k=20
# Non-thinking: temp=0.7, top_p=0.8, top_k=20
# Model card recommends: temp=0.6, top_p=0.95, top_k=20
# We omit top_k to match vllm eval (which doesn't set it).
# Non-thinking: temp=0.7, top_p=0.8
# max_tokens: 32768 general, 38912 for complex math/code
[[model]]
@@ -83,6 +89,7 @@ patterns = ["Qwen3-0.6B"]
reasoning = true
temperature = 0.6
top_p = 0.95
enable_thinking = true
max_tokens = 38912
[[model]]
@@ -91,6 +98,7 @@ patterns = ["Qwen3-30B-A3B"]
reasoning = true
temperature = 0.6
top_p = 0.95
enable_thinking = true
max_tokens = 38912
[[model]]
@@ -99,6 +107,7 @@ patterns = ["Qwen3-235B-A22B"]
reasoning = true
temperature = 0.6
top_p = 0.95
enable_thinking = true
max_tokens = 38912
[[model]]
@@ -107,6 +116,7 @@ patterns = ["Qwen3-Next-80B-A3B-Thinking"]
reasoning = true
temperature = 0.6
top_p = 0.95
enable_thinking = true
max_tokens = 38912
[[model]]
@@ -129,9 +139,9 @@ max_tokens = 16384
name = "Qwen3 Coder Next"
patterns = ["Qwen3-Coder-Next"]
reasoning = false
temperature = 0.7
top_p = 0.8
max_tokens = 16384
temperature = 1.0
top_p = 0.95
max_tokens = 121072
# ─── GPT-OSS (OpenAI) ───────────────────────────────────────────────
# Source: OpenAI GitHub README + HuggingFace discussion #21
@@ -165,10 +175,38 @@ patterns = ["DeepSeek-V3.1"]
reasoning = true
temperature = 0.0
[[model]]
name = "DeepSeek V3.2"
patterns = ["DeepSeek-V3.2"]
reasoning = true
temperature = 1.0
top_p = 0.95
enable_thinking = true
# ─── NVIDIA Nemotron ───────────────────────────────────────────────────
# Source: HuggingFace model cards
# All variants: temp=1.0, top_p=0.95, enable_thinking=true
[[model]]
name = "Nemotron Cascade 2 30B A3B"
patterns = ["Nemotron-Cascade-2-30B-A3B"]
reasoning = true
temperature = 1.0
top_p = 0.95
enable_thinking = true
[[model]]
name = "Nemotron 3 Super 120B A12B"
patterns = ["Nemotron-3-Super-120B-A12B", "NVIDIA-Nemotron-3-Super-120B-A12B"]
reasoning = true
temperature = 1.0
top_p = 0.95
enable_thinking = true
# ─── GLM (ZhipuAI / THUDM) ──────────────────────────────────────────
# Source: HuggingFace model cards + generation_config.json + docs.z.ai
# GLM 4.5+: temp=1.0, top_p=0.95
# Reasoning tasks: 131072 max_tokens; coding/SWE tasks: temp=0.7
# max_tokens=121072 to match vllm eval (131072 context - 10000 safety margin)
[[model]]
name = "GLM-5"
@@ -176,7 +214,8 @@ patterns = ["GLM-5"]
reasoning = true
temperature = 1.0
top_p = 0.95
max_tokens = 131072
enable_thinking = true
max_tokens = 121072
[[model]]
name = "GLM 4.5 Air"
@@ -191,7 +230,8 @@ patterns = ["GLM-4.7-"]
reasoning = true
temperature = 1.0
top_p = 0.95
max_tokens = 131072
enable_thinking = true
max_tokens = 121072
# Note: matches both GLM-4.7 and GLM-4.7-Flash
# ─── Kimi (Moonshot AI) ─────────────────────────────────────────────
@@ -213,7 +253,8 @@ patterns = ["Kimi-K2.5"]
reasoning = true
temperature = 1.0
top_p = 0.95
max_tokens = 131072
enable_thinking = true
max_tokens = 121072
[[model]]
name = "Kimi K2 Instruct"
@@ -223,7 +264,17 @@ temperature = 0.6
# ─── MiniMax ─────────────────────────────────────────────────────────
# Source: HuggingFace model cards + generation_config.json
# All models: temp=1.0, top_p=0.95, top_k=40
# All models: temp=1.0, top_p=0.95
# max_tokens=90000 to match vllm eval (100000 context - 10000 safety margin)
[[model]]
name = "MiniMax M2.7"
patterns = ["MiniMax-M2.7"]
reasoning = true
temperature = 1.0
top_p = 0.95
enable_thinking = true
max_tokens = 90000
[[model]]
name = "MiniMax M2.5"
@@ -231,6 +282,8 @@ patterns = ["MiniMax-M2.5"]
reasoning = true
temperature = 1.0
top_p = 0.95
enable_thinking = true
max_tokens = 90000
[[model]]
name = "MiniMax M2.1"
@@ -251,6 +304,8 @@ patterns = ["Step-3.5-Flash"]
reasoning = true
temperature = 1.0
top_p = 0.95
enable_thinking = true
max_tokens = 121072
# ─── Llama (Meta) ───────────────────────────────────────────────────
# Source: generation_config.json + meta-llama/llama-models generation.py
+189 -96
View File
@@ -35,6 +35,7 @@ from harness import (
ExoHttpError,
add_common_instance_args,
capture_cluster_snapshot,
find_existing_instance,
instance_id_from_instance,
node_ids_from_instance,
nodes_used_in_instance,
@@ -79,7 +80,7 @@ def load_tokenizer_for_bench(model_id: str) -> Any:
model_path = Path(
snapshot_download(
model_id,
allow_patterns=["*.json", "*.py", "*.tiktoken", "*.model"],
allow_patterns=["*.json", "*.py", "*.tiktoken", "*.model", "*.jinja"],
)
)
@@ -277,28 +278,72 @@ def run_one_completion(
prompt_sizer: PromptSizer,
*,
use_prefix_cache: bool = False,
stream: bool = False,
) -> tuple[dict[str, Any], int]:
content, pp_tokens = prompt_sizer.build(pp_hint)
payload: dict[str, Any] = {
"model": model_id,
"messages": [{"role": "user", "content": content}],
"stream": False,
"max_tokens": tg,
"logprobs": False,
"use_prefix_cache": use_prefix_cache,
}
t0 = time.perf_counter()
out = client.post_bench_chat_completions(payload)
elapsed = time.perf_counter() - t0
if not stream:
payload["stream"] = False
t0 = time.perf_counter()
out = client.post_bench_chat_completions(payload)
elapsed = time.perf_counter() - t0
stats = out.get("generation_stats")
stats = out.get("generation_stats")
choices = out.get("choices") or [{}]
message = choices[0].get("message", {}) if choices else {}
content = message.get("content") or ""
preview = content[:200] if content else ""
else:
tokens = 0
first_token_time = None
t0 = time.perf_counter()
text_parts: list[str] = []
stats = None
# Extract preview, handling None content (common for thinking models)
choices = out.get("choices") or [{}]
message = choices[0].get("message", {}) if choices else {}
content = message.get("content") or ""
preview = content[:200] if content else ""
for raw_line in client.stream_bench_chat_completions(payload):
line = raw_line.strip()
if line.startswith(": generation_stats "):
with contextlib.suppress(json.JSONDecodeError):
stats = json.loads(line[len(": generation_stats ") :])
continue
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
break
try:
chunk = json.loads(data)
delta = chunk.get("choices", [{}])[0].get("delta", {})
if delta.get("content"):
if first_token_time is None:
first_token_time = time.perf_counter()
tokens += 1
text_parts.append(delta["content"])
except json.JSONDecodeError:
pass
elapsed = time.perf_counter() - t0
preview = "".join(text_parts)[:200]
if not stats:
ttft = (first_token_time - t0) if first_token_time else elapsed
gen_time = elapsed - ttft if tokens > 1 else elapsed
gen_tps = (tokens - 1) / gen_time if tokens > 1 and gen_time > 0 else 0.0
prompt_tps = pp_tokens / ttft if ttft > 0 else 0.0
stats = {
"prompt_tokens": pp_tokens,
"generation_tokens": tokens,
"prompt_tps": round(prompt_tps, 2),
"generation_tps": round(gen_tps, 2),
"peak_memory_usage": {"inBytes": 0},
}
return {
"elapsed_s": elapsed,
@@ -425,6 +470,11 @@ def main() -> int:
action="store_true",
help="Force all pp×tg combinations (cartesian product) even when lists have equal length.",
)
ap.add_argument(
"--stream",
action="store_true",
help="Use /bench/chat/completions with streaming SSE response (bench=True still applies: no EOS detection, no KV cache).",
)
ap.add_argument(
"--no-system-metrics",
action="store_true",
@@ -490,81 +540,124 @@ def main() -> int:
logger.error("[exo-bench] tokenizer usable but prompt sizing failed")
raise
selected = settle_and_fetch_placements(
client, full_model_id, args, settle_timeout=args.settle_timeout
)
# Optionally reuse a running instance for this model
reused_instance_id: str | None = None
if args.reuse_instance:
existing = find_existing_instance(client, full_model_id)
if existing:
reused_instance_id = existing
logger.info(f"Reusing existing instance {reused_instance_id}")
else:
logger.warning(
"--reuse-instance: no existing instance found, creating a new one"
)
if not selected:
logger.error("No valid placements matched your filters.")
return 1
selected.sort(
key=lambda p: (
str(p.get("instance_meta", "")),
str(p.get("sharding", "")),
-nodes_used_in_instance(p["instance"]),
),
reverse=True,
)
logger.debug(f"exo-bench model: short_id={short_id} full_id={full_model_id}")
logger.info(f"placements: {len(selected)}")
for p in selected:
logger.info(
f" - {p['sharding']} / {p['instance_meta']} / nodes={nodes_used_in_instance(p['instance'])}"
if reused_instance_id is not None:
# Use the existing instance directly — skip placement iteration
selected = []
download_duration_s = None
else:
selected = settle_and_fetch_placements(
client, full_model_id, args, settle_timeout=args.settle_timeout
)
if args.dry_run:
return 0
if not selected:
logger.error("No valid placements matched your filters.")
return 1
settle_deadline = (
time.monotonic() + args.settle_timeout if args.settle_timeout > 0 else None
)
selected.sort(
key=lambda p: (
str(p.get("instance_meta", "")),
str(p.get("sharding", "")),
nodes_used_in_instance(p["instance"]),
),
reverse=True,
)
logger.info("Planning phase: checking downloads...")
download_duration_s = run_planning_phase(
client,
full_model_id,
selected[0],
args.danger_delete_downloads,
args.timeout,
settle_deadline,
)
if download_duration_s is not None:
logger.info(f"Download: {download_duration_s:.1f}s (freshly downloaded)")
else:
logger.info("Download: model already cached")
logger.debug(f"exo-bench model: short_id={short_id} full_id={full_model_id}")
logger.info(f"placements: {len(selected)}")
for p in selected:
logger.info(
f" - {p['sharding']} / {p['instance_meta']} / nodes={nodes_used_in_instance(p['instance'])}"
)
if args.dry_run:
return 0
settle_deadline = (
time.monotonic() + args.settle_timeout if args.settle_timeout > 0 else None
)
logger.info("Planning phase: checking downloads...")
download_duration_s = run_planning_phase(
client,
full_model_id,
selected[0],
args.danger_delete_downloads,
args.timeout,
settle_deadline,
)
if download_duration_s is not None:
logger.info(f"Download: {download_duration_s:.1f}s (freshly downloaded)")
else:
logger.info("Download: model already cached")
cluster_snapshot = capture_cluster_snapshot(client)
all_rows: list[dict[str, Any]] = []
all_system_metrics: dict[str, dict[str, dict[str, float]]] = {}
# If reusing an existing instance, run a single benchmark pass against it
if reused_instance_id is not None:
selected = [None]
for preview in selected:
instance = preview["instance"]
instance_id = instance_id_from_instance(instance)
created_instance = False
if preview is not None:
instance = preview["instance"]
instance_id = instance_id_from_instance(instance)
sharding = str(preview["sharding"])
instance_meta = str(preview["instance_meta"])
n_nodes = nodes_used_in_instance(instance)
sharding = str(preview["sharding"])
instance_meta = str(preview["instance_meta"])
n_nodes = nodes_used_in_instance(instance)
logger.info("=" * 80)
logger.info(
f"PLACEMENT: {sharding} / {instance_meta} / nodes={n_nodes} / instance_id={instance_id}"
)
logger.info("=" * 80)
logger.info(
f"PLACEMENT: {sharding} / {instance_meta} / nodes={n_nodes} / instance_id={instance_id}"
)
client.request_json("POST", "/instance", body={"instance": instance})
try:
wait_for_instance_ready(client, instance_id)
except (RuntimeError, TimeoutError) as e:
logger.error(f"Failed to initialize placement: {e}")
with contextlib.suppress(ExoHttpError):
client.request_json("DELETE", f"/instance/{instance_id}")
continue
# Delete any existing instances to free resources before placing
try:
state = client.request_json("GET", "/state")
for old_id in list(state.get("instances", {}).keys()):
logger.info(f"Deleting stale instance {old_id}")
with contextlib.suppress(ExoHttpError):
client.request_json("DELETE", f"/instance/{old_id}")
if state.get("instances"):
time.sleep(2)
except Exception as e:
logger.warning(f"Failed to clean up stale instances: {e}")
time.sleep(1)
client.request_json("POST", "/instance", body={"instance": instance})
try:
wait_for_instance_ready(client, instance_id)
except (RuntimeError, TimeoutError) as e:
logger.error(f"Failed to initialize placement: {e}")
with contextlib.suppress(ExoHttpError):
client.request_json("DELETE", f"/instance/{instance_id}")
continue
time.sleep(1)
created_instance = True
else:
instance_id = reused_instance_id
sharding = "reused"
instance_meta = "reused"
n_nodes = 0
logger.info("=" * 80)
logger.info(f"Using existing instance {instance_id}")
sampler: SystemMetricsSampler | None = None
if not args.no_system_metrics:
if not args.no_system_metrics and preview is not None:
nids = node_ids_from_instance(instance)
sampler = SystemMetricsSampler(
ExoClient(args.host, args.port, timeout_s=30),
@@ -573,16 +666,20 @@ def main() -> int:
)
sampler.start()
def _do_one(c: ExoClient, pp: int, tg: int) -> tuple[dict[str, Any], int]:
return run_one_completion(
c,
full_model_id,
pp,
tg,
prompt_sizer,
use_prefix_cache=args.use_prefix_cache,
stream=args.stream,
)
try:
for i in range(args.warmup):
run_one_completion(
client,
full_model_id,
pp_list[0],
tg_list[0],
prompt_sizer,
use_prefix_cache=args.use_prefix_cache,
)
_do_one(client, pp_list[0], tg_list[0])
logger.debug(f" warmup {i + 1}/{args.warmup} done")
# If pp and tg lists have same length, run in tandem (zip)
@@ -604,14 +701,7 @@ def main() -> int:
# Sequential: single request
try:
inf_t0 = time.monotonic()
row, actual_pp_tokens = run_one_completion(
client,
full_model_id,
pp,
tg,
prompt_sizer,
use_prefix_cache=args.use_prefix_cache,
)
row, actual_pp_tokens = _do_one(client, pp, tg)
inference_windows.append((inf_t0, time.monotonic()))
except Exception as e:
logger.error(e)
@@ -760,10 +850,12 @@ def main() -> int:
gen_tps = per_req_tps * concurrency
ptok = mean(x["stats"]["prompt_tokens"] for x in runs)
gtok = mean(x["stats"]["generation_tokens"] for x in runs)
peak = mean(
x["stats"]["peak_memory_usage"]["inBytes"] for x in runs
)
def _peak_bytes(s: dict[str, Any]) -> float:
pm = s["peak_memory_usage"]
return pm.get("inBytes") or pm.get("in_bytes", 0)
peak = mean(_peak_bytes(x["stats"]) for x in runs)
summary = (
f"prompt_tps={prompt_tps:.2f} gen_tps={gen_tps:.2f} "
f"prompt_tokens={ptok} gen_tokens={gtok} "
@@ -788,15 +880,16 @@ def main() -> int:
if placement_metrics:
all_system_metrics.update(placement_metrics)
try:
client.request_json("DELETE", f"/instance/{instance_id}")
except ExoHttpError as e:
if e.status != 404:
raise
wait_for_instance_gone(client, instance_id)
logger.debug(f"Deleted instance {instance_id}")
if created_instance and instance_id is not None:
try:
client.request_json("DELETE", f"/instance/{instance_id}")
except ExoHttpError as e:
if e.status != 404:
raise
wait_for_instance_gone(client, instance_id)
logger.debug(f"Deleted instance {instance_id}")
time.sleep(5)
time.sleep(5)
output: dict[str, Any] = {"runs": all_rows}
if cluster_snapshot:
+427 -56
View File
@@ -47,6 +47,7 @@ from harness import (
ExoHttpError,
add_common_instance_args,
capture_cluster_snapshot,
find_existing_instance,
instance_id_from_instance,
nodes_used_in_instance,
resolve_model_short_id,
@@ -62,6 +63,15 @@ from loguru import logger
# ---------------------------------------------------------------------------
MAX_RETRIES = 30
INSTANCE_HEALTH_CHECK_AFTER = (
3 # Check instance health after this many consecutive failures
)
class InstanceFailedError(RuntimeError):
"""Raised when the exo instance is detected as failed/gone."""
DEFAULT_MAX_TOKENS = 16_384
REASONING_MAX_TOKENS = 131_072
TEMPERATURE_NON_REASONING = 0.0
@@ -271,7 +281,7 @@ def run_humaneval_test(
@dataclass
class QuestionResult:
question_id: int
question_id: int | str
prompt: str
response: str
extracted_answer: str | None
@@ -281,7 +291,11 @@ class QuestionResult:
prompt_tokens: int = 0
completion_tokens: int = 0
reasoning_tokens: int = 0
reasoning_content: str = ""
finish_reason: str = ""
elapsed_s: float = 0.0
power_watts: float = 0.0
energy_joules: float = 0.0
@dataclass
@@ -517,6 +531,10 @@ class ApiResult:
prompt_tokens: int
completion_tokens: int
reasoning_tokens: int
reasoning_content: str = ""
finish_reason: str = ""
power_watts: float = 0.0
energy_joules: float = 0.0
async def _call_api(
@@ -530,6 +548,9 @@ async def _call_api(
system_message: str | None = None,
reasoning_effort: str | None = None,
top_p: float | None = None,
top_k: int | None = None,
min_p: float | None = None,
enable_thinking: bool | None = None,
) -> ApiResult:
messages = []
if system_message:
@@ -546,6 +567,12 @@ async def _call_api(
body["reasoning_effort"] = reasoning_effort
if top_p is not None:
body["top_p"] = top_p
if top_k is not None:
body["top_k"] = top_k
if min_p is not None:
body["min_p"] = min_p
if enable_thinking is not None:
body["enable_thinking"] = enable_thinking
resp = await client.post(
f"{base_url}/v1/chat/completions",
@@ -554,19 +581,40 @@ async def _call_api(
)
resp.raise_for_status()
data = resp.json()
content = data["choices"][0]["message"]["content"]
if not content or not content.strip():
choice = data["choices"][0]
message = choice["message"]
content = message.get("content") or ""
reasoning_content = message.get("reasoning_content") or ""
finish_reason = choice.get("finish_reason") or ""
# For thinking models, empty content is expected when finish_reason is "length"
if not content.strip() and finish_reason != "length" and not reasoning_content:
raise ValueError("Empty response from model")
usage = data.get("usage", {})
details = usage.get("completion_tokens_details", {})
power = data.get("power_usage") or {}
return ApiResult(
content=content,
prompt_tokens=usage.get("prompt_tokens", 0),
completion_tokens=usage.get("completion_tokens", 0),
reasoning_tokens=details.get("reasoning_tokens", 0) if details else 0,
reasoning_content=reasoning_content,
finish_reason=finish_reason,
power_watts=power.get("total_avg_sys_power_watts", 0.0),
energy_joules=power.get("total_energy_joules", 0.0),
)
async def _check_instance_health(base_url: str) -> bool:
"""Return True if the exo instance is still reachable."""
try:
async with httpx.AsyncClient() as c:
resp = await c.get(f"{base_url}/models", timeout=5.0)
return resp.status_code == 200
except Exception:
return False
async def call_with_retries(
client: httpx.AsyncClient,
base_url: str,
@@ -578,8 +626,14 @@ async def call_with_retries(
system_message: str | None = None,
reasoning_effort: str | None = None,
top_p: float | None = None,
top_k: int | None = None,
min_p: float | None = None,
enable_thinking: bool | None = None,
instance_failed: asyncio.Event | None = None,
) -> ApiResult | None:
for attempt in range(MAX_RETRIES):
if instance_failed and instance_failed.is_set():
raise InstanceFailedError("Instance already marked as failed")
try:
return await _call_api(
client,
@@ -592,8 +646,30 @@ async def call_with_retries(
system_message,
reasoning_effort,
top_p,
top_k,
min_p,
enable_thinking,
)
except Exception as e:
is_conn_error = isinstance(
e,
(
httpx.ConnectError,
httpx.RemoteProtocolError,
ConnectionRefusedError,
OSError,
),
)
if (
is_conn_error
and attempt >= INSTANCE_HEALTH_CHECK_AFTER
and not await _check_instance_health(base_url)
):
if instance_failed:
instance_failed.set()
raise InstanceFailedError(
f"Instance is down after {attempt + 1} failures: {e}"
) from e
if attempt < MAX_RETRIES - 1:
wait = min(2**attempt, 60)
logger.warning(
@@ -618,10 +694,16 @@ async def evaluate_benchmark(
max_tokens: int,
concurrency: int = 1,
limit: int | None = None,
offset: int = 0,
timeout: float | None = None,
reasoning_effort: str | None = None,
top_p: float | None = None,
top_k: int | None = None,
min_p: float | None = None,
enable_thinking: bool | None = None,
difficulty: str | None = None,
checkpoint_path: Path | None = None,
release_version: str | None = None,
) -> list[QuestionResult]:
"""Run a benchmark. Returns per-question results."""
import datasets
@@ -652,7 +734,21 @@ async def evaluate_benchmark(
ds = ds.filter(lambda x: x["difficulty"] == difficulty)
logger.info(f"Filtered to {len(ds)} {difficulty} problems")
if release_version and "release_version" in ds.column_names:
ds = ds.filter(lambda x: x["release_version"] == release_version)
logger.info(
f"Filtered to {len(ds)} problems with release_version={release_version}"
)
# Sort by question_id to match LCB runner ordering (scenario_router.py:60).
# This ensures [offset:offset+limit] slices select the same problems as vllm.
if "question_id" in ds.column_names:
ds = ds.sort("question_id")
total = len(ds)
if offset > 0:
ds = ds.select(range(min(offset, total), total))
total = len(ds)
if limit and limit < total:
ds = ds.select(range(limit))
total = limit
@@ -660,6 +756,13 @@ async def evaluate_benchmark(
logger.info(
f"Evaluating {benchmark_name}: {total} questions, concurrency={concurrency}, "
f"temperature={temperature}, max_tokens={max_tokens}"
+ (f", top_k={top_k}" if top_k is not None else "")
+ (f", min_p={min_p}" if min_p is not None else "")
+ (
f", enable_thinking={enable_thinking}"
if enable_thinking is not None
else ""
)
)
if config.kind == "code":
@@ -667,16 +770,64 @@ async def evaluate_benchmark(
"Code benchmarks execute model-generated code. Use a sandboxed environment."
)
# Load checkpoint for resume
checkpoint_data: dict[str | int, dict[str, Any]] = {}
if checkpoint_path and checkpoint_path.exists():
with open(checkpoint_path) as f:
for line in f:
entry = json.loads(line)
checkpoint_data[entry["question_id"]] = entry
logger.info(f"Loaded {len(checkpoint_data)} checkpointed results")
semaphore = asyncio.Semaphore(concurrency)
instance_failed = asyncio.Event()
results: list[QuestionResult | None] = [None] * total
completed = 0
lock = asyncio.Lock()
def _get_question_id(idx: int, doc: dict) -> str | int:
"""Get a stable question ID for checkpointing."""
if benchmark_name == "livecodebench":
return doc.get("question_id", idx)
elif benchmark_name == "humaneval":
return doc.get("task_id", idx)
return idx
async def process_question(
idx: int, doc: dict, http_client: httpx.AsyncClient
) -> None:
nonlocal completed
system_msg = None
question_id = _get_question_id(idx, doc)
# Bail out early if instance is already dead
if instance_failed.is_set():
return
# Check checkpoint
if question_id in checkpoint_data:
cached = checkpoint_data[question_id]
results[idx] = QuestionResult(
question_id=question_id,
prompt=cached.get("prompt", ""),
response=cached.get("response", ""),
extracted_answer=cached.get("extracted_answer"),
gold_answer=cached.get("gold_answer", ""),
correct=cached.get("correct", False),
error=cached.get("error"),
prompt_tokens=cached.get("prompt_tokens", 0),
completion_tokens=cached.get("completion_tokens", 0),
reasoning_tokens=cached.get("reasoning_tokens", 0),
reasoning_content=cached.get("reasoning_content", ""),
finish_reason=cached.get("finish_reason", ""),
elapsed_s=cached.get("elapsed_s", 0.0),
power_watts=cached.get("power_watts", 0.0),
energy_joules=cached.get("energy_joules", 0.0),
)
async with lock:
completed += 1
logger.info(f" [{completed}/{total}] {question_id} (cached)")
return
if benchmark_name == "gpqa_diamond":
prompt, gold = format_gpqa_question(doc, idx)
@@ -697,24 +848,50 @@ async def evaluate_benchmark(
raise ValueError(f"Unknown benchmark: {benchmark_name}")
async with semaphore:
if instance_failed.is_set():
return
t0 = time.monotonic()
api_result = await call_with_retries(
http_client,
base_url,
model,
prompt,
temperature,
max_tokens,
timeout,
system_message=system_msg,
reasoning_effort=reasoning_effort,
top_p=top_p,
)
try:
# Race the API call against the instance_failed event
api_task = asyncio.create_task(
call_with_retries(
http_client,
base_url,
model,
prompt,
temperature,
max_tokens,
timeout,
system_message=system_msg,
reasoning_effort=reasoning_effort,
top_p=top_p,
top_k=top_k,
min_p=min_p,
enable_thinking=enable_thinking,
instance_failed=instance_failed,
)
)
failed_waiter = asyncio.create_task(instance_failed.wait())
done, pending = await asyncio.wait(
[api_task, failed_waiter],
return_when=asyncio.FIRST_COMPLETED,
)
for p in pending:
p.cancel()
with contextlib.suppress(asyncio.CancelledError):
await p
if instance_failed.is_set() and api_task not in done:
logger.error(f"Instance failed, aborting {question_id}")
return
api_result = api_task.result()
except InstanceFailedError:
logger.error(f"Instance failed, skipping {question_id}")
return
elapsed = time.monotonic() - t0
if api_result is None:
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response="",
extracted_answer=None,
@@ -729,13 +906,17 @@ async def evaluate_benchmark(
"prompt_tokens": api_result.prompt_tokens,
"completion_tokens": api_result.completion_tokens,
"reasoning_tokens": api_result.reasoning_tokens,
"reasoning_content": api_result.reasoning_content,
"finish_reason": api_result.finish_reason,
"elapsed_s": elapsed,
"power_watts": api_result.power_watts,
"energy_joules": api_result.energy_joules,
}
if config.kind == "mc":
extracted = extract_mc_answer(response, valid_letters)
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response=response,
extracted_answer=extracted,
@@ -749,7 +930,7 @@ async def evaluate_benchmark(
check_aime_answer(extracted, int(gold)) if extracted else False
)
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response=response,
extracted_answer=extracted,
@@ -763,7 +944,7 @@ async def evaluate_benchmark(
code = extract_code_block(response, preserve_indent=keep_indent)
if code is None:
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response=response,
extracted_answer=None,
@@ -778,7 +959,7 @@ async def evaluate_benchmark(
code,
)
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response=response,
extracted_answer="pass" if passed else "fail",
@@ -793,7 +974,7 @@ async def evaluate_benchmark(
exec_meta["sample"],
)
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response=response,
extracted_answer="pass" if passed else "fail",
@@ -804,7 +985,7 @@ async def evaluate_benchmark(
)
else:
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response=response,
extracted_answer=None,
@@ -815,7 +996,7 @@ async def evaluate_benchmark(
)
else:
result = QuestionResult(
question_id=idx,
question_id=question_id,
prompt=prompt,
response=response,
extracted_answer=None,
@@ -827,24 +1008,82 @@ async def evaluate_benchmark(
results[idx] = result
# Write checkpoint (skip infra failures so they get retried on resume,
# but keep wrong answers — they are legitimate results)
if checkpoint_path is not None and result.response:
_write_checkpoint(checkpoint_path, result)
async with lock:
completed += 1
n = completed
if n % max(1, total // 20) == 0 or n == total:
correct_so_far = sum(1 for r in results if r is not None and r.correct)
answered = sum(1 for r in results if r is not None)
logger.info(
f" [{n}/{total}] {correct_so_far}/{answered} correct "
f"({correct_so_far / max(answered, 1):.1%})"
)
# Log progress
thinking_info = ""
if result.reasoning_content:
thinking_info = f", {len(result.reasoning_content)} chars thinking"
logger.info(
f" [{n}/{total}] {question_id}: {len(result.response)} chars{thinking_info}, "
f"tokens: {result.prompt_tokens}+{result.completion_tokens} "
f"[{result.finish_reason}]"
+ (f" {result.extracted_answer}" if result.extracted_answer else "")
)
async def _health_monitor() -> None:
"""Periodically check if the instance is still alive."""
# Wait a bit before first check to let things start
await asyncio.sleep(10)
while not instance_failed.is_set():
if not await _check_instance_health(base_url):
# Double-check to avoid false positives
await asyncio.sleep(2)
if not await _check_instance_health(base_url):
logger.error("Health monitor: instance is down!")
instance_failed.set()
return
await asyncio.sleep(5)
async with httpx.AsyncClient() as http_client:
monitor = asyncio.create_task(_health_monitor())
tasks = [process_question(i, doc, http_client) for i, doc in enumerate(ds)]
await asyncio.gather(*tasks)
monitor.cancel()
with contextlib.suppress(asyncio.CancelledError):
await monitor
if instance_failed.is_set():
completed_count = sum(1 for r in results if r is not None)
logger.error(
f"Instance failed! Completed {completed_count}/{total} problems. "
f"Checkpoint saved — restart to resume remaining problems."
)
raise InstanceFailedError("Instance failed during evaluation")
return [r for r in results if r is not None]
def _write_checkpoint(path: Path, result: QuestionResult) -> None:
"""Append a single result to the JSONL checkpoint file."""
entry = {
"question_id": result.question_id,
"prompt": result.prompt,
"response": result.response,
"extracted_answer": result.extracted_answer,
"gold_answer": result.gold_answer,
"correct": result.correct,
"error": result.error,
"prompt_tokens": result.prompt_tokens,
"completion_tokens": result.completion_tokens,
"reasoning_tokens": result.reasoning_tokens,
"reasoning_content": result.reasoning_content,
"finish_reason": result.finish_reason,
"elapsed_s": round(result.elapsed_s, 2),
"power_watts": round(result.power_watts, 2),
"energy_joules": round(result.energy_joules, 2),
}
with open(path, "a") as f:
f.write(json.dumps(entry) + "\n")
# ---------------------------------------------------------------------------
# Results display
# ---------------------------------------------------------------------------
@@ -867,6 +1106,8 @@ def print_results(
total_elapsed = sum(r.elapsed_s for r in results)
wall_clock = max(r.elapsed_s for r in results) if results else 0.0
avg_gen_tps = total_completion_tokens / total_elapsed if total_elapsed > 0 else 0.0
total_energy = sum(r.energy_joules for r in results)
avg_power = sum(r.power_watts for r in results) / max(total, 1)
label = f"[c={concurrency}] " if concurrency is not None else ""
print(f"\n{label}{benchmark_name}: {correct}/{total} ({accuracy:.1%})")
@@ -878,6 +1119,10 @@ def print_results(
f" | total time: {total_elapsed:.1f}s wall clock: {wall_clock:.1f}s"
)
print(tok_line)
if total_energy > 0:
print(
f" power: avg {avg_power:.1f}W | total energy: {total_energy:.1f}J ({total_energy / 3600:.2f}Wh)"
)
if errors:
print(f" API errors: {errors}")
if no_extract:
@@ -896,6 +1141,8 @@ def print_results(
"total_elapsed_s": total_elapsed,
"wall_clock_s": wall_clock,
"avg_gen_tps": avg_gen_tps,
"avg_power_watts": avg_power,
"total_energy_joules": total_energy,
}
@@ -1053,7 +1300,11 @@ def save_results(
"prompt_tokens": r.prompt_tokens,
"completion_tokens": r.completion_tokens,
"reasoning_tokens": r.reasoning_tokens,
"reasoning_content": r.reasoning_content,
"finish_reason": r.finish_reason,
"elapsed_s": round(r.elapsed_s, 2),
"power_watts": round(r.power_watts, 2),
"energy_joules": round(r.energy_joules, 2),
}
for r in results
],
@@ -1069,6 +1320,15 @@ def save_results(
# ---------------------------------------------------------------------------
def _checkpoint_path(
results_dir: str, benchmark: str, model: str, concurrency: int
) -> Path:
"""Return the JSONL checkpoint path for a benchmark run."""
out_dir = Path(results_dir) / model.replace("/", "_") / benchmark
out_dir.mkdir(parents=True, exist_ok=True)
return out_dir / f"c{concurrency}.checkpoint.jsonl"
def parse_int_list(values: list[str]) -> list[int]:
items: list[int] = []
for v in values:
@@ -1096,6 +1356,12 @@ def main() -> int:
default=None,
help="Max questions per benchmark (for fast iteration).",
)
ap.add_argument(
"--offset",
type=int,
default=0,
help="Skip first N questions (0-based).",
)
reasoning_group = ap.add_mutually_exclusive_group()
reasoning_group.add_argument(
@@ -1115,6 +1381,8 @@ def main() -> int:
"--temperature", type=float, default=None, help="Override temperature."
)
ap.add_argument("--top-p", type=float, default=None, help="Override top_p.")
ap.add_argument("--top-k", type=int, default=None, help="Override top_k.")
ap.add_argument("--min-p", type=float, default=None, help="Override min_p.")
ap.add_argument(
"--max-tokens", type=int, default=None, help="Override max output tokens."
)
@@ -1148,15 +1416,31 @@ def main() -> int:
choices=["easy", "medium", "hard"],
help="Filter by difficulty (livecodebench only). E.g. --difficulty hard",
)
ap.add_argument(
"--release-version",
default=None,
help="LCB dataset release version (livecodebench only). E.g. release_v5",
)
ap.add_argument(
"--results-dir",
default="eval_results",
help="Directory for result JSON files (default: eval_results).",
)
ap.add_argument(
"--skip-instance-setup",
"--enable-thinking",
type=lambda v: v.lower() in ("true", "1", "yes"),
default=None,
help="Enable thinking mode for models that support it.",
)
ap.add_argument(
"--force",
action="store_true",
help="Skip exo instance management (assumes model is already running).",
help="Discard any existing checkpoint and run from scratch.",
)
ap.add_argument(
"--keep-instance",
action="store_true",
help="Skip deleting the instance after eval (for chaining runs).",
)
args, _ = ap.parse_known_args()
@@ -1177,13 +1461,26 @@ def main() -> int:
# Instance management
client = ExoClient(args.host, args.port, timeout_s=args.timeout)
instance_id: str | None = None
created_instance = False
if not args.skip_instance_setup:
short_id, full_model_id = resolve_model_short_id(
client,
args.model,
force_download=args.force_download,
)
_short_id, full_model_id = resolve_model_short_id(
client,
args.model,
force_download=args.force_download,
)
# Optionally reuse a running instance for this model
if args.reuse_instance:
existing = find_existing_instance(client, full_model_id)
if existing:
instance_id = existing
logger.info(f"Reusing existing instance {instance_id}")
else:
logger.warning(
"--reuse-instance: no existing instance found, creating a new one"
)
if instance_id is None:
selected = settle_and_fetch_placements(
client,
full_model_id,
@@ -1198,7 +1495,7 @@ def main() -> int:
key=lambda p: (
str(p.get("instance_meta", "")),
str(p.get("sharding", "")),
-nodes_used_in_instance(p["instance"]),
nodes_used_in_instance(p["instance"]),
),
reverse=True,
)
@@ -1225,6 +1522,18 @@ def main() -> int:
if download_duration is not None:
logger.info(f"Download: {download_duration:.1f}s")
# Delete any existing instances to free resources before placing
try:
state = client.request_json("GET", "/state")
for old_id in list(state.get("instances", {}).keys()):
logger.info(f"Deleting stale instance {old_id}")
with contextlib.suppress(ExoHttpError):
client.request_json("DELETE", f"/instance/{old_id}")
if state.get("instances"):
time.sleep(2)
except Exception as e:
logger.warning(f"Failed to clean up stale instances: {e}")
client.request_json("POST", "/instance", body={"instance": instance})
try:
wait_for_instance_ready(client, instance_id)
@@ -1234,10 +1543,9 @@ def main() -> int:
client.request_json("DELETE", f"/instance/{instance_id}")
return 1
time.sleep(1)
cluster_snapshot = capture_cluster_snapshot(client)
else:
full_model_id = args.model
cluster_snapshot = None
created_instance = True
cluster_snapshot = capture_cluster_snapshot(client)
# Auto-detect reasoning from model config
model_config = load_model_config(full_model_id)
@@ -1291,16 +1599,57 @@ def main() -> int:
reasoning_effort = str(cfg["reasoning_effort"])
else:
reasoning_effort = "high" if is_reasoning else None
if args.top_k is not None:
top_k: int | None = args.top_k
elif "top_k" in cfg:
top_k = int(cfg["top_k"])
else:
top_k = None
if args.min_p is not None:
min_p: float | None = args.min_p
elif "min_p" in cfg:
min_p = float(cfg["min_p"])
else:
min_p = None
if args.enable_thinking is not None:
enable_thinking: bool | None = args.enable_thinking
elif "enable_thinking" in cfg:
enable_thinking = bool(cfg["enable_thinking"])
else:
enable_thinking = None
base_url = f"http://{args.host}:{args.port}"
logger.info(f"Model: {full_model_id}")
logger.info(
f"Settings: temperature={temperature}, max_tokens={max_tokens}, "
+ (f"top_p={top_p}, " if top_p is not None else "")
+ (f"top_k={top_k}, " if top_k is not None else "")
+ (f"min_p={min_p}, " if min_p is not None else "")
+ f"reasoning={'yes' if is_reasoning else 'no'}"
+ (f", reasoning_effort={reasoning_effort}" if reasoning_effort else "")
+ (
f", enable_thinking={enable_thinking}"
if enable_thinking is not None
else ""
)
)
# Common kwargs for evaluate_benchmark
eval_kwargs: dict[str, Any] = {
"reasoning_effort": reasoning_effort,
"top_p": top_p,
"top_k": top_k,
"min_p": min_p,
"enable_thinking": enable_thinking,
"difficulty": args.difficulty,
"offset": args.offset,
"release_version": args.release_version,
}
try:
if args.compare_concurrency:
concurrency_levels = parse_int_list(args.compare_concurrency)
@@ -1309,6 +1658,11 @@ def main() -> int:
for c in concurrency_levels:
logger.info(f"\n{'=' * 50}")
logger.info(f"Running {task_name} at concurrency={c}")
checkpoint_path = _checkpoint_path(
args.results_dir, task_name, full_model_id, c
)
if args.force and checkpoint_path.exists():
checkpoint_path.unlink()
results = asyncio.run(
evaluate_benchmark(
task_name,
@@ -1319,9 +1673,8 @@ def main() -> int:
concurrency=c,
limit=args.limit,
timeout=args.request_timeout,
reasoning_effort=reasoning_effort,
top_p=top_p,
difficulty=args.difficulty,
checkpoint_path=checkpoint_path,
**eval_kwargs,
)
)
if results:
@@ -1336,10 +1689,18 @@ def main() -> int:
cluster=cluster_snapshot,
)
results_by_c[c] = results
# Clean up checkpoint on success
if checkpoint_path.exists():
checkpoint_path.unlink()
if len(results_by_c) >= 2:
print_comparison(task_name, results_by_c)
else:
for task_name in task_names:
checkpoint_path = _checkpoint_path(
args.results_dir, task_name, full_model_id, args.num_concurrent
)
if args.force and checkpoint_path.exists():
checkpoint_path.unlink()
results = asyncio.run(
evaluate_benchmark(
task_name,
@@ -1350,9 +1711,8 @@ def main() -> int:
concurrency=args.num_concurrent,
limit=args.limit,
timeout=args.request_timeout,
reasoning_effort=reasoning_effort,
top_p=top_p,
difficulty=args.difficulty,
checkpoint_path=checkpoint_path,
**eval_kwargs,
)
)
if results:
@@ -1366,14 +1726,25 @@ def main() -> int:
scores,
cluster=cluster_snapshot,
)
# Clean up checkpoint on success
if checkpoint_path.exists():
checkpoint_path.unlink()
finally:
if instance_id is not None:
try:
client.request_json("DELETE", f"/instance/{instance_id}")
except ExoHttpError as e:
if e.status != 404:
raise
wait_for_instance_gone(client, instance_id)
if created_instance and instance_id is not None:
if args.keep_instance:
logger.info(f"Keeping instance {instance_id} (--keep-instance)")
else:
try:
client.request_json("DELETE", f"/instance/{instance_id}")
except ExoHttpError as e:
if e.status != 404:
raise
try:
wait_for_instance_gone(client, instance_id)
except TimeoutError:
logger.warning(
f"Timed out waiting for instance {instance_id} to be deleted"
)
return 0
+65 -12
View File
@@ -6,6 +6,7 @@ import http.client
import json
import os
import time
from collections.abc import Iterator
from typing import Any
from urllib.parse import urlencode
@@ -69,6 +70,30 @@ class ExoClient:
def post_bench_chat_completions(self, payload: dict[str, Any]) -> dict[str, Any]:
return self.request_json("POST", "/bench/chat/completions", body=payload)
def stream_bench_chat_completions(self, payload: dict[str, Any]) -> Iterator[str]:
"""POST /bench/chat/completions with stream=True, yielding raw SSE lines."""
payload = {**payload, "stream": True}
data = json.dumps(payload).encode("utf-8")
conn = http.client.HTTPConnection(self.host, self.port, timeout=self.timeout_s)
try:
conn.request(
"POST",
"/bench/chat/completions",
body=data,
headers={
"Content-Type": "application/json",
"Accept": "text/event-stream",
},
)
resp = conn.getresponse()
if resp.status >= 400:
raw = resp.read().decode("utf-8", errors="replace")
raise ExoHttpError(resp.status, resp.reason, raw[:300])
for line in resp:
yield line.decode("utf-8", errors="replace")
finally:
conn.close()
def get_state_path(self, path: str) -> Any:
try:
return self.request_json("GET", f"/state/{path}")
@@ -268,11 +293,15 @@ def sharding_filter(sharding: str, wanted: str) -> bool:
def fetch_and_filter_placements(
client: ExoClient, full_model_id: str, args: argparse.Namespace
client: ExoClient,
full_model_id: str,
args: argparse.Namespace,
node_id: str | None = None,
) -> list[dict[str, Any]]:
previews_resp = client.request_json(
"GET", "/instance/previews", params={"model_id": full_model_id}
)
params: dict[str, str] = {"model_id": full_model_id}
if node_id is not None:
params["node_ids"] = node_id
previews_resp = client.request_json("GET", "/instance/previews", params=params)
previews = previews_resp.get("previews") or []
selected: list[dict[str, Any]] = []
@@ -332,8 +361,9 @@ def settle_and_fetch_placements(
full_model_id: str,
args: argparse.Namespace,
settle_timeout: float = 0,
node_id: str | None = None,
) -> list[dict[str, Any]]:
selected = fetch_and_filter_placements(client, full_model_id, args)
selected = fetch_and_filter_placements(client, full_model_id, args, node_id=node_id)
if not selected and settle_timeout > 0:
backoff = _SETTLE_INITIAL_BACKOFF_S
@@ -346,7 +376,9 @@ def settle_and_fetch_placements(
)
time.sleep(min(backoff, remaining))
backoff = min(backoff * _SETTLE_BACKOFF_MULTIPLIER, _SETTLE_MAX_BACKOFF_S)
selected = fetch_and_filter_placements(client, full_model_id, args)
selected = fetch_and_filter_placements(
client, full_model_id, args, node_id=node_id
)
return selected
@@ -462,9 +494,8 @@ def run_planning_phase(
)
logger.info(f"Started download on {node_id}")
# Wait for downloads
start = time.time()
while time.time() - start < timeout:
# Wait for downloads (no timeout — poll until complete or failed)
while True:
all_done = True
for node_id in node_ids:
node_downloads = client.get_node_downloads(node_id) or []
@@ -514,9 +545,24 @@ def run_planning_phase(
if download_t0 is not None:
return time.perf_counter() - download_t0
return None
time.sleep(1)
time.sleep(10)
raise TimeoutError("Downloads did not complete in time")
def find_existing_instance(client: ExoClient, model_id: str) -> str | None:
"""Find an existing running instance for the given model."""
try:
state = client.request_json("GET", "/state")
except Exception:
return None
for inst_id, inst in state.get("instances", {}).items():
# Instance structure is nested: {"MlxJacclInstance": {"shardAssignments": {"modelId": ...}}}
for _inst_type, inner in inst.items():
if not isinstance(inner, dict):
continue
sa = inner.get("shardAssignments", {})
if sa.get("modelId") == model_id:
return inst_id
return None
def add_common_instance_args(ap: argparse.ArgumentParser) -> None:
@@ -543,7 +589,9 @@ def add_common_instance_args(ap: argparse.ArgumentParser) -> None:
help="Only consider placements using >= this many nodes.",
)
ap.add_argument(
"--instance-meta", choices=["ring", "jaccl", "both"], default="both"
"--instance-meta",
choices=["ring", "jaccl", "vllm", "both"],
default="both",
)
ap.add_argument(
"--sharding", choices=["pipeline", "tensor", "both"], default="both"
@@ -572,3 +620,8 @@ def add_common_instance_args(ap: argparse.ArgumentParser) -> None:
action="store_true",
help="Delete existing models from smallest to largest to make room for benchmark model.",
)
ap.add_argument(
"--reuse-instance",
action="store_true",
help="Reuse an existing running instance for this model instead of creating a new one.",
)
+36
View File
@@ -0,0 +1,36 @@
# Prefill/Decode disaggregation benchmark config.
#
# Top-level keys are bench-wide. [prefill] and [decode] sections set per-side
# placement filters and (optionally) per-side model.
#
# Example:
# uv run python bench/prefill_decode_bench.py --config bench/prefill-decode.toml
host = "james"
port = 52415
timeout = 7200.0
settle_timeout = 60.0
# Workload
pp = [4096, 8192]
tg = [128]
repeat = 1
warmup = 0
json_out = "bench/prefill_decode_results.json"
[prefill]
model = "sakamakismile/Qwen3.6-27B-NVFP4"
node = "gx10-de89"
instance_meta = "vllm"
sharding = "pipeline"
min_nodes = 1
max_nodes = 1
[decode]
model = "mlx-community/Qwen3.6-27B-4bit"
node = "Ryuichis MacBook Pro"
instance_meta = "ring"
sharding = "pipeline"
min_nodes = 1
max_nodes = 1
+869
View File
@@ -0,0 +1,869 @@
# type: ignore
#!/usr/bin/env python3
"""Disaggregated prefill-decode benchmark for exo (MLX → MLX).
Spins up two MLX instances on the cluster, marks one as Prefill source and
the other as Decode target via /v1/instance-links, then sends chat
completions to the API. The master routes the request to the decode
instance and stamps `prefill_endpoint` pointing at the prefill instance —
the worker decides per-request whether to ship prefill remotely
(uncached_count > REMOTE_PREFILL_MIN_TOKENS).
Usage:
uv run python bench/prefill_decode_bench.py --model <id> --pp 2048,8192 --tg 128
uv run python bench/prefill_decode_bench.py --model <id> --pp 4096 --tg 128 --repeat 3
uv run python bench/prefill_decode_bench.py --model <id> --pp 2048 --tg 128 --dry-run
"""
from __future__ import annotations
import argparse
import contextlib
import copy
import itertools
import json
import sys
import time
import tomllib
from pathlib import Path
from statistics import mean
from typing import Any
from exo_bench import (
PromptSizer,
SystemMetricsSampler,
format_peak_memory,
load_tokenizer_for_bench,
parse_int_list,
)
from harness import (
ExoClient,
ExoHttpError,
add_common_instance_args,
instance_id_from_instance,
node_ids_from_instance,
nodes_used_in_instance,
resolve_model_short_id,
run_planning_phase,
settle_and_fetch_placements,
unwrap_instance,
wait_for_instance_gone,
wait_for_instance_ready,
)
from loguru import logger
def _node_id_to_friendly(client: ExoClient) -> dict[str, str]:
identities = client.get_node_identities() or {}
out: dict[str, str] = {}
for node_id, identity in identities.items():
if isinstance(identity, dict):
name = identity.get("friendlyName") or identity.get("friendly_name")
if isinstance(name, str):
out[str(node_id)] = name
return out
def _placement_node_friendly_names(
placement: dict[str, Any], id_to_friendly: dict[str, str]
) -> list[str]:
instance = placement["instance"]
return [id_to_friendly.get(nid, nid) for nid in node_ids_from_instance(instance)]
def _filter_by_node(
placements: list[dict[str, Any]],
friendly_name: str,
id_to_friendly: dict[str, str],
) -> list[dict[str, Any]]:
target = friendly_name.lower()
matched: list[dict[str, Any]] = []
for p in placements:
names = [n.lower() for n in _placement_node_friendly_names(p, id_to_friendly)]
if any(target == n or target in n for n in names):
matched.append(p)
return matched
def _node_id_by_friendly(id_to_friendly: dict[str, str], target: str) -> str | None:
target_lc = target.lower()
for nid, name in id_to_friendly.items():
if target_lc == name.lower() or target_lc in name.lower():
return nid
return None
def _load_toml(path: str) -> dict[str, Any]:
with Path(path).open("rb") as f:
return tomllib.load(f)
_TOP_LEVEL_TOML_KEYS = {
"host",
"port",
"timeout",
"settle_timeout",
"model",
"pp",
"tg",
"repeat",
"warmup",
"json_out",
"instance_meta",
"sharding",
"min_nodes",
"max_nodes",
"force_download",
"danger_delete_downloads",
"all_combinations",
}
def _inject_toml_into_argv() -> None:
"""If --config X is in sys.argv, pre-load it and inject required CLI args
(--model, --pp, --tg) so argparse's required=True checks pass."""
argv = sys.argv
if "--config" not in argv:
return
idx = argv.index("--config")
if idx + 1 >= len(argv):
return
cfg_path = argv[idx + 1]
cfg = _load_toml(cfg_path)
decode = cfg.get("decode", {})
def _has(flag: str) -> bool:
return any(a == flag or a.startswith(flag + "=") for a in argv)
# --model: prefer top-level, then [decode].model
if not _has("--model"):
model = cfg.get("model") or decode.get("model")
if model:
argv += ["--model", str(model)]
if not _has("--pp"):
pp = cfg.get("pp")
if pp:
argv += (
["--pp", *(str(x) for x in pp)]
if isinstance(pp, list)
else [
"--pp",
str(pp),
]
)
if not _has("--tg"):
tg = cfg.get("tg")
if tg:
argv += (
["--tg", *(str(x) for x in tg)]
if isinstance(tg, list)
else [
"--tg",
str(tg),
]
)
def _merge_toml_into_args(args: argparse.Namespace, cfg: dict[str, Any]) -> None:
"""Apply top-level toml keys onto args namespace where args has a default."""
for key, value in cfg.items():
if key in {"prefill", "decode"}:
continue
if key not in _TOP_LEVEL_TOML_KEYS:
continue
attr = key
current = getattr(args, attr, None)
if current in (None, [], False):
setattr(args, attr, value)
def _side_args(
base: argparse.Namespace, overrides: dict[str, Any]
) -> argparse.Namespace:
out = copy.copy(base)
for k in (
"instance_meta",
"sharding",
"min_nodes",
"max_nodes",
"skip_pipeline_jaccl",
"skip_tensor_ring",
):
if k in overrides:
setattr(out, k, overrides[k])
return out
def _pick_two_distinct_placements(
placements: list[dict[str, Any]],
) -> tuple[dict[str, Any], dict[str, Any]] | None:
if len(placements) < 2:
return None
seen_nodes: set[tuple[str, ...]] = set()
chosen: list[dict[str, Any]] = []
for p in placements:
nodes = tuple(sorted(str(n) for n in p.get("nodes", [])))
if nodes in seen_nodes:
continue
seen_nodes.add(nodes)
chosen.append(p)
if len(chosen) == 2:
return chosen[0], chosen[1]
return None
def _create_instance_link(
client: ExoClient,
prefill_instance_id: str,
decode_instance_id: str,
) -> str:
out = client.request_json(
"POST",
"/v1/instance-links",
body={
"prefill_instances": [prefill_instance_id],
"decode_instances": [decode_instance_id],
},
)
return str(out.get("commandId", ""))
def _list_instance_links(client: ExoClient) -> list[dict[str, Any]]:
out = client.request_json("GET", "/v1/instance-links")
return out if isinstance(out, list) else []
def _delete_instance_link(client: ExoClient, link_id: str) -> None:
client.request_json("DELETE", f"/v1/instance-links/{link_id}")
def run_one(
client: ExoClient,
model_id: str,
pp_hint: int,
tg: int,
prompt_sizer: PromptSizer,
) -> tuple[dict[str, Any], int]:
content, pp_tokens = prompt_sizer.build(pp_hint)
payload: dict[str, Any] = {
"model": model_id,
"messages": [{"role": "user", "content": content}],
"stream": False,
"max_tokens": tg,
}
t0 = time.perf_counter()
out = client.post_bench_chat_completions(payload)
elapsed = time.perf_counter() - t0
stats = out.get("generation_stats")
choices = out.get("choices") or [{}]
message = choices[0].get("message", {}) if choices else {}
text = message.get("content") or ""
preview = text[:200] if text else ""
return {
"elapsed_s": elapsed,
"output_text_preview": preview,
"stats": stats,
}, pp_tokens
def _run_phase(
*,
client: ExoClient,
label: str,
pp_tg_pairs: list[tuple[int, int]],
model_id: str,
prompt_sizer: PromptSizer,
warmup: int,
repeat: int,
common_meta: dict[str, Any],
sampler: SystemMetricsSampler | None = None,
) -> list[dict[str, Any]]:
logger.info(f"=== phase: {label} (model={model_id}) ===")
rows: list[dict[str, Any]] = []
for i in range(warmup):
run_one(client, model_id, pp_tg_pairs[0][0], pp_tg_pairs[0][1], prompt_sizer)
logger.debug(f" warmup {i + 1}/{warmup} done")
for pp, tg in pp_tg_pairs:
logger.info(f"--- {label}: pp={pp} tg={tg} ---")
runs: list[dict[str, Any]] = []
inference_windows: list[tuple[float, float]] = []
for r in range(repeat):
time.sleep(2)
try:
inf_t0 = time.monotonic()
row, actual_pp_tokens = run_one(client, model_id, pp, tg, prompt_sizer)
inference_windows.append((inf_t0, time.monotonic()))
except Exception as e:
logger.error(e)
continue
row.update(common_meta)
row.update(
{
"phase": label,
"phase_model_id": model_id,
"pp_tokens": actual_pp_tokens,
"tg": tg,
"repeat_index": r,
}
)
runs.append(row)
rows.append(row)
if runs:
prompt_tps = mean(x["stats"]["prompt_tps"] for x in runs)
gen_tps = mean(x["stats"]["generation_tps"] for x in runs)
ptok = mean(x["stats"]["prompt_tokens"] for x in runs)
gtok = mean(x["stats"]["generation_tokens"] for x in runs)
peak = mean(x["stats"]["peak_memory_usage"]["inBytes"] for x in runs)
avg_elapsed = mean(x["elapsed_s"] for x in runs)
energy_str = ""
if sampler is not None and inference_windows:
joules = sum(
sampler.energy_between(t0, t1) for t0, t1 in inference_windows
)
inf_seconds = sum(t1 - t0 for t0, t1 in inference_windows)
avg_watts = joules / inf_seconds if inf_seconds > 0 else 0.0
energy_per_run = joules / len(runs) if runs else 0.0
energy_str = (
f" energy={joules:.1f}J ({avg_watts:.1f}W avg over "
f"{inf_seconds:.1f}s inference, {energy_per_run:.1f}J/run)"
)
for run_row, (t0, t1) in zip(runs, inference_windows, strict=False):
run_row["energy_joules"] = sampler.energy_between(t0, t1)
run_row["inference_window_s"] = t1 - t0
logger.info(
f"[{label}] prompt_tps={prompt_tps:.2f} gen_tps={gen_tps:.2f} "
f"prompt_tokens={ptok} gen_tokens={gtok} "
f"peak_memory={format_peak_memory(peak)} "
f"avg_elapsed={avg_elapsed:.2f}s{energy_str}"
)
time.sleep(2)
return rows
def _summarise(rows: list[dict[str, Any]]) -> dict[tuple[int, int], dict[str, float]]:
grouped: dict[tuple[int, int], list[dict[str, Any]]] = {}
for r in rows:
key = (int(r["pp_tokens"]), int(r["tg"]))
grouped.setdefault(key, []).append(r)
out: dict[tuple[int, int], dict[str, float]] = {}
for key, runs in grouped.items():
energy_runs = [x.get("energy_joules") for x in runs if "energy_joules" in x]
window_runs = [
x.get("inference_window_s") for x in runs if "inference_window_s" in x
]
out[key] = {
"prompt_tps": mean(x["stats"]["prompt_tps"] for x in runs),
"gen_tps": mean(x["stats"]["generation_tps"] for x in runs),
"elapsed_s": mean(x["elapsed_s"] for x in runs),
"prompt_tokens": mean(x["stats"]["prompt_tokens"] for x in runs),
"gen_tokens": mean(x["stats"]["generation_tokens"] for x in runs),
"energy_j": mean(energy_runs) if energy_runs else 0.0,
"inference_window_s": mean(window_runs) if window_runs else 0.0,
}
return out
def _normalised_seconds(summary: dict[str, float], pp: int, tg: int) -> float | None:
"""Wall-clock time implied by reported tps for the *configured* pp/tg.
elapsed_s is not comparable across phases when models EOS at different
lengths. This formula reconstructs "what would this phase take to do
pp prompt tokens + tg generation tokens" using its own reported rates.
"""
p_tps = summary.get("prompt_tps", 0.0)
g_tps = summary.get("gen_tps", 0.0)
if p_tps <= 0 or g_tps <= 0:
return None
return pp / p_tps + tg / g_tps
def _print_diff(
disagg_rows: list[dict[str, Any]],
decode_alone_rows: list[dict[str, Any]],
prefill_alone_rows: list[dict[str, Any]],
) -> None:
disagg = _summarise(disagg_rows)
decode_alone = _summarise(decode_alone_rows)
prefill_alone = _summarise(prefill_alone_rows)
keys = set(disagg.keys()) | set(decode_alone.keys()) | set(prefill_alone.keys())
width = 110
for key in sorted(keys):
pp, tg = key
logger.info("" * width)
logger.info(f" pp={pp} tg={tg}")
logger.info("" * width)
logger.info(
f" {'phase':<16} {'elapsed':>9} {'norm':>9} "
f"{'prompt_tps':>11} {'gen_tps':>8} "
f"{'p_tok':>6} {'g_tok':>6} "
f"{'energy':>9} {'avg_W':>7}"
)
for label, summary in (
("disaggregated", disagg.get(key)),
("decode_alone", decode_alone.get(key)),
("prefill_alone", prefill_alone.get(key)),
):
if summary is None:
logger.info(
f" {label:<16} {'':>9} {'':>9} "
f"{'':>11} {'':>8} {'':>6} {'':>6} "
f"{'':>9} {'':>7}"
)
continue
norm = _normalised_seconds(summary, pp, tg)
norm_str = f"{norm:>8.2f}s" if norm is not None else f"{'':>9}"
energy = summary.get("energy_j", 0.0)
window = summary.get("inference_window_s", 0.0)
energy_str = f"{energy:>8.1f}J" if energy > 0 else f"{'':>9}"
avg_w = energy / window if window > 0 else 0.0
avg_w_str = f"{avg_w:>6.1f}W" if avg_w > 0 else f"{'':>7}"
logger.info(
f" {label:<16} "
f"{summary['elapsed_s']:>8.2f}s "
f"{norm_str} "
f"{summary['prompt_tps']:>11.1f} "
f"{summary['gen_tps']:>8.2f} "
f"{summary['prompt_tokens']:>6.0f} "
f"{summary['gen_tokens']:>6.0f} "
f"{energy_str} "
f"{avg_w_str}"
)
d = disagg.get(key)
da = decode_alone.get(key)
pa = prefill_alone.get(key)
d_norm = _normalised_seconds(d, pp, tg) if d else None
if d_norm and da:
da_norm = _normalised_seconds(da, pp, tg)
if da_norm:
logger.info(
f" norm speedup vs decode_alone: {da_norm / d_norm:.2f}x "
f"(prefill {d['prompt_tps'] / da['prompt_tps']:.2f}x, "
f"decode {d['gen_tps'] / da['gen_tps']:.2f}x)"
)
if d_norm and pa:
pa_norm = _normalised_seconds(pa, pp, tg)
if pa_norm:
logger.info(
f" norm speedup vs prefill_alone: {pa_norm / d_norm:.2f}x "
f"(prefill {d['prompt_tps'] / pa['prompt_tps']:.2f}x, "
f"decode {d['gen_tps'] / pa['gen_tps']:.2f}x)"
)
logger.info("" * width)
def main() -> int:
_inject_toml_into_argv()
ap = argparse.ArgumentParser(
prog="prefill-decode-bench",
description="Benchmark MLX-MLX disaggregated prefill/decode via instance links.",
)
add_common_instance_args(ap)
ap.add_argument(
"--pp",
nargs="+",
required=True,
help="Prompt-size hints (ints, must be >1000). Accepts commas.",
)
ap.add_argument(
"--tg",
nargs="+",
required=True,
help="Generation lengths (ints). Accepts commas.",
)
ap.add_argument(
"--repeat", type=int, default=1, help="Repetitions per (pp,tg) pair."
)
ap.add_argument(
"--warmup",
type=int,
default=0,
help="Warmup runs (uses first pp/tg).",
)
ap.add_argument(
"--json-out",
default="bench/prefill_decode_results.json",
help="Write raw per-run results JSON to this path.",
)
ap.add_argument("--stdout", action="store_true", help="Write results to stdout")
ap.add_argument(
"--dry-run", action="store_true", help="List selected placements and exit."
)
ap.add_argument(
"--all-combinations",
action="store_true",
help="Force all pp×tg combinations even when lists have equal length.",
)
ap.add_argument(
"--prefill-model",
default=None,
help="Model id for the prefill instance. Defaults to --model.",
)
ap.add_argument(
"--prefill-node",
default=None,
help="friendly_name of the node hosting the prefill instance.",
)
ap.add_argument(
"--decode-node",
default=None,
help="friendly_name of the node hosting the decode instance.",
)
ap.add_argument(
"--config",
default=None,
help="TOML config file. CLI flags override toml values.",
)
ap.add_argument(
"--compare-baseline",
action="store_true",
help="Also run each (pp,tg) pair without the prefill/decode link "
"(decode instance does its own prefill) and report the diff.",
)
args = ap.parse_args()
cfg = _load_toml(args.config) if args.config else {}
_merge_toml_into_args(args, cfg)
prefill_overrides = cfg.get("prefill", {}) if cfg else {}
decode_overrides = cfg.get("decode", {}) if cfg else {}
if args.prefill_model is None and "model" in prefill_overrides:
args.prefill_model = prefill_overrides["model"]
if args.prefill_node is None and "node" in prefill_overrides:
args.prefill_node = prefill_overrides["node"]
if args.decode_node is None and "node" in decode_overrides:
args.decode_node = decode_overrides["node"]
if "model" in decode_overrides and not args.model:
args.model = decode_overrides["model"]
pp_list = parse_int_list(args.pp)
tg_list = parse_int_list(args.tg)
if not pp_list or not tg_list:
logger.error("pp and tg lists must be non-empty")
return 2
for pp in pp_list:
if pp <= 1000:
logger.error(
f"pp={pp} must be >1000 (remote prefill triggers when uncached >1000)"
)
return 2
if args.repeat <= 0:
logger.error("--repeat must be >= 1")
return 2
use_combinations = args.all_combinations or len(pp_list) != len(tg_list)
if use_combinations:
logger.info(
f"pp/tg mode: combinations (product) — {len(pp_list) * len(tg_list)} pairs"
)
else:
logger.info(f"pp/tg mode: tandem (zip) — {len(pp_list)} pairs")
client = ExoClient(args.host, args.port, timeout_s=args.timeout)
decode_short_id, decode_full_id = resolve_model_short_id(
client, args.model, force_download=args.force_download
)
if args.prefill_model:
prefill_short_id, prefill_full_id = resolve_model_short_id(
client, args.prefill_model, force_download=args.force_download
)
else:
prefill_short_id, prefill_full_id = decode_short_id, decode_full_id
tokenizer = load_tokenizer_for_bench(decode_full_id)
if tokenizer is None:
raise RuntimeError("[prefill-decode-bench] decode tokenizer load failed")
try:
decode_prompt_sizer = PromptSizer(tokenizer)
except Exception:
logger.error("[prefill-decode-bench] decode prompt sizing failed")
raise
if prefill_full_id == decode_full_id:
prefill_prompt_sizer = decode_prompt_sizer
else:
prefill_tokenizer = load_tokenizer_for_bench(prefill_full_id)
if prefill_tokenizer is None:
raise RuntimeError("[prefill-decode-bench] prefill tokenizer load failed")
prefill_prompt_sizer = PromptSizer(prefill_tokenizer)
id_to_friendly = _node_id_to_friendly(client)
prefill_args = _side_args(args, prefill_overrides)
decode_args = _side_args(args, decode_overrides)
if prefill_full_id == decode_full_id and prefill_overrides == decode_overrides:
placements = settle_and_fetch_placements(
client, decode_full_id, args, settle_timeout=args.settle_timeout
)
prefill_candidates = (
_filter_by_node(placements, args.prefill_node, id_to_friendly)
if args.prefill_node
else placements
)
decode_candidates = (
_filter_by_node(placements, args.decode_node, id_to_friendly)
if args.decode_node
else placements
)
if args.prefill_node and not prefill_candidates:
logger.error(f"No placement on prefill node {args.prefill_node!r}.")
return 1
if args.decode_node and not decode_candidates:
logger.error(f"No placement on decode node {args.decode_node!r}.")
return 1
if args.prefill_node and args.decode_node:
prefill_p = prefill_candidates[0]
decode_p = decode_candidates[0]
else:
pair = _pick_two_distinct_placements(placements)
if pair is None:
logger.error(
"Need at least two distinct-node MLX placements for the same model."
)
return 1
prefill_p, decode_p = pair
if args.prefill_node:
prefill_p = prefill_candidates[0]
if args.decode_node:
decode_p = decode_candidates[0]
else:
prefill_node_id = (
_node_id_by_friendly(id_to_friendly, args.prefill_node)
if args.prefill_node
else None
)
decode_node_id = (
_node_id_by_friendly(id_to_friendly, args.decode_node)
if args.decode_node
else None
)
if args.prefill_node and prefill_node_id is None:
logger.error(f"Unknown node {args.prefill_node!r}.")
return 1
if args.decode_node and decode_node_id is None:
logger.error(f"Unknown node {args.decode_node!r}.")
return 1
prefill_placements = settle_and_fetch_placements(
client,
prefill_full_id,
prefill_args,
settle_timeout=args.settle_timeout,
node_id=prefill_node_id,
)
decode_placements = settle_and_fetch_placements(
client,
decode_full_id,
decode_args,
settle_timeout=args.settle_timeout,
node_id=decode_node_id,
)
if not prefill_placements:
logger.error(
f"No placement found for prefill model {prefill_full_id}"
f"{f' on node {args.prefill_node!r}' if args.prefill_node else ''}."
)
return 1
if not decode_placements:
logger.error(
f"No placement found for decode model {decode_full_id}"
f"{f' on node {args.decode_node!r}' if args.decode_node else ''}."
)
return 1
prefill_p = prefill_placements[0]
decode_p = decode_placements[0]
prefill_node_names = _placement_node_friendly_names(prefill_p, id_to_friendly)
decode_node_names = _placement_node_friendly_names(decode_p, id_to_friendly)
_ = unwrap_instance
prefill_instance = prefill_p["instance"]
decode_instance = decode_p["instance"]
prefill_id = instance_id_from_instance(prefill_instance)
decode_id = instance_id_from_instance(decode_instance)
prefill_meta = str(prefill_p.get("instance_meta", ""))
decode_meta = str(decode_p.get("instance_meta", ""))
prefill_nodes = nodes_used_in_instance(prefill_instance)
decode_nodes = nodes_used_in_instance(decode_instance)
logger.info("=" * 80)
logger.info(
f"PREFILL: {prefill_meta} / nodes={prefill_nodes} ({','.join(prefill_node_names)}) "
f"/ {prefill_short_id} ({prefill_full_id}) / instance_id={prefill_id}"
)
logger.info(
f"DECODE: {decode_meta} / nodes={decode_nodes} ({','.join(decode_node_names)}) "
f"/ {decode_short_id} ({decode_full_id}) / instance_id={decode_id}"
)
if args.dry_run:
return 0
settle_deadline = (
time.monotonic() + args.settle_timeout if args.settle_timeout > 0 else None
)
logger.info("Planning phase: prefill...")
run_planning_phase(
client,
prefill_full_id,
prefill_p,
args.danger_delete_downloads,
args.timeout,
settle_deadline,
)
logger.info("Planning phase: decode...")
run_planning_phase(
client,
decode_full_id,
decode_p,
args.danger_delete_downloads,
args.timeout,
settle_deadline,
)
if use_combinations:
pp_tg_pairs = list(itertools.product(pp_list, tg_list))
else:
pp_tg_pairs = list(zip(pp_list, tg_list, strict=True))
common_meta = {
"decode_model_short_id": decode_short_id,
"decode_model_id": decode_full_id,
"prefill_model_short_id": prefill_short_id,
"prefill_model_id": prefill_full_id,
"prefill_instance_id": prefill_id,
"prefill_instance_meta": prefill_meta,
"prefill_nodes": prefill_nodes,
"decode_instance_id": decode_id,
"decode_instance_meta": decode_meta,
"decode_nodes": decode_nodes,
}
all_rows: list[dict[str, Any]] = []
disagg_rows: list[dict[str, Any]] = []
decode_alone_rows: list[dict[str, Any]] = []
prefill_alone_rows: list[dict[str, Any]] = []
link_id = ""
prefill_alive = False
decode_alive = False
sampler_nodes = sorted(
{
*node_ids_from_instance(prefill_instance),
*node_ids_from_instance(decode_instance),
}
)
sampler = SystemMetricsSampler(
ExoClient(args.host, args.port, timeout_s=30), sampler_nodes
)
sampler.start()
try:
logger.info("Creating prefill instance...")
client.request_json("POST", "/instance", body={"instance": prefill_instance})
wait_for_instance_ready(client, prefill_id)
prefill_alive = True
logger.info("Prefill instance ready")
if args.compare_baseline:
time.sleep(2)
prefill_alone_rows = _run_phase(
client=client,
label="prefill_alone",
pp_tg_pairs=pp_tg_pairs,
model_id=prefill_full_id,
prompt_sizer=prefill_prompt_sizer,
warmup=args.warmup,
repeat=args.repeat,
common_meta=common_meta,
sampler=sampler,
)
all_rows.extend(prefill_alone_rows)
logger.info("Creating decode instance...")
client.request_json("POST", "/instance", body={"instance": decode_instance})
wait_for_instance_ready(client, decode_id)
decode_alive = True
logger.info("Decode instance ready")
logger.info("Linking instances (prefill → decode)...")
_create_instance_link(client, prefill_id, decode_id)
time.sleep(1)
links = _list_instance_links(client)
if not links:
logger.error("Link did not appear in state.")
return 1
link_id = str(links[-1].get("linkId") or links[-1].get("link_id") or "")
logger.info(f"Link created: {link_id}")
time.sleep(2)
disagg_rows = _run_phase(
client=client,
label="disaggregated",
pp_tg_pairs=pp_tg_pairs,
model_id=decode_full_id,
prompt_sizer=decode_prompt_sizer,
warmup=args.warmup,
repeat=args.repeat,
common_meta=common_meta,
sampler=sampler,
)
all_rows.extend(disagg_rows)
if args.compare_baseline:
logger.info("Removing link and prefill instance to isolate decode_alone.")
with contextlib.suppress(ExoHttpError):
if link_id:
_delete_instance_link(client, link_id)
link_id = ""
with contextlib.suppress(ExoHttpError):
client.request_json("DELETE", f"/instance/{prefill_id}")
wait_for_instance_gone(client, prefill_id)
prefill_alive = False
time.sleep(2)
decode_alone_rows = _run_phase(
client=client,
label="decode_alone",
pp_tg_pairs=pp_tg_pairs,
model_id=decode_full_id,
prompt_sizer=decode_prompt_sizer,
warmup=args.warmup,
repeat=args.repeat,
common_meta=common_meta,
sampler=sampler,
)
all_rows.extend(decode_alone_rows)
_print_diff(disagg_rows, decode_alone_rows, prefill_alone_rows)
finally:
sampler.stop()
with contextlib.suppress(ExoHttpError):
if link_id:
_delete_instance_link(client, link_id)
if decode_alive:
with contextlib.suppress(ExoHttpError):
client.request_json("DELETE", f"/instance/{decode_id}")
wait_for_instance_gone(client, decode_id)
if prefill_alive:
with contextlib.suppress(ExoHttpError):
client.request_json("DELETE", f"/instance/{prefill_id}")
wait_for_instance_gone(client, prefill_id)
logger.debug("Deleted both instances")
if args.stdout:
json.dump(all_rows, sys.stdout, indent=2, ensure_ascii=False)
elif args.json_out:
with open(args.json_out, "w", encoding="utf-8") as f:
json.dump(all_rows, f, indent=2, ensure_ascii=False)
logger.debug(f"\nWrote results JSON: {args.json_out}")
return 0
if __name__ == "__main__":
sys.exit(main())
@@ -202,6 +202,7 @@
let instanceType: string | null = null;
if (instanceTag === "MlxRingInstance") instanceType = "MLX Ring";
else if (instanceTag === "MlxJacclInstance") instanceType = "MLX RDMA";
else if (instanceTag === "VllmInstance") instanceType = "vLLM";
let sharding: string | null = null;
const inst = instance as {
+292 -2
View File
@@ -9,7 +9,7 @@
*/
interface Props {
/** "macbook pro" | "mac studio" | "mac mini" etc. */
/** "macbook pro" | "mac studio" | "mac mini" | "dgx spark" | "linux" etc. */
deviceType: string;
/** Center X coordinate in SVG space */
cx: number;
@@ -38,10 +38,43 @@
const LOGO_NATIVE_WIDTH = 814;
const LOGO_NATIVE_HEIGHT = 1000;
// NVIDIA logo SVG path
const NVIDIA_LOGO_PATH =
"M0.81 0.429V0.299c0.013 -0.001 0.026 -0.002 0.038 -0.002 0.355 -0.011 0.588 0.306 0.588 0.306S1.186 0.952 0.916 0.952c-0.036 0 -0.071 -0.006 -0.105 -0.017V0.542c0.138 0.017 0.166 0.078 0.249 0.216l0.185 -0.155s-0.135 -0.177 -0.362 -0.177c-0.024 -0.001 -0.048 0.001 -0.072 0.003m0 -0.429v0.194l0.038 -0.002c0.494 -0.017 0.816 0.405 0.816 0.405s-0.37 0.45 -0.754 0.45c-0.034 0 -0.066 -0.003 -0.099 -0.009v0.12c0.027 0.003 0.055 0.006 0.082 0.006 0.358 0 0.618 -0.183 0.869 -0.399 0.042 0.034 0.212 0.114 0.247 0.15 -0.238 0.2 -0.794 0.361 -1.11 0.361 -0.03 0 -0.059 -0.002 -0.088 -0.005v0.169h1.362V0zm0 0.935v0.102c-0.331 -0.059 -0.423 -0.404 -0.423 -0.404s0.159 -0.176 0.423 -0.205v0.112h-0.001C0.671 0.524 0.562 0.654 0.562 0.654s0.062 0.218 0.248 0.282m-0.588 -0.316s0.196 -0.29 0.589 -0.32V0.194C0.376 0.229 0 0.597 0 0.597s0.213 0.616 0.81 0.672v-0.112c-0.438 -0.054 -0.588 -0.538 -0.588 -0.538";
const wireColor = "rgba(179,179,179,0.8)";
const strokeWidth = 1.5;
const modelLower = $derived(deviceType.toLowerCase());
const isSpark = $derived(
modelLower.includes("dgx") || modelLower.includes("gx10"),
);
const isLinux = $derived(!isSpark && modelLower.startsWith("linux"));
const isLinuxLaptop = $derived(isLinux && modelLower.includes("laptop"));
// ── DGX Spark dimensions ──
const dgxW = $derived(size * 1.55);
const dgxH = $derived(size * 0.58);
const dgxX = $derived(cx - dgxW / 2);
const dgxY = $derived(cy - dgxH / 2);
const dgxChassisX = $derived(dgxX - dgxW * 0.03);
const dgxChassisW = $derived(dgxW * 1.05);
const dgxHandleW = $derived(dgxW * 0.27);
const dgxHandleGap = $derived(dgxH * 0.05);
const dgxHandleH = $derived(dgxH - dgxHandleGap * 2);
const dgxHandleY = $derived(dgxY + dgxHandleGap);
const dgxInnerHandleW = $derived(dgxW * 0.12);
const dgxInnerHandleH = $derived(dgxHandleH - dgxH * 0.06);
const dgxLeftHandleX = $derived(dgxX + 4);
const dgxRightHandleX = $derived(dgxX + dgxW - dgxHandleW - 4);
const dgxClipId = $derived(`di-dgx-${uid}`);
const dgxTextureId = $derived(`di-dgx-tex-${uid}`);
// ── Linux Desktop dimensions (reuses Mac Studio proportions) ──
const linuxDesktopClipId = $derived(`di-linux-desktop-${uid}`);
// ── Linux Laptop dimensions (reuses MacBook proportions) ──
const linuxScreenClipId = $derived(`di-linux-screen-${uid}`);
// ── Mac Studio dimensions (same ratios as TopologyGraph) ──
const studioW = $derived(size * 1.25);
@@ -114,7 +147,264 @@
const studioClipId = $derived(`di-studio-${uid}`);
</script>
{#if modelLower === "mac studio" || modelLower === "mac mini"}
{#if isSpark}
<!-- DGX Spark -->
<defs>
<clipPath id={dgxClipId}>
<rect x={dgxX} y={dgxY} width={dgxW} height={dgxH} rx="3" />
</clipPath>
<pattern
id={dgxTextureId}
patternUnits="userSpaceOnUse"
width="8"
height="8"
>
<rect width="8" height="8" fill="#6f6248" />
<circle cx="2" cy="2" r="1" fill="#5a4f3b" opacity="0.5" />
<circle cx="6" cy="6" r="1" fill="#4a4232" opacity="0.45" />
</pattern>
</defs>
<!-- Main body -->
<rect
x={dgxChassisX}
y={dgxY}
width={dgxChassisW}
height={dgxH}
rx="3"
fill="url(#{dgxTextureId})"
stroke={wireColor}
stroke-width={strokeWidth}
/>
<!-- Side border accents -->
<rect
x={dgxChassisX}
y={dgxY}
width={dgxW * 0.02}
height={dgxH}
fill="#8a7a56"
/>
<rect
x={dgxChassisX + dgxChassisW - dgxW * 0.02}
y={dgxY}
width={dgxW * 0.02}
height={dgxH}
fill="#8a7a56"
/>
<!-- Memory fill -->
{#if ramPercent > 0}
<rect
x={dgxX}
y={dgxY + dgxH - (ramPercent / 100) * dgxH}
width={dgxW}
height={(ramPercent / 100) * dgxH}
fill="rgba(255,215,0,0.45)"
clip-path="url(#{dgxClipId})"
/>
{/if}
<!-- Left handle -->
<rect
x={dgxLeftHandleX}
y={dgxHandleY}
width={dgxHandleW}
height={dgxHandleH}
rx="2.4"
fill="#b3a170"
stroke="#403723"
stroke-width="0.7"
/>
<rect
x={dgxLeftHandleX + dgxHandleW * 0.06}
y={dgxHandleY + dgxH * 0.03}
width={dgxInnerHandleW}
height={dgxInnerHandleH}
rx="1.6"
fill="#8a7a56"
/>
<!-- Right handle -->
<rect
x={dgxRightHandleX}
y={dgxHandleY}
width={dgxHandleW}
height={dgxHandleH}
rx="2.4"
fill="#b3a170"
stroke="#403723"
stroke-width="0.7"
/>
<rect
x={dgxRightHandleX + dgxHandleW - dgxInnerHandleW - dgxHandleW * 0.08}
y={dgxHandleY + dgxH * 0.03}
width={dgxInnerHandleW}
height={dgxInnerHandleH}
rx="1.6"
fill="#8a7a56"
/>
<!-- NVIDIA logo (rotated 90deg on left handle) -->
{@const badgeW = dgxW * 0.09}
{@const badgeH = dgxHandleH * 0.5}
{@const badgeX = dgxLeftHandleX + dgxHandleW - badgeW - dgxHandleW * 0.06}
{@const badgeYPos = dgxHandleY + (dgxHandleH - badgeH) / 2}
{@const textSz = badgeW * 0.58}
{@const logoW = textSz * 1.2}
{@const logoH = logoW * (1.438 / 2.174)}
{@const ctrX = badgeX + badgeW / 2 - badgeW * 0.03}
{@const ctrY = badgeYPos + badgeH / 2}
{@const labelGap = badgeW * 0.15}
{@const totalW = logoW + labelGap + textSz * 3.6}
<g transform="rotate(90 {ctrX} {ctrY})">
<svg
x={ctrX - totalW / 2}
y={ctrY - logoH / 2}
width={logoW}
height={logoH}
viewBox="0 0 2.174 1.438"
>
<path d={NVIDIA_LOGO_PATH} fill="#76b900" />
</svg>
<text
x={ctrX - totalW / 2 + logoW + labelGap}
y={ctrY}
text-anchor="start"
dominant-baseline="middle"
fill="#8a7a56"
font-size={textSz}
font-family="monospace"
font-weight="700">NVIDIA</text
>
</g>
{:else if isLinuxLaptop}
<!-- Linux Laptop — MacBook shape with Tux logo -->
<defs>
<clipPath id={linuxScreenClipId}>
<rect
x={mbScreenX + mbBezel}
y={mbY + mbBezel}
width={mbScreenW - mbBezel * 2}
height={mbScreenH - mbBezel * 2}
rx="2"
/>
</clipPath>
</defs>
<rect
x={mbScreenX}
y={mbY}
width={mbScreenW}
height={mbScreenH}
rx="3"
fill="#1a1a1a"
stroke={wireColor}
stroke-width={strokeWidth}
/>
<rect
x={mbScreenX + mbBezel}
y={mbY + mbBezel}
width={mbScreenW - mbBezel * 2}
height={mbScreenH - mbBezel * 2}
rx="2"
fill="#0a0a12"
/>
{#if ramPercent > 0}
<rect
x={mbScreenX + mbBezel}
y={mbY + mbBezel + (mbMemTotalH - mbMemH)}
width={mbScreenW - mbBezel * 2}
height={mbMemH}
fill="rgba(255,215,0,0.85)"
clip-path="url(#{linuxScreenClipId})"
/>
{/if}
<!-- Terminal prompt on screen -->
<text
x={cx}
y={mbY + mbScreenH / 2}
text-anchor="middle"
dominant-baseline="middle"
fill="#FFFFFF"
opacity="0.9"
font-size={mbScreenH * 0.25}
font-family="SF Mono, Monaco, monospace"
font-weight="700">{">_"}</text
>
<path
d="M {mbBaseTopX} {mbBaseY} L {mbBaseTopX +
mbBaseTopW} {mbBaseY} L {mbBaseBottomX + mbBaseBottomW} {mbBaseY +
mbBaseH} L {mbBaseBottomX} {mbBaseY + mbBaseH} Z"
fill="#2c2c2c"
stroke={wireColor}
stroke-width="1"
/>
<rect
x={mbKbX}
y={mbKbY}
width={mbKbW}
height={mbKbH}
fill="rgba(0,0,0,0.2)"
rx="2"
/>
<rect
x={mbTpX}
y={mbTpY}
width={mbTpW}
height={mbTpH}
fill="rgba(255,255,255,0.08)"
rx="2"
/>
{:else if isLinux}
<!-- Linux Desktop — Mac Studio shape with Tux logo -->
<defs>
<clipPath id={linuxDesktopClipId}>
<rect
x={studioX}
y={studioY + studioTopH}
width={studioW}
height={studioH - studioTopH}
rx={studioCorner - 1}
/>
</clipPath>
</defs>
<rect
x={studioX}
y={studioY}
width={studioW}
height={studioH}
rx={studioCorner}
fill="#1a1a1a"
stroke={wireColor}
stroke-width={strokeWidth}
/>
{#if ramPercent > 0}
<rect
x={studioX}
y={studioY + studioTopH + (studioMemTotalH - studioMemH)}
width={studioW}
height={studioMemH}
fill="rgba(255,215,0,0.75)"
clip-path="url(#{linuxDesktopClipId})"
/>
{/if}
<!-- Terminal prompt on front face -->
<text
x={cx}
y={studioY + studioTopH + (studioH - studioTopH) / 2}
text-anchor="middle"
dominant-baseline="middle"
fill="rgba(255,255,255,0.5)"
font-size={(studioH - studioTopH) * 0.4}
font-family="SF Mono, Monaco, monospace"
font-weight="700">{">_"}</text
>
{:else if modelLower === "mac studio" || modelLower === "mac mini"}
<!-- Mac Studio / Mac Mini -->
<defs>
<clipPath id={studioClipId}>
@@ -1,5 +1,8 @@
<script lang="ts">
import { browser } from "$app/environment";
import { featureFlags } from "$lib/stores/app.svelte";
const showAdvanced = $derived(featureFlags()["disaggregation"] === true);
interface Props {
showHome?: boolean;
@@ -297,5 +300,28 @@
</svg>
<span class="hidden sm:inline">Integrations</span>
</a>
{#if showAdvanced}
<a
href="/#/advanced"
class="text-xs md:text-sm text-white/70 hover:text-exo-yellow transition-colors tracking-wider uppercase flex items-center gap-1.5 md:gap-2 cursor-pointer"
title="Advanced cluster settings"
>
<svg
class="w-4 h-4"
viewBox="0 0 24 24"
fill="none"
stroke="currentColor"
stroke-width="2"
stroke-linecap="round"
stroke-linejoin="round"
>
<circle cx="12" cy="12" r="3" />
<path
d="M19.4 15a1.65 1.65 0 0 0 .33 1.82l.06.06a2 2 0 0 1 0 2.83 2 2 0 0 1-2.83 0l-.06-.06a1.65 1.65 0 0 0-1.82-.33 1.65 1.65 0 0 0-1 1.51V21a2 2 0 0 1-4 0v-.09A1.65 1.65 0 0 0 9 19.4a1.65 1.65 0 0 0-1.82.33l-.06.06a2 2 0 0 1-2.83 0 2 2 0 0 1 0-2.83l.06-.06a1.65 1.65 0 0 0 .33-1.82 1.65 1.65 0 0 0-1.51-1H3a2 2 0 0 1 0-4h.09A1.65 1.65 0 0 0 4.6 9a1.65 1.65 0 0 0-.33-1.82l-.06-.06a2 2 0 0 1 0-2.83 2 2 0 0 1 2.83 0l.06.06a1.65 1.65 0 0 0 1.82.33H9a1.65 1.65 0 0 0 1-1.51V3a2 2 0 0 1 4 0v.09a1.65 1.65 0 0 0 1 1.51 1.65 1.65 0 0 0 1.82-.33l.06-.06a2 2 0 0 1 2.83 0 2 2 0 0 1 0 2.83l-.06.06a1.65 1.65 0 0 0-.33 1.82V9a1.65 1.65 0 0 0 1.51 1H21a2 2 0 0 1 0 4h-.09a1.65 1.65 0 0 0-1.51 1z"
/>
</svg>
<span class="hidden sm:inline">Advanced</span>
</a>
{/if}
</nav>
</header>
+85 -4
View File
@@ -23,7 +23,7 @@
} | null;
nodes?: Record<string, NodeInfo>;
sharding?: "Pipeline" | "Tensor";
runtime?: "MlxRing" | "MlxJaccl";
runtime?: "MlxRing" | "MlxJaccl" | "Vllm";
onLaunch?: () => void;
tags?: string[];
apiPreview?: PlacementPreview | null;
@@ -168,8 +168,10 @@
function getDeviceType(
name: string,
): "macbook" | "studio" | "mini" | "unknown" {
): "macbook" | "studio" | "mini" | "dgx" | "linux" | "unknown" {
const lower = name.toLowerCase();
if (lower.includes("dgx") || lower.includes("gx10")) return "dgx";
if (lower.includes("linux")) return "linux";
if (lower.includes("macbook")) return "macbook";
if (lower.includes("studio")) return "studio";
if (lower.includes("mini")) return "mini";
@@ -576,13 +578,17 @@
class="px-1.5 py-0.5 text-xs font-mono tracking-wider uppercase bg-exo-medium-gray/30 text-exo-light-gray border border-exo-medium-gray/40"
title={runtime === "MlxRing"
? "Ring: standard networking. Works over any connection (Wi-Fi, Ethernet, Thunderbolt)."
: "RDMA: direct memory access over Thunderbolt. Significantly faster for multi-device inference."}
: runtime === "MlxJaccl"
? "RDMA: direct memory access over Thunderbolt. Significantly faster for multi-device inference."
: "vLLM: NVIDIA CUDA inference engine."}
>
{runtime === "MlxRing"
? "MLX Ring"
: runtime === "MlxJaccl"
? "MLX RDMA"
: runtime}
: runtime === "Vllm"
? "vLLM"
: runtime}
</span>
</div>
@@ -990,6 +996,81 @@
/>
{/if}
</g>
{:else if node.deviceType === "dgx"}
<!-- DGX Spark icon -->
{@const s = node.iconSize}
{@const dgxW = s * 1.4}
{@const dgxH = s * 0.52}
<g transform="translate({-dgxW / 2}, {-dgxH / 2})">
<!-- Chassis -->
<rect
x="0"
y="0"
width={dgxW}
height={dgxH}
rx="2"
fill="#6f6248"
stroke={node.isUsed ? "#FFD700" : "#4B5563"}
stroke-width="1.5"
/>
<!-- Side accents -->
<rect
x="0"
y="0"
width={dgxW * 0.02}
height={dgxH}
fill="#8a7a56"
/>
<rect
x={dgxW - dgxW * 0.02}
y="0"
width={dgxW * 0.02}
height={dgxH}
fill="#8a7a56"
/>
<!-- Left handle -->
<rect
x={dgxW * 0.04}
y={dgxH * 0.08}
width={dgxW * 0.22}
height={dgxH * 0.84}
rx="2"
fill="#b3a170"
stroke="#403723"
stroke-width="0.5"
/>
<!-- Right handle -->
<rect
x={dgxW - dgxW * 0.04 - dgxW * 0.22}
y={dgxH * 0.08}
width={dgxW * 0.22}
height={dgxH * 0.84}
rx="2"
fill="#b3a170"
stroke="#403723"
stroke-width="0.5"
/>
<!-- Memory fill -->
<rect
x="2"
y={dgxH - dgxH * (node.currentPercent / 100)}
width={dgxW - 4}
height={dgxH * (node.currentPercent / 100)}
fill="rgba(255,215,0,0.35)"
/>
{#if node.modelUsageGB > 0 && node.isUsed}
<rect
x="2"
y={dgxH - dgxH * (node.newPercent / 100)}
width={dgxW - 4}
height={dgxH *
((node.newPercent - node.currentPercent) / 100)}
fill="#FFD700"
filter="url(#memGlow-{filterId})"
class="animate-pulse-slow"
/>
{/if}
</g>
{:else}
<!-- Unknown device - hexagon -->
<g
@@ -9,6 +9,7 @@
capabilities?: string[];
family?: string;
is_custom?: boolean;
requires_vllm?: boolean;
}
interface ModelGroup {
@@ -19,6 +20,7 @@
variants: ModelInfo[];
smallestVariant: ModelInfo;
hasMultipleVariants: boolean;
requiresVllm: boolean;
}
type DownloadAvailability = {
@@ -213,6 +215,14 @@
<span class="font-mono text-sm text-white truncate">
{group.name}
</span>
{#if group.requiresVllm}
<span
class="text-[10px] font-mono px-1.5 py-0.5 rounded bg-orange-500/15 text-orange-300 border border-orange-400/30 flex-shrink-0 tracking-wider uppercase"
title="Requires vLLM runtime"
>
vLLM
</span>
{/if}
<!-- Capability icons -->
{#each group.capabilities.filter((c) => c !== "text") as cap}
{#if cap === "thinking"}
@@ -523,6 +533,15 @@
{variant.quantization || "default"}
</span>
{#if variant.requires_vllm}
<span
class="text-[10px] font-mono px-1.5 py-0.5 rounded bg-orange-500/15 text-orange-300 border border-orange-400/30 flex-shrink-0 tracking-wider uppercase"
title="Requires vLLM runtime"
>
vLLM
</span>
{/if}
<!-- Size -->
<span
class="text-xs font-mono flex-1 {getSizeClassForFitStatus(
@@ -628,6 +647,7 @@
variants: [variant],
smallestVariant: variant,
hasMultipleVariants: false,
requiresVllm: variant.requires_vllm === true,
});
}}
title="View variant details"
@@ -22,6 +22,7 @@
is_custom?: boolean;
tasks?: string[];
hugging_face_id?: string;
requires_vllm?: boolean;
}
interface ModelGroup {
@@ -32,6 +33,7 @@
variants: ModelInfo[];
smallestVariant: ModelInfo;
hasMultipleVariants: boolean;
requiresVllm: boolean;
}
interface FilterState {
@@ -396,6 +398,7 @@
variants: [],
smallestVariant: model,
hasMultipleVariants: false,
requiresVllm: true,
});
}
@@ -430,6 +433,7 @@
(a.storage_size_megabytes || 0) - (b.storage_size_megabytes || 0),
);
group.hasMultipleVariants = group.variants.length > 1;
group.requiresVllm = group.variants.every((v) => v.requires_vllm);
}
// Convert to array and sort by smallest variant size (biggest first)
@@ -587,6 +591,7 @@
variants: [model],
smallestVariant: model,
hasMultipleVariants: false,
requiresVllm: model.requires_vllm === true,
});
}
}
@@ -1165,6 +1170,17 @@
<span class="text-white/40">Variants:</span>
<span class="text-white/70">{infoGroup.variants.length}</span>
</div>
{#if infoGroup.requiresVllm}
<div class="flex items-center gap-2">
<span class="text-white/40">Runtime:</span>
<span
class="text-[10px] font-mono px-1.5 py-0.5 rounded bg-orange-500/15 text-orange-300 border border-orange-400/30 tracking-wider uppercase"
>
vLLM
</span>
<span class="text-white/40 text-[11px]">required</span>
</div>
{/if}
{#if infoGroup.variants.length > 0}
<div class="mt-3 pt-3 border-t border-exo-yellow/10">
<span class="text-white/40">Available quantizations:</span>
@@ -0,0 +1,565 @@
<script lang="ts">
import { onMount, onDestroy } from "svelte";
import FamilyLogos from "$lib/components/FamilyLogos.svelte";
import {
instances,
instanceLinks,
nodeIdentities,
refreshState,
createInstanceLink,
updateInstanceLink,
deleteInstanceLink,
type Instance,
} from "$lib/stores/app.svelte";
import { deriveBaseModel, deriveFamily } from "$lib/utils/model_family";
type InstanceWrapper = {
MlxRingInstance?: Instance;
MlxJacclInstance?: Instance;
VllmInstance?: Instance;
};
let interval: ReturnType<typeof setInterval> | null = null;
onMount(() => {
refreshState();
interval = setInterval(refreshState, 3000);
});
onDestroy(() => {
if (interval) clearInterval(interval);
});
type InstanceRow = {
id: string;
modelId: string;
family: string;
baseModel: string;
nodeNames: string[];
nodeCount: number;
};
const instanceRows = $derived.by<InstanceRow[]>(() => {
const rows: InstanceRow[] = [];
const ids = nodeIdentities();
for (const [id, raw] of Object.entries(instances())) {
const wrapper = raw as InstanceWrapper;
const inst =
wrapper.MlxRingInstance ??
wrapper.MlxJacclInstance ??
wrapper.VllmInstance;
const modelId = inst?.shardAssignments?.modelId ?? "";
const nodeToRunner = inst?.shardAssignments?.nodeToRunner ?? {};
const nodeIds = Object.keys(nodeToRunner);
const nodeNames = nodeIds
.map((nodeId) => ids[nodeId]?.friendlyName ?? nodeId.slice(0, 6))
.filter((name) => !!name);
rows.push({
id,
modelId,
family: deriveFamily(modelId),
baseModel: deriveBaseModel(modelId),
nodeNames,
nodeCount: nodeIds.length,
});
}
rows.sort((a, b) => a.modelId.localeCompare(b.modelId));
return rows;
});
const instanceById = $derived(
Object.fromEntries(instanceRows.map((r) => [r.id, r])),
);
type LinkRow = {
linkId: string;
prefill: string[];
decode: string[];
families: string[];
multiNode: boolean;
};
const linkRows = $derived.by<LinkRow[]>(() => {
const rows: LinkRow[] = [];
for (const [, link] of Object.entries(instanceLinks())) {
const fams = new Set<string>();
let multiNode = false;
for (const id of [...link.prefillInstances, ...link.decodeInstances]) {
const r = instanceById[id];
if (r && r.baseModel) fams.add(r.baseModel.toLowerCase());
if (r && r.nodeCount > 1) multiNode = true;
}
rows.push({
linkId: link.linkId,
prefill: link.prefillInstances,
decode: link.decodeInstances,
families: Array.from(fams),
multiNode,
});
}
return rows;
});
let editingLinkId = $state<string | null>(null);
let editingPrefill = $state<Set<string>>(new Set());
let editingDecode = $state<Set<string>>(new Set());
let saving = $state(false);
let errorMessage = $state<string | null>(null);
function startCreate() {
editingLinkId = "new";
editingPrefill = new Set();
editingDecode = new Set();
errorMessage = null;
}
function startEdit(row: LinkRow) {
editingLinkId = row.linkId;
editingPrefill = new Set(row.prefill);
editingDecode = new Set(row.decode);
errorMessage = null;
}
function cancelEdit() {
editingLinkId = null;
editingPrefill = new Set();
editingDecode = new Set();
errorMessage = null;
}
type Role = "prefill" | "decode" | "none";
function roleOf(id: string): Role {
if (editingPrefill.has(id)) return "prefill";
if (editingDecode.has(id)) return "decode";
return "none";
}
function setRole(id: string, role: Role) {
const p = new Set(editingPrefill);
const d = new Set(editingDecode);
p.delete(id);
d.delete(id);
if (role === "prefill") p.add(id);
if (role === "decode") d.add(id);
editingPrefill = p;
editingDecode = d;
}
const editingFamilies = $derived.by<string[]>(() => {
const fams = new Set<string>();
for (const id of [...editingPrefill, ...editingDecode]) {
const r = instanceById[id];
if (r && r.baseModel) fams.add(r.baseModel.toLowerCase());
}
return Array.from(fams);
});
const editingMultiNode = $derived.by<string[]>(() => {
const names: string[] = [];
for (const id of [...editingPrefill, ...editingDecode]) {
const r = instanceById[id];
if (r && r.nodeCount > 1) {
names.push(r.baseModel || r.modelId);
}
}
return names;
});
const editingMismatch = $derived(editingFamilies.length > 1);
const canSave = $derived(
editingLinkId !== null &&
editingPrefill.size > 0 &&
editingDecode.size > 0 &&
!saving,
);
async function save() {
if (editingLinkId === null) return;
saving = true;
errorMessage = null;
try {
const prefill = Array.from(editingPrefill);
const decode = Array.from(editingDecode);
if (editingLinkId === "new") {
await createInstanceLink(prefill, decode);
} else {
await updateInstanceLink(editingLinkId, prefill, decode);
}
cancelEdit();
await refreshState();
} catch (err) {
errorMessage = err instanceof Error ? err.message : String(err);
} finally {
saving = false;
}
}
async function remove(linkId: string) {
if (!confirm("Remove this routing?")) return;
try {
await deleteInstanceLink(linkId);
if (editingLinkId === linkId) cancelEdit();
await refreshState();
} catch (err) {
errorMessage = err instanceof Error ? err.message : String(err);
}
}
</script>
<div class="font-mono text-foreground">
<div class="mb-6 space-y-4">
<details open class="group [&_summary::-webkit-details-marker]:hidden">
<summary
class="cursor-pointer list-none text-exo-yellow text-xs font-mono tracking-widest uppercase flex items-center gap-2 hover:opacity-80 transition-opacity"
>
<span
class="inline-block transition-transform group-open:rotate-90 text-exo-light-gray"
>▶</span
>
Prefill vs Decode
</summary>
<div class="mt-2 text-white/80 text-sm leading-relaxed">
Prefill is the compute-bound pass that consumes the entire prompt and
builds a KV cache. Decode is the memory-bandwidth-bound loop that emits
tokens sequentially from that cache. The two phases have very different
bottlenecks, so running them on different hardware can be substantially
faster than doing both on one node.
</div>
</details>
<details class="group [&_summary::-webkit-details-marker]:hidden">
<summary
class="cursor-pointer list-none text-exo-yellow text-xs font-mono tracking-widest uppercase flex items-center gap-2 hover:opacity-80 transition-opacity"
>
<span
class="inline-block transition-transform group-open:rotate-90 text-exo-light-gray"
>▶</span
>
Linking Instances
</summary>
<div class="mt-2 text-white/80 text-sm leading-relaxed space-y-2">
<p>
A linked route here tells the cluster: when a request is sent to a
model in that cluster, the decode node (or the least active one if
there are multiple) will handle it. If it decides it must do a lot of
prefill not already cached in the prefix cache, it routes the request
to the prefill node over TCP IP. The prefill node streams the KV cache
back to the decode node which picks up from there.
</p>
<p>
Linked instances must be running the same model family — KV layouts
differ across architectures. More on the <a
class="text-exo-yellow underline underline-offset-2 hover:text-exo-yellow-darker transition-colors"
href="https://blog.exolabs.net/nvidia-dgx-spark/"
target="_blank"
rel="noreferrer noopener">blog</a
>.
</p>
</div>
</details>
</div>
{#if errorMessage}
<div
class="mb-4 px-4 py-3 bg-red-500/10 border border-red-500/40 text-red-300 text-sm"
>
{errorMessage}
</div>
{/if}
<section class="mt-12">
<h2
class="text-exo-yellow text-xs font-mono tracking-widest uppercase m-0 mb-3"
>
Existing routes
</h2>
{#if linkRows.length === 0}
{#if editingLinkId === null}
<div class="flex items-center justify-between">
<p class="text-exo-light-gray italic text-sm m-0">
No routes yet. Create one to enable remote prefill.
</p>
<button
class="px-3 py-1.5 text-xs font-mono tracking-wider uppercase bg-exo-yellow/15 border border-exo-yellow/50 text-exo-yellow hover:bg-exo-yellow/25 hover:border-exo-yellow/80 transition-colors"
onclick={startCreate}
>
+ New route
</button>
</div>
{/if}
{:else}
{#if editingLinkId === null}
<div class="flex justify-end mb-3">
<button
class="px-3 py-1.5 text-xs font-mono tracking-wider uppercase bg-exo-yellow/15 border border-exo-yellow/50 text-exo-yellow hover:bg-exo-yellow/25 hover:border-exo-yellow/80 transition-colors"
onclick={startCreate}
>
+ New route
</button>
</div>
{/if}
<div
class="bg-exo-dark-gray/60 border border-exo-medium-gray/40 flex flex-col"
>
{#each linkRows as row (row.linkId)}
{#if editingLinkId !== row.linkId}
<article
class="p-4 border-b border-exo-light-gray/25 last:border-b-0"
>
{#if row.multiNode}
<div
class="mb-3 px-3 py-2 bg-red-500/10 border border-red-500/40 text-red-300 text-xs tracking-wide"
>
⚠ Multi-node instance detected. Remote prefill currently only
works on single-node (rank-0) instances. This route will not
function until that's supported.
</div>
{/if}
{#if row.families.length > 1}
<div
class="mb-3 px-3 py-2 bg-amber-500/10 border border-amber-500/40 text-amber-300 text-xs tracking-wide"
>
⚠ Mixed model families: {row.families.join(", ")}
</div>
{/if}
<div
class="grid grid-cols-[1fr_auto_1fr_auto] items-center gap-x-3 gap-y-2"
>
<span
class="inline-block justify-self-start text-[10px] font-mono tracking-widest uppercase px-2 py-0.5 bg-exo-yellow/15 border border-exo-yellow/40 text-exo-yellow"
>Prefill</span
>
<span></span>
<span
class="inline-block justify-self-start text-[10px] font-mono tracking-widest uppercase px-2 py-0.5 bg-exo-medium-gray/40 border border-exo-medium-gray/60 text-foreground"
>Decode</span
>
<span></span>
<div class="min-w-0">
<ul class="list-none p-0 m-0 flex flex-col gap-2">
{#each row.prefill as id (id)}
{@const r = instanceById[id]}
{#if r}
<li
class="flex items-center gap-2 px-2.5 py-2 bg-exo-medium-gray/20 border border-exo-medium-gray/40"
>
<FamilyLogos family={r.family} />
<div class="min-w-0 flex-1">
<div
class="text-exo-yellow text-xs font-mono truncate"
>
{r.baseModel || r.modelId}
</div>
<div
class="text-exo-light-gray text-[11px] truncate"
>
{r.nodeNames.join(", ") || "?"}{r.nodeCount > 1
? ` (${r.nodeCount} nodes)`
: ""}
</div>
<div
class="text-exo-light-gray/40 text-[10px] font-mono truncate"
title={r.id}
>
{r.id.slice(0, 8)}
</div>
</div>
</li>
{/if}
{/each}
</ul>
</div>
<div class="text-exo-yellow/60 text-xl px-2" aria-hidden="true">
</div>
<div class="min-w-0">
<ul class="list-none p-0 m-0 flex flex-col gap-2">
{#each row.decode as id (id)}
{@const r = instanceById[id]}
{#if r}
<li
class="flex items-center gap-2 px-2.5 py-2 bg-exo-medium-gray/20 border border-exo-medium-gray/40"
>
<FamilyLogos family={r.family} />
<div class="min-w-0 flex-1">
<div
class="text-exo-yellow text-xs font-mono truncate"
>
{r.baseModel || r.modelId}
</div>
<div
class="text-exo-light-gray text-[11px] truncate"
>
{r.nodeNames.join(", ") || "?"}{r.nodeCount > 1
? ` (${r.nodeCount} nodes)`
: ""}
</div>
<div
class="text-exo-light-gray/40 text-[10px] font-mono truncate"
title={r.id}
>
{r.id.slice(0, 8)}
</div>
</div>
</li>
{/if}
{/each}
</ul>
</div>
<div class="flex gap-2 pl-3">
<button
class="px-2 py-0.5 text-[11px] font-mono tracking-wider uppercase bg-exo-medium-gray/30 border border-exo-medium-gray/60 rounded text-foreground hover:border-exo-yellow/60 hover:text-exo-yellow disabled:opacity-40 disabled:cursor-not-allowed transition-colors"
onclick={() => startEdit(row)}
disabled={editingLinkId !== null}
>
Edit
</button>
<button
class="px-2 py-0.5 text-[11px] font-mono tracking-wider uppercase bg-red-500/15 border border-red-500/40 rounded text-red-300 hover:bg-red-500/25 transition-colors"
onclick={() => remove(row.linkId)}
>
Remove
</button>
</div>
</div>
</article>
{/if}
{/each}
</div>
{/if}
</section>
{#if editingLinkId !== null && instanceRows.length === 0}
<section
class="mt-6 bg-exo-dark-gray/60 border border-exo-yellow/30 px-4 py-2.5 flex items-center justify-between gap-3"
>
<span class="text-exo-light-gray italic text-sm font-mono"
>No instances available.</span
>
<button
class="px-3 py-1 text-xs font-mono tracking-wider uppercase bg-exo-medium-gray/30 border border-exo-medium-gray/60 rounded text-foreground hover:border-exo-yellow/60 transition-colors"
onclick={cancelEdit}
>
Cancel
</button>
</section>
{:else if editingLinkId !== null}
<section class="mt-6 bg-exo-dark-gray/60 border border-exo-yellow/30 p-5">
<h2
class="text-exo-yellow text-xs font-mono tracking-widest uppercase m-0 mb-3"
>
{editingLinkId === "new" ? "New route" : "Edit route"}
</h2>
{#if editingMismatch}
<div
class="mb-3 px-3 py-2 bg-amber-500/10 border border-amber-500/40 text-amber-300 text-xs tracking-wide"
>
⚠ Selected instances span multiple model families: <strong
>{editingFamilies.join(", ")}</strong
>. Linking across families produces a corrupt KV cache.
</div>
{/if}
{#if editingMultiNode.length > 0}
<div
class="mb-3 px-3 py-2 bg-red-500/10 border border-red-500/40 text-red-300 text-xs tracking-wide"
>
⚠ Multi-node instance(s) selected: <strong
>{editingMultiNode.join(", ")}</strong
>. Remote prefill currently only works on single-node instances. This
route will not function until multi-node support lands.
</div>
{/if}
<p class="text-exo-light-gray text-xs mb-4">
Pick a role for each instance:
<span class="text-exo-yellow">Prefill</span>
serves KV cache,
<span class="text-foreground">Decode</span> consumes it.
</p>
<div
class="grid gap-2.5"
style="grid-template-columns: repeat(auto-fill, minmax(360px, 1fr));"
>
{#each instanceRows as row (row.id)}
{@const role = roleOf(row.id)}
<div
class="border p-3 flex flex-col gap-2.5 transition-colors {role ===
'prefill'
? 'border-exo-yellow/60 bg-exo-dark-gray/60'
: role === 'decode'
? 'border-exo-light-gray/60 bg-exo-dark-gray/60'
: 'border-exo-medium-gray/40 bg-exo-dark-gray/40'}"
>
<div class="flex items-center gap-2">
<FamilyLogos family={row.family} />
<div class="min-w-0 flex-1">
<div class="text-exo-yellow text-xs font-mono truncate">
{row.baseModel || row.modelId}
</div>
<div class="text-exo-light-gray text-[11px] truncate">
{row.nodeNames.join(", ") || "?"}{row.nodeCount > 1
? ` (${row.nodeCount} nodes)`
: ""}
</div>
<div
class="text-exo-light-gray/40 text-[10px] font-mono truncate"
title={row.id}
>
{row.id.slice(0, 8)}
</div>
</div>
{#if row.nodeCount > 1}
<span
class="text-[9px] font-mono tracking-widest uppercase px-1.5 py-0.5 bg-red-500/15 border border-red-500/40 text-red-300"
title="Multi-node instances are not supported by remote prefill yet."
>Unsupported</span
>
{/if}
</div>
<div
class="flex rounded-md overflow-hidden border border-exo-light-gray/40 divide-x divide-exo-light-gray/40"
>
<button
class="flex-1 px-2 py-1 text-[11px] font-mono tracking-wider uppercase transition-colors {role ===
'prefill'
? 'bg-exo-yellow/20 text-exo-yellow'
: 'bg-transparent text-white/80 hover:text-exo-yellow'}"
onclick={() =>
setRole(row.id, role === "prefill" ? "none" : "prefill")}
>Prefill</button
>
<button
class="flex-1 px-2 py-1 text-[11px] font-mono tracking-wider uppercase transition-colors {role ===
'decode'
? 'bg-exo-medium-gray/50 text-foreground'
: 'bg-transparent text-white/80 hover:text-foreground'}"
onclick={() =>
setRole(row.id, role === "decode" ? "none" : "decode")}
>Decode</button
>
</div>
</div>
{/each}
</div>
<div class="flex gap-2 mt-5 justify-end">
<button
class="px-3 py-1.5 text-xs font-mono tracking-wider uppercase bg-exo-yellow/15 border border-exo-yellow/50 text-exo-yellow hover:bg-exo-yellow/25 hover:border-exo-yellow/80 disabled:opacity-40 disabled:cursor-not-allowed transition-colors"
onclick={save}
disabled={!canSave}
>
{saving ? "Saving..." : "Save route"}
</button>
<button
class="px-3 py-1.5 text-xs font-mono tracking-wider uppercase bg-exo-medium-gray/30 border border-exo-medium-gray/60 text-foreground hover:border-exo-yellow/60 disabled:opacity-40 disabled:cursor-not-allowed transition-colors"
onclick={cancelEdit}
disabled={saving}
>
Cancel
</button>
</div>
</section>
{/if}
</div>
@@ -117,6 +117,10 @@
const LOGO_NATIVE_WIDTH = 814;
const LOGO_NATIVE_HEIGHT = 1000;
// NVIDIA logo SVG path (from exo-nvidia)
const NVIDIA_LOGO_PATH =
"M0.81 0.429V0.299c0.013 -0.001 0.026 -0.002 0.038 -0.002 0.355 -0.011 0.588 0.306 0.588 0.306S1.186 0.952 0.916 0.952c-0.036 0 -0.071 -0.006 -0.105 -0.017V0.542c0.138 0.017 0.166 0.078 0.249 0.216l0.185 -0.155s-0.135 -0.177 -0.362 -0.177c-0.024 -0.001 -0.048 0.001 -0.072 0.003m0 -0.429v0.194l0.038 -0.002c0.494 -0.017 0.816 0.405 0.816 0.405s-0.37 0.45 -0.754 0.45c-0.034 0 -0.066 -0.003 -0.099 -0.009v0.12c0.027 0.003 0.055 0.006 0.082 0.006 0.358 0 0.618 -0.183 0.869 -0.399 0.042 0.034 0.212 0.114 0.247 0.15 -0.238 0.2 -0.794 0.361 -1.11 0.361 -0.03 0 -0.059 -0.002 -0.088 -0.005v0.169h1.362V0zm0 0.935v0.102c-0.331 -0.059 -0.423 -0.404 -0.423 -0.404s0.159 -0.176 0.423 -0.205v0.112h-0.001C0.671 0.524 0.562 0.654 0.562 0.654s0.062 0.218 0.248 0.282m-0.588 -0.316s0.196 -0.29 0.589 -0.32V0.194C0.376 0.229 0 0.597 0 0.597s0.213 0.616 0.81 0.672v-0.112c-0.438 -0.054 -0.588 -0.538 -0.588 -0.538";
function formatBytes(bytes: number, decimals = 1): string {
if (!bytes || bytes === 0) return "0B";
const k = 1024;
@@ -554,6 +558,13 @@
const clipPathId = `clip-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
const modelLower = modelId.toLowerCase();
const identity = identitiesData[nodeInfo.id];
const nameLower = (friendlyName || "").toLowerCase();
const isSpark = modelLower.includes("dgx") || modelLower.includes("gx10");
const isLinux =
!isSpark &&
(modelLower.startsWith("linux") || identity?.osVersion === "Linux");
const isLinuxLaptop = isLinux && modelLower.includes("laptop");
// Check node states for styling
const isHighlighted = highlightedNodes.has(nodeInfo.id);
@@ -623,7 +634,382 @@
`${friendlyName}\nID: ${nodeInfo.id.slice(-8)}\nMemory: ${formatBytes(ramUsed)}/${formatBytes(ramTotal)}`,
);
if (modelLower === "mac studio") {
if (isSpark) {
// NVIDIA DGX Spark — gold chassis with textured front, side handles, and NVIDIA badge
iconBaseWidth = nodeRadius * 1.55;
iconBaseHeight = nodeRadius * 0.58;
const x = nodeInfo.x - iconBaseWidth / 2;
const y = nodeInfo.y - iconBaseHeight / 2;
const chassisX = x - iconBaseWidth * 0.03;
const chassisWidth = iconBaseWidth * 1.05;
const cornerRadius = 3;
const dgxClipId = `dgx-clip-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
defs
.append("clipPath")
.attr("id", dgxClipId)
.append("rect")
.attr("x", x)
.attr("y", y)
.attr("width", iconBaseWidth)
.attr("height", iconBaseHeight)
.attr("rx", cornerRadius);
// Chassis texture pattern
const textureId = `chassis-texture-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
defs
.append("pattern")
.attr("id", textureId)
.attr("patternUnits", "userSpaceOnUse")
.attr("width", 8)
.attr("height", 8);
const texturePattern = defs.select(`#${textureId}`);
texturePattern
.append("rect")
.attr("width", 8)
.attr("height", 8)
.attr("fill", "#6f6248");
texturePattern
.append("circle")
.attr("cx", 2)
.attr("cy", 2)
.attr("r", 1)
.attr("fill", "#5a4f3b")
.attr("opacity", 0.5);
texturePattern
.append("circle")
.attr("cx", 6)
.attr("cy", 6)
.attr("r", 1)
.attr("fill", "#4a4232")
.attr("opacity", 0.45);
// Main body
nodeG
.append("rect")
.attr("class", "node-outline")
.attr("x", chassisX)
.attr("y", y)
.attr("width", chassisWidth)
.attr("height", iconBaseHeight)
.attr("rx", cornerRadius)
.attr("fill", `url(#${textureId})`)
.attr("stroke", wireColor)
.attr("stroke-width", strokeWidth);
// Side border accents
const sideThickness = iconBaseWidth * 0.02;
nodeG
.append("rect")
.attr("x", chassisX)
.attr("y", y)
.attr("width", sideThickness)
.attr("height", iconBaseHeight)
.attr("fill", "#8a7a56");
nodeG
.append("rect")
.attr("x", chassisX + chassisWidth - sideThickness)
.attr("y", y)
.attr("width", sideThickness)
.attr("height", iconBaseHeight)
.attr("fill", "#8a7a56");
// Memory fill (bottom up)
if (ramUsagePercent > 0) {
const memFillHeight = (ramUsagePercent / 100) * iconBaseHeight;
nodeG
.append("rect")
.attr("x", x)
.attr("y", y + iconBaseHeight - memFillHeight)
.attr("width", iconBaseWidth)
.attr("height", memFillHeight)
.attr("fill", "rgba(255,215,0,0.45)")
.attr("clip-path", `url(#${dgxClipId})`);
}
// Side handles with inner recess
const handleWidth = iconBaseWidth * 0.27;
const handleGap = iconBaseHeight * 0.05;
const handleHeight = iconBaseHeight - handleGap * 2;
const handleY = y + handleGap;
const innerHandleWidth = iconBaseWidth * 0.12;
const innerHandleHeight = handleHeight - iconBaseHeight * 0.06;
const leftHandleX = x + 4;
const rightHandleX = x + iconBaseWidth - handleWidth - 4;
// Left handle
nodeG
.append("rect")
.attr("x", leftHandleX)
.attr("y", handleY)
.attr("width", handleWidth)
.attr("height", handleHeight)
.attr("rx", 2.4)
.attr("fill", "#b3a170")
.attr("stroke", "#403723")
.attr("stroke-width", 0.7);
nodeG
.append("rect")
.attr("x", leftHandleX + handleWidth * 0.06)
.attr("y", handleY + iconBaseHeight * 0.03)
.attr("width", innerHandleWidth)
.attr("height", innerHandleHeight)
.attr("rx", 1.6)
.attr("fill", "#8a7a56");
// Right handle
nodeG
.append("rect")
.attr("x", rightHandleX)
.attr("y", handleY)
.attr("width", handleWidth)
.attr("height", handleHeight)
.attr("rx", 2.4)
.attr("fill", "#b3a170")
.attr("stroke", "#403723")
.attr("stroke-width", 0.7);
nodeG
.append("rect")
.attr(
"x",
rightHandleX + handleWidth - innerHandleWidth - handleWidth * 0.08,
)
.attr("y", handleY + iconBaseHeight * 0.03)
.attr("width", innerHandleWidth)
.attr("height", innerHandleHeight)
.attr("rx", 1.6)
.attr("fill", "#8a7a56");
// NVIDIA logo + text label (rotated 90 deg on left handle)
const badgeWidth = iconBaseWidth * 0.09;
const badgeHeight = handleHeight * 0.5;
const badgeX =
leftHandleX + handleWidth - badgeWidth - handleWidth * 0.06;
const badgeY = handleY + (handleHeight - badgeHeight) / 2;
const textSize = badgeWidth * 0.58;
const logoWidth = textSize * 1.2;
const logoHeight = logoWidth * (1.438 / 2.174);
const centerX = badgeX + badgeWidth / 2 - badgeWidth * 0.03;
const centerY = badgeY + badgeHeight / 2;
const gap = badgeWidth * 0.15;
const totalWidth = logoWidth + gap + textSize * 3.6;
const labelGroup = nodeG
.append("g")
.attr("transform", `rotate(90 ${centerX} ${centerY})`);
labelGroup
.append("svg")
.attr("x", centerX - totalWidth / 2)
.attr("y", centerY - logoHeight / 2)
.attr("width", logoWidth)
.attr("height", logoHeight)
.attr("viewBox", "0 0 2.174 1.438")
.append("path")
.attr("d", NVIDIA_LOGO_PATH)
.attr("fill", "#76b900");
labelGroup
.append("text")
.attr("x", centerX - totalWidth / 2 + logoWidth + gap)
.attr("y", centerY)
.attr("text-anchor", "start")
.attr("dominant-baseline", "middle")
.attr("fill", "#8a7a56")
.attr("font-size", textSize)
.attr("font-family", "monospace")
.attr("font-weight", "700")
.text("NVIDIA");
} else if (isLinuxLaptop) {
// Linux Laptop — same shape as MacBook but with Tux logo
iconBaseWidth = nodeRadius * 1.6;
iconBaseHeight = nodeRadius * 1.15;
const x = nodeInfo.x - iconBaseWidth / 2;
const y = nodeInfo.y - iconBaseHeight / 2;
const screenHeight = iconBaseHeight * 0.7;
const baseHeight = iconBaseHeight * 0.3;
const screenWidth = iconBaseWidth * 0.85;
const screenX = nodeInfo.x - screenWidth / 2;
const screenBezel = 3;
const linuxScreenClipId = `linux-screen-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
defs
.append("clipPath")
.attr("id", linuxScreenClipId)
.append("rect")
.attr("x", screenX + screenBezel)
.attr("y", y + screenBezel)
.attr("width", screenWidth - screenBezel * 2)
.attr("height", screenHeight - screenBezel * 2)
.attr("rx", 2);
// Screen outer frame
nodeG
.append("rect")
.attr("class", "node-outline")
.attr("x", screenX)
.attr("y", y)
.attr("width", screenWidth)
.attr("height", screenHeight)
.attr("rx", 3)
.attr("fill", "#1a1a1a")
.attr("stroke", wireColor)
.attr("stroke-width", strokeWidth);
// Screen inner
nodeG
.append("rect")
.attr("x", screenX + screenBezel)
.attr("y", y + screenBezel)
.attr("width", screenWidth - screenBezel * 2)
.attr("height", screenHeight - screenBezel * 2)
.attr("rx", 2)
.attr("fill", "#0a0a12");
// Memory fill on screen
if (ramUsagePercent > 0) {
const memFillTotalHeight = screenHeight - screenBezel * 2;
const memFillActualHeight =
(ramUsagePercent / 100) * memFillTotalHeight;
nodeG
.append("rect")
.attr("x", screenX + screenBezel)
.attr(
"y",
y + screenBezel + (memFillTotalHeight - memFillActualHeight),
)
.attr("width", screenWidth - screenBezel * 2)
.attr("height", memFillActualHeight)
.attr("fill", "rgba(255,215,0,0.85)")
.attr("clip-path", `url(#${linuxScreenClipId})`);
}
// Terminal prompt on screen
nodeG
.append("text")
.attr("x", nodeInfo.x)
.attr("y", y + screenHeight / 2)
.attr("text-anchor", "middle")
.attr("dominant-baseline", "middle")
.attr("fill", "#FFFFFF")
.attr("opacity", 0.9)
.attr("font-size", screenHeight * 0.25)
.attr("font-family", "SF Mono, Monaco, monospace")
.attr("font-weight", "700")
.text(">_");
// Keyboard base (trapezoidal)
const baseY = y + screenHeight;
const baseTopWidth = screenWidth;
const baseBottomWidth = iconBaseWidth;
const baseTopX = nodeInfo.x - baseTopWidth / 2;
const baseBottomX = nodeInfo.x - baseBottomWidth / 2;
nodeG
.append("path")
.attr(
"d",
`M ${baseTopX} ${baseY} L ${baseTopX + baseTopWidth} ${baseY} L ${baseBottomX + baseBottomWidth} ${baseY + baseHeight} L ${baseBottomX} ${baseY + baseHeight} Z`,
)
.attr("fill", "#2c2c2c")
.attr("stroke", wireColor)
.attr("stroke-width", 1);
// Keyboard area
const keyboardX = baseTopX + 6;
const keyboardY = baseY + 3;
const keyboardWidth = baseTopWidth - 12;
const keyboardHeight = baseHeight * 0.55;
nodeG
.append("rect")
.attr("x", keyboardX)
.attr("y", keyboardY)
.attr("width", keyboardWidth)
.attr("height", keyboardHeight)
.attr("fill", "rgba(0,0,0,0.2)")
.attr("rx", 2);
// Trackpad
const trackpadWidth = baseTopWidth * 0.4;
const trackpadX = nodeInfo.x - trackpadWidth / 2;
const trackpadY = baseY + keyboardHeight + 5;
const trackpadHeight = baseHeight * 0.3;
nodeG
.append("rect")
.attr("x", trackpadX)
.attr("y", trackpadY)
.attr("width", trackpadWidth)
.attr("height", trackpadHeight)
.attr("fill", "rgba(255,255,255,0.08)")
.attr("rx", 2);
} else if (isLinux) {
// Linux Desktop — same shape as Mac Studio but with Tux logo
iconBaseWidth = nodeRadius * 1.25;
iconBaseHeight = nodeRadius * 0.85;
const x = nodeInfo.x - iconBaseWidth / 2;
const y = nodeInfo.y - iconBaseHeight / 2;
const cornerRadius = 4;
const topSurfaceHeight = iconBaseHeight * 0.15;
const linuxDesktopClipId = `linux-desktop-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
defs
.append("clipPath")
.attr("id", linuxDesktopClipId)
.append("rect")
.attr("x", x)
.attr("y", y + topSurfaceHeight)
.attr("width", iconBaseWidth)
.attr("height", iconBaseHeight - topSurfaceHeight)
.attr("rx", cornerRadius - 1);
// Main body
nodeG
.append("rect")
.attr("class", "node-outline")
.attr("x", x)
.attr("y", y)
.attr("width", iconBaseWidth)
.attr("height", iconBaseHeight)
.attr("rx", cornerRadius)
.attr("fill", "#1a1a1a")
.attr("stroke", wireColor)
.attr("stroke-width", strokeWidth);
// Memory fill
if (ramUsagePercent > 0) {
const memFillTotalHeight = iconBaseHeight - topSurfaceHeight;
const memFillActualHeight =
(ramUsagePercent / 100) * memFillTotalHeight;
nodeG
.append("rect")
.attr("x", x)
.attr(
"y",
y + topSurfaceHeight + (memFillTotalHeight - memFillActualHeight),
)
.attr("width", iconBaseWidth)
.attr("height", memFillActualHeight)
.attr("fill", "rgba(255,215,0,0.75)")
.attr("clip-path", `url(#${linuxDesktopClipId})`);
}
// Terminal prompt on front face
nodeG
.append("text")
.attr("x", nodeInfo.x)
.attr(
"y",
y + topSurfaceHeight + (iconBaseHeight - topSurfaceHeight) / 2,
)
.attr("text-anchor", "middle")
.attr("dominant-baseline", "middle")
.attr("fill", "rgba(255,255,255,0.5)")
.attr("font-size", (iconBaseHeight - topSurfaceHeight) * 0.4)
.attr("font-family", "SF Mono, Monaco, monospace")
.attr("font-weight", "700")
.text(">_");
} else if (modelLower === "mac studio") {
// Mac Studio - classic cube with memory fill
iconBaseWidth = nodeRadius * 1.25;
iconBaseHeight = nodeRadius * 0.85;
@@ -1182,8 +1568,12 @@
debugLabelY += debugLineHeight;
}
const identity = identitiesData[nodeInfo.id];
if (identity?.osVersion) {
const dbgIdentity = identitiesData[nodeInfo.id];
if (dbgIdentity?.osVersion) {
const osLabel =
dbgIdentity.osVersion === "Linux"
? "Linux"
: `macOS ${dbgIdentity.osVersion}${dbgIdentity.osBuildVersion ? ` (${dbgIdentity.osBuildVersion})` : ""}`;
nodeG
.append("text")
.attr("x", nodeInfo.x)
@@ -1192,9 +1582,7 @@
.attr("fill", "rgba(179,179,179,0.7)")
.attr("font-size", debugFontSize)
.attr("font-family", "SF Mono, Monaco, monospace")
.text(
`macOS ${identity.osVersion}${identity.osBuildVersion ? ` (${identity.osBuildVersion})` : ""}`,
);
.text(osLabel);
}
}
});
+128 -10
View File
@@ -74,6 +74,12 @@ export interface Instance {
};
}
export interface RawInstanceLink {
linkId: string;
prefillInstances: string[];
decodeInstances: string[];
}
// Granular node state types from the new state structure
interface RawNodeIdentity {
modelId?: string;
@@ -223,6 +229,7 @@ interface RawStateResponse {
}
>;
runners?: Record<string, unknown>;
instanceLinks?: Record<string, RawInstanceLink>;
downloads?: Record<string, unknown[]>;
// New granular node state fields
nodeIdentities?: Record<string, RawNodeIdentity>;
@@ -541,6 +548,8 @@ class AppStore {
topologyData = $state<TopologyData | null>(null);
instances = $state<Record<string, unknown>>({});
runners = $state<Record<string, unknown>>({});
instanceLinks = $state<Record<string, RawInstanceLink>>({});
featureFlags = $state<Record<string, boolean>>({});
downloads = $state<Record<string, unknown[]>>({});
nodeDisk = $state<
Record<
@@ -1274,6 +1283,7 @@ class AppStore {
startPolling() {
this.fetchState();
this.fetchFeatureFlags();
this.fetchInterval = setInterval(() => this.fetchState(), 1000);
}
@@ -1285,6 +1295,16 @@ class AppStore {
this.stopPreviewsPolling();
}
async fetchFeatureFlags() {
try {
const response = await fetch("/v1/feature-flags");
if (!response.ok) return;
this.featureFlags = await response.json();
} catch {
// Silently ignore — defaults to all-disabled.
}
}
async fetchState() {
try {
const response = await fetch("/state");
@@ -1310,6 +1330,11 @@ class AppStore {
if (data.runners) {
this.runners = data.runners;
}
if (data.instanceLinks) {
this.instanceLinks = data.instanceLinks;
} else {
this.instanceLinks = {};
}
if (data.downloads) {
this.downloads = data.downloads;
}
@@ -1670,7 +1695,15 @@ class AppStore {
}
}
}
return { role: m.role, content: msgContent };
const out: {
role: string;
content: string;
reasoning_content?: string;
} = { role: m.role, content: msgContent };
if (m.role === "assistant" && m.thinking) {
out.reasoning_content = m.thinking;
}
return out;
}),
];
@@ -1877,7 +1910,15 @@ class AppStore {
const apiMessages = [
systemPrompt,
...targetConversation.messages.slice(0, -1).map((m) => {
return { role: m.role, content: m.content };
const out: {
role: string;
content: string;
reasoning_content?: string;
} = { role: m.role, content: m.content };
if (m.role === "assistant" && m.thinking) {
out.reasoning_content = m.thinking;
}
return out;
}),
];
@@ -2408,10 +2449,15 @@ class AppStore {
contentParts.push({ type: "text", text: textContent });
}
return {
role: m.role,
content: contentParts,
};
const out: {
role: string;
content: typeof contentParts;
reasoning_content?: string;
} = { role: m.role, content: contentParts };
if (m.role === "assistant" && m.thinking) {
out.reasoning_content = m.thinking;
}
return out;
}
// Text-only message (original path)
@@ -2429,10 +2475,15 @@ class AppStore {
}
}
return {
role: m.role,
content: msgContent,
};
const out: {
role: string;
content: string;
reasoning_content?: string;
} = { role: m.role, content: msgContent };
if (m.role === "assistant" && m.thinking) {
out.reasoning_content = m.thinking;
}
return out;
}),
];
@@ -3281,6 +3332,60 @@ class AppStore {
}
}
async createInstanceLink(
prefillInstances: string[],
decodeInstances: string[],
): Promise<void> {
const response = await fetch("/v1/instance-links", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({
prefill_instances: prefillInstances,
decode_instances: decodeInstances,
}),
});
if (!response.ok) {
throw new Error(
`Failed to create instance link: ${response.status} ${await response.text()}`,
);
}
}
async updateInstanceLink(
linkId: string,
prefillInstances: string[],
decodeInstances: string[],
): Promise<void> {
const response = await fetch(
`/v1/instance-links/${encodeURIComponent(linkId)}`,
{
method: "PUT",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({
prefill_instances: prefillInstances,
decode_instances: decodeInstances,
}),
},
);
if (!response.ok) {
throw new Error(
`Failed to update instance link: ${response.status} ${await response.text()}`,
);
}
}
async deleteInstanceLink(linkId: string): Promise<void> {
const response = await fetch(
`/v1/instance-links/${encodeURIComponent(linkId)}`,
{ method: "DELETE" },
);
if (!response.ok) {
throw new Error(
`Failed to delete instance link: ${response.status} ${await response.text()}`,
);
}
}
/**
* Delete a downloaded model from a specific node
*/
@@ -3379,6 +3484,19 @@ export const prefillProgress = () => appStore.prefillProgress;
export const topologyData = () => appStore.topologyData;
export const instances = () => appStore.instances;
export const runners = () => appStore.runners;
export const instanceLinks = () => appStore.instanceLinks;
export const featureFlags = () => appStore.featureFlags;
export const createInstanceLink = (
prefillInstances: string[],
decodeInstances: string[],
) => appStore.createInstanceLink(prefillInstances, decodeInstances);
export const updateInstanceLink = (
linkId: string,
prefillInstances: string[],
decodeInstances: string[],
) => appStore.updateInstanceLink(linkId, prefillInstances, decodeInstances);
export const deleteInstanceLink = (linkId: string) =>
appStore.deleteInstanceLink(linkId);
export const downloads = () => appStore.downloads;
export const nodeDisk = () => appStore.nodeDisk;
export const placementPreviews = () => appStore.placementPreviews;
+44
View File
@@ -0,0 +1,44 @@
// Mirrors src/exo/shared/models/model_cards.py:derive_base_model
const QUANT_SUFFIXES = new RegExp(
"[-_ ](?:MLX|MXFP[0-9]+|NVFP[0-9]+|GPTQ|AWQ|GGUF|fp16|bf16|fp8|int[0-9]+|[0-9]+(?:\\.[0-9]+)?bit|Q[0-9]+(?:_[A-Z0-9]+)?|gs[0-9]+)" +
"(?:[-_ ](?:MLX|Q[0-9]+|Int[0-9]+|[A-Z0-9]+|gs[0-9]+))*$",
"i",
);
function normalize(s: string): string {
return s
.replaceAll("-", " ")
.replaceAll("_", " ")
.replaceAll(" ", " ")
.trim();
}
export function deriveBaseModel(modelId: string): string {
const short = modelId.includes("/")
? (modelId.split("/").pop() ?? modelId)
: modelId;
const stripped = short.replace(QUANT_SUFFIXES, "");
return normalize(stripped);
}
export function baseModelsCompatible(a: string, b: string): boolean {
return deriveBaseModel(a).toLowerCase() === deriveBaseModel(b).toLowerCase();
}
// Mirrors src/exo/shared/models/model_cards.py:derive_family
export function deriveFamily(modelId: string): string {
const short = modelId.includes("/")
? (modelId.split("/").pop() ?? modelId)
: modelId;
const stripped = short
.replace(QUANT_SUFFIXES, "")
.toLowerCase()
.replaceAll("_", "-");
const parts = stripped.split(/[-.]/);
const familyParts: string[] = [];
for (const p of parts) {
if (/^\d+$/.test(p) || /^\d+[bm]?$/i.test(p)) break;
familyParts.push(p);
}
return familyParts.length > 0 ? familyParts.join("-") : stripped;
}
+82 -14
View File
@@ -65,6 +65,7 @@
nodeThunderboltBridge,
nodeIdentities,
isConnected,
featureFlags,
type DownloadProgress,
type PlacementPreview,
} from "$lib/stores/app.svelte";
@@ -702,7 +703,10 @@
? Object.keys(topologyData()!.nodes).length
: 1;
const sharding = nodeCount <= 1 ? "Pipeline" : selectedSharding;
const instanceType = nodeCount <= 1 ? "MlxRing" : selectedInstanceType;
const instanceType =
nodeCount <= 1 && selectedInstanceType === "MlxJaccl"
? "MlxRing"
: selectedInstanceType;
try {
const placementResponse = await fetch(
`/instance/placement?model_id=${encodeURIComponent(modelId)}&sharding=${sharding}&instance_meta=${instanceType}&min_nodes=1`,
@@ -783,6 +787,7 @@
quantization?: string;
base_model?: string;
capabilities?: string[];
requires_vllm?: boolean;
}>
>([]);
type ModelMemoryFitStatus =
@@ -886,7 +891,7 @@
}
let selectedSharding = $state<"Pipeline" | "Tensor">("Pipeline");
type InstanceMeta = "MlxRing" | "MlxJaccl";
type InstanceMeta = "MlxRing" | "MlxJaccl" | "Vllm";
// Launch defaults persistence
const LAUNCH_DEFAULTS_KEY = "exo-launch-defaults-v2";
@@ -932,7 +937,12 @@
// Apply sharding and instance type unconditionally
selectedSharding = defaults.sharding;
selectedInstanceType =
defaults.instanceType === "MlxRing" ? "MlxRing" : "MlxJaccl";
defaults.instanceType === "MlxRing"
? "MlxRing"
: defaults.instanceType === "Vllm"
? "Vllm"
: "MlxJaccl";
userPickedInstanceType = true;
// Apply minNodes if valid (between 1 and maxNodes)
if (
@@ -954,6 +964,23 @@
}
let selectedInstanceType = $state<InstanceMeta>("MlxRing");
let userPickedInstanceType = $state(false);
$effect(() => {
if (!userPickedInstanceType && featureFlags()["vllm_available"]) {
selectedInstanceType = "Vllm";
}
});
const selectedModelRequiresVllm = $derived.by((): boolean => {
const id = selectedPreviewModelId();
if (!id) return false;
const model = models.find((m) => m.id === id);
return model?.requires_vllm === true;
});
$effect(() => {
if (selectedModelRequiresVllm) {
selectedInstanceType = "Vllm";
}
});
let selectedMinNodes = $state<number>(1);
let minNodesInitialized = $state(false);
let launchingModelId = $state<string | null>(null);
@@ -1146,9 +1173,7 @@
}
const matchesSelectedRuntime = (runtime: InstanceMeta): boolean =>
selectedInstanceType === "MlxRing"
? runtime === "MlxRing"
: runtime === "MlxJaccl";
runtime === selectedInstanceType;
// Helper to check if a model can be launched (has valid placement with >= minNodes)
function canModelFit(modelId: string): boolean {
@@ -2063,6 +2088,7 @@
let instanceType = "Unknown";
if (instanceTag === "MlxRingInstance") instanceType = "MLX Ring";
else if (instanceTag === "MlxJacclInstance") instanceType = "MLX RDMA";
else if (instanceTag === "VllmInstance") instanceType = "vLLM";
const inst = instance as {
shardAssignments?: {
@@ -5769,14 +5795,18 @@
</div>
<div class="flex gap-2">
<button
disabled={selectedModelRequiresVllm}
onclick={() => {
if (selectedModelRequiresVllm) return;
selectedInstanceType = "MlxRing";
userPickedInstanceType = true;
saveLaunchDefaults();
}}
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 cursor-pointer {selectedInstanceType ===
'MlxRing'
? 'bg-transparent text-exo-yellow border-exo-yellow'
: 'bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 {selectedModelRequiresVllm
? 'opacity-40 cursor-not-allowed bg-transparent text-white/40 border-exo-medium-gray/30'
: selectedInstanceType === 'MlxRing'
? 'cursor-pointer bg-transparent text-exo-yellow border-exo-yellow'
: 'cursor-pointer bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
>
<span
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
@@ -5792,14 +5822,18 @@
TCP/IP
</button>
<button
disabled={selectedModelRequiresVllm}
onclick={() => {
if (selectedModelRequiresVllm) return;
selectedInstanceType = "MlxJaccl";
userPickedInstanceType = true;
saveLaunchDefaults();
}}
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 cursor-pointer {selectedInstanceType ===
'MlxJaccl'
? 'bg-transparent text-exo-yellow border-exo-yellow'
: 'bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 {selectedModelRequiresVllm
? 'opacity-40 cursor-not-allowed bg-transparent text-white/40 border-exo-medium-gray/30'
: selectedInstanceType === 'MlxJaccl'
? 'cursor-pointer bg-transparent text-exo-yellow border-exo-yellow'
: 'cursor-pointer bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
>
<span
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
@@ -5814,7 +5848,41 @@
</span>
RDMA (Fast)
</button>
{#if featureFlags()["vllm_available"] || selectedModelRequiresVllm}
<button
onclick={() => {
selectedInstanceType = "Vllm";
userPickedInstanceType = true;
saveLaunchDefaults();
}}
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 cursor-pointer {selectedInstanceType ===
'Vllm'
? 'bg-transparent text-exo-yellow border-exo-yellow'
: 'bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
>
<span
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
'Vllm'
? 'border-exo-yellow'
: 'border-exo-medium-gray'}"
>
{#if selectedInstanceType === "Vllm"}
<span
class="w-1.5 h-1.5 rounded-full bg-exo-yellow"
></span>
{/if}
</span>
vLLM (CUDA)
</button>
{/if}
</div>
{#if selectedModelRequiresVllm}
<div
class="mt-2 text-[11px] font-mono text-orange-300/80"
>
This model requires vLLM.
</div>
{/if}
</div>
<!-- Minimum Devices -->
@@ -0,0 +1,81 @@
<script lang="ts">
import { browser } from "$app/environment";
import HeaderNav from "$lib/components/HeaderNav.svelte";
import PrefillDecodeDisaggregation from "$lib/components/PrefillDecodeDisaggregation.svelte";
import { featureFlags, refreshState } from "$lib/stores/app.svelte";
import { onMount } from "svelte";
type TabId = "prefill-decode";
const tabs: { id: TabId; label: string }[] = [
{ id: "prefill-decode", label: "Prefill / Decode" },
];
let activeTab = $state<TabId>(tabs[0].id);
let flagsLoaded = $state(false);
onMount(() => {
refreshState().finally(() => {
flagsLoaded = true;
});
});
const flags = $derived(featureFlags());
const enabled = $derived(flags["disaggregation"] === true);
$effect(() => {
if (browser && flagsLoaded && !enabled) {
// No advanced features enabled — bounce home.
window.location.hash = "/";
}
});
</script>
<div class="min-h-screen bg-exo-dark-gray flex flex-col">
<HeaderNav />
<main class="flex-1 max-w-[1100px] mx-auto w-full px-4 md:px-6 py-8">
{#if !flagsLoaded}
<div class="text-exo-light-gray/60 text-sm">Loading…</div>
{:else if !enabled}
<div class="text-exo-light-gray/60 text-sm">
No advanced features enabled. Set <code
class="text-exo-yellow font-mono">ENABLE_DISAGGREGATION=true</code
> on the cluster to access prefill/decode disaggregation.
</div>
{:else}
<div class="mb-4">
<h1
class="text-white text-xl md:text-2xl font-semibold tracking-wide mb-2"
>
Advanced
</h1>
<p class="text-exo-light-gray/60 text-sm">
Cluster-level configuration. Most users don't need anything here.
</p>
</div>
<div
class="flex flex-wrap gap-2 mb-6 border-b border-exo-light-gray/10 pb-3"
>
{#each tabs as tab (tab.id)}
<button
onclick={() => (activeTab = tab.id)}
class="px-3 py-1.5 text-xs rounded-md transition-all cursor-pointer
{activeTab === tab.id
? 'bg-exo-yellow/15 text-exo-yellow border border-exo-yellow/30'
: 'text-exo-light-gray/60 hover:text-white/80 border border-transparent hover:border-exo-light-gray/20'}"
>
{tab.label}
</button>
{/each}
</div>
<div class="space-y-4">
{#if activeTab === "prefill-decode"}
<PrefillDecodeDisaggregation />
{/if}
</div>
{/if}
</main>
</div>
+33 -1
View File
@@ -14,6 +14,7 @@
let modelCapabilities = $state<Record<string, string[]>>({});
let modelContextLengths = $state<Record<string, number>>({});
let modelReasoningDialects = $state<Record<string, string>>({});
const runningModels = $derived.by(() => {
const models: string[] = [];
@@ -132,6 +133,7 @@
for (const modelId of runningModels) {
const caps = modelCapabilities[modelId] || [];
const ctxLen = modelContextLengths[modelId] || 0;
const dialect = modelReasoningDialects[modelId];
const entry: Record<string, unknown> = { name: modelId };
if (ctxLen > 0) {
entry.limit = { context: ctxLen, output: Math.min(ctxLen, 16384) };
@@ -139,6 +141,27 @@
if (caps.includes("vision")) {
entry.modalities = { input: ["text", "image"], output: ["text"] };
}
// Reasoning round-trip: opencode's `interleaved` field tells the
// openai-compatible adapter to send the assistant's prior
// reasoning_content back in subsequent turns. Emit it for dialects
// whose chat templates use prior reasoning:
// - `tool_conditional` (DeepSeek V3.2 / V4): wrapper preserves all
// reasoning when tools are present.
// - `post_last_user` (Qwen3-Thinking, GLM 4.5+, MiniMax M2.x):
// Jinja template reads reasoning_content for assistant turns since
// the last user message — exactly the tool-chain window.
// - `channel` (gpt-oss / Harmony): the model's Jinja template reads
// `message.thinking` rather than `message.reasoning_content`, but
// the server bridges `reasoning_content` → `thinking` before
// rendering, so the round-trip works through the standard field.
// `suffix` (Kimi): reasoning lives in content; no separate field path.
if (
dialect === "tool_conditional" ||
dialect === "post_last_user" ||
dialect === "channel"
) {
entry.interleaved = { field: "reasoning_content" };
}
models[modelId] = entry;
}
if (Object.keys(models).length === 0) {
@@ -350,16 +373,25 @@
try {
const resp = await fetch("/v1/models");
const data = (await resp.json()) as {
data: { id: string; capabilities: string[]; context_length: number }[];
data: {
id: string;
capabilities: string[];
context_length: number;
reasoning_dialect?: string;
}[];
};
const caps: Record<string, string[]> = {};
const ctxs: Record<string, number> = {};
const dialects: Record<string, string> = {};
for (const model of data.data) {
caps[model.id] = model.capabilities || [];
if (model.context_length > 0) ctxs[model.id] = model.context_length;
if (model.reasoning_dialect)
dialects[model.id] = model.reasoning_dialect;
}
modelCapabilities = caps;
modelContextLengths = ctxs;
modelReasoningDialects = dialects;
} catch {
/* ignore */
}
+1 -1
View File
@@ -146,7 +146,7 @@
config.treefmt.build.wrapper
# PYTHON
self'.packages.editableVenv
self'.packages.exo.passthru.evenv
uv
# RUST
+13
View File
@@ -40,6 +40,19 @@ build-app: rust-rebuild sync-clean package
xcodebuild build -project app/EXO/EXO.xcodeproj -scheme EXO -configuration Debug -derivedDataPath app/EXO/build
@echo "\nBuild complete. Run with:\n open {{justfile_directory()}}/app/EXO/build/Build/Products/Debug/EXO.app"
sync-cuda:
#!/usr/bin/env bash
set -euo pipefail
uv sync --extra vllm-cuda13 --extra mlx-cpu --no-install-package vllm
dest=".venv/lib/python3.13/site-packages"
[[ -d $dest/vllm ]] || {
nix build .#exo-cuda-13.passthru.evenv
# will also grab vllm-0.19.1-distinfo
cp -aL result/lib/python3.13/site-packages/vllm* .venv/lib/python3.13/site-packages
chmod -R u+rwX .venv/lib/python3.13/site-packages/vllm*
rm result
}
clean:
rm -rf **/__pycache__
rm -rf target/
+26
View File
@@ -0,0 +1,26 @@
diff --git a/setup.py b/setup.py
index 6dc2ed028..bdcc6354a 100644
--- a/setup.py
+++ b/setup.py
@@ -18,6 +18,13 @@ from setuptools import Extension, setup
from setuptools.command.build_ext import build_ext
+if "NIX_ATTRS_JSON_FILE" in os.environ:
+ with open(os.environ["NIX_ATTRS_JSON_FILE"], "r") as f:
+ NIX_ATTRS = json.load(f)
+else:
+ NIX_ATTRS = { "cmakeFlags": os.environ.get("cmakeFlags", "").split() }
+
+
def load_module_from_path(module_name, path):
spec = importlib.util.spec_from_file_location(module_name, path)
module = importlib.util.module_from_spec(spec)
@@ -184,6 +191,7 @@ class cmake_build_ext(build_ext):
cmake_args = [
"-DCMAKE_BUILD_TYPE={}".format(cfg),
"-DVLLM_TARGET_DEVICE={}".format(VLLM_TARGET_DEVICE),
+ *NIX_ATTRS["cmakeFlags"],
]
verbose = envs.VERBOSE
+57 -34
View File
@@ -15,26 +15,22 @@ dependencies = [
"huggingface-hub>=1.8.0",
"psutil>=7.0.0",
"loguru>=0.7.3",
"exo-pyo3-bindings", # rust bindings
"exo-pyo3-bindings", # rust bindings
"anyio==4.11.0",
"mlx==0.31.2; sys_platform == 'darwin'",
"mlx-lm; sys_platform=='darwin'",
"tiktoken>=0.12.0", # required for kimi k2 tokenizer
"tiktoken>=0.12.0", # required for kimi k2 tokenizer
"hypercorn>=0.18.0",
"openai-harmony>=0.0.8",
"httpx>=0.28.1",
"tomlkit>=0.14.0",
"mflux==0.17.2; sys_platform == 'darwin'",
"python-multipart>=0.0.21",
"msgspec>=0.19.0",
"zstandard>=0.23.0",
"mlx-vlm>=0.3.11",
"transformers>=5.6.2",
"nvidia-ml-py>=13.595.45",
]
[project.scripts]
exo = "exo.main:main"
exo-reasoning-proxy = "exo.reasoning_proxy.main:main"
# dependencies only required for development
[dependency-groups]
@@ -49,23 +45,30 @@ dev = [
[project.optional-dependencies]
build = ["nanobind"]
cpu = [
"mlx==0.31.1; sys_platform == 'linux'",
"mlx-cpu==0.31.1; sys_platform == 'linux'",
"mlx-lm; sys_platform == 'linux'",
"torch>=2.10.0; sys_platform == 'linux'",
mlx-none = ["anyio"]
mlx = [
"mlx==0.31.2",
"mlx-lm",
"mlx-vlm>=0.3.11",
"mflux==0.17.5",
# pinning vllms versions for consistency.
"torch==2.10.0; sys_platform == 'darwin'",
"torch==2.10.0; sys_platform == 'linux'",
"torchaudio==2.10.0; sys_platform == 'darwin'",
"torchaudio==2.10.0; sys_platform == 'linux'",
"torchvision==0.25.0; sys_platform == 'darwin'",
"torchvision==0.25.0; sys_platform == 'linux'",
]
cuda12 = [
"mlx==0.31.1; sys_platform == 'linux'",
"mlx-cuda-12==0.31.1; sys_platform == 'linux'",
"mlx-lm; sys_platform == 'linux'",
"torch>=2.10.0; sys_platform == 'linux'",
]
cuda13 = [
"mlx==0.31.1; sys_platform == 'linux'",
"mlx-cuda-13==0.31.1; sys_platform == 'linux'",
"mlx-lm; sys_platform == 'linux'",
"torch>=2.10.0; sys_platform == 'linux'",
mlx-cpu = ["exo[mlx]", "mlx-cpu==0.31.2; sys_platform == 'linux'"]
mlx-cuda12 = ["exo[mlx]", "mlx-cuda-12==0.31.1; sys_platform == 'linux'"]
mlx-cuda13 = ["exo[mlx]", "mlx-cuda-13==0.31.1; sys_platform == 'linux'"]
vllm-none = ["anyio"]
vllm-cuda13 = [
"vllm[cuda13, fastsafetensors]; sys_platform == 'linux'",
"torch==2.10.0; sys_platform == 'linux'",
"torchaudio==2.10.0; sys_platform == 'linux'",
"torchvision==0.25.0; sys_platform == 'linux'",
]
###
@@ -79,12 +82,23 @@ members = ["rust/exo_pyo3_bindings", "bench"]
exo-pyo3-bindings = { workspace = true }
mlx = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git", branch = "address-rdma-gpu-locks", marker = "sys_platform == 'darwin'" }
mlx-lm = { git = "https://github.com/rltakashige/mlx-lm", branch = "leo/deepseek-v4" }
mflux = { git = "http://github.com/evanev7/mflux", branch = "exo" }
vllm = { git = "http://github.com/evanev7/vllm", branch = "exo2" }
torch = [
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and extra == 'cuda13' and extra != 'cpu' and extra != 'cuda12'" },
{ index = "pytorch-cu120", marker = "sys_platform == 'linux' and extra == 'cuda12' and extra != 'cpu' and extra != 'cuda13'" },
{ index = "pytorch-cpu", marker = "(extra != 'cuda12' and extra != 'cuda13' and sys_platform == 'linux') or sys_platform == 'darwin'" },
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and extra == 'mlx-cpu' and extra != 'vllm-cuda13' and extra != 'mlx-cuda13' and extra != 'mlx-cuda12'" },
{ index = "pytorch-cu128", marker = "sys_platform == 'linux' and extra == 'mlx-cuda12' and extra != 'mlx-cuda13' and extra != 'vllm-cuda13'" },
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and (extra == 'mlx-cuda13' or extra == 'vllm-cuda13')" },
]
torchvision = [
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and extra == 'mlx-cpu' and extra != 'vllm-cuda13' and extra != 'mlx-cuda13' and extra != 'mlx-cuda12'" },
{ index = "pytorch-cu128", marker = "sys_platform == 'linux' and extra == 'mlx-cuda12' and extra != 'mlx-cuda13' and extra != 'vllm-cuda13'" },
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and (extra == 'mlx-cuda13' or extra == 'vllm-cuda13')" },
]
torchaudio = [
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and extra == 'mlx-cpu' and extra != 'vllm-cuda13' and extra != 'mlx-cuda13' and extra != 'mlx-cuda12'" },
{ index = "pytorch-cu128", marker = "sys_platform == 'linux' and extra == 'mlx-cuda12' and extra != 'mlx-cuda13' and extra != 'vllm-cuda13'" },
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and (extra == 'mlx-cuda13' or extra == 'vllm-cuda13')" },
]
vllm = { git = "https://github.com/hmellor/vllm.git", branch = "transformers-v5" }
[[tool.uv.index]]
name = "pytorch-cu130"
@@ -92,8 +106,8 @@ url = "https://download.pytorch.org/whl/cu130"
explicit = true
[[tool.uv.index]]
name = "pytorch-cu120"
url = "https://download.pytorch.org/whl/cu120"
name = "pytorch-cu128"
url = "https://download.pytorch.org/whl/cu128"
explicit = true
[[tool.uv.index]]
@@ -154,11 +168,19 @@ root = "src"
required-version = ">=0.8.6"
prerelease = "allow"
environments = ["sys_platform == 'darwin'", "sys_platform == 'linux'"]
conflicts = [[{ extra = "cuda12" }, { extra = "cuda13" }, { extra = "cpu" }]]
constraint-dependencies = ["transformers>=5.6.2"]
override-dependencies = [
"mlx==0.31.1; sys_platform=='linux'",
"mlx; sys_platform=='darwin'",
override-dependencies = ["opencv-python; python_version < '0'"]
conflicts = [
[
{ extra = "mlx-cuda13" },
{ extra = "mlx-cuda12" },
{ extra = "mlx-cpu" },
{ extra = "mlx-none" },
],
[
{ extra = "vllm-cuda13" },
{ extra = "mlx-cuda12" },
{ extra = "vllm-none" },
],
]
[tool.uv.extra-build-dependencies]
@@ -173,6 +195,7 @@ mlx = [
"ninja",
]
mlx-lm = ["setuptools"]
mflux = ["uv_build"]
xgrammar = [
"nanobind",
"setuptools",
+202 -20
View File
@@ -10,8 +10,10 @@ let
inherit (pkgs.stdenv.hostPlatform) isLinux isDarwin isx86_64;
inherit (pkgs.config) cudaSupport;
inherit (pkgs) cudaPackages;
cuda13Support = cudaSupport && cudaPackages.cudaMajorVersion == "13";
libmlx_source = if cuda13Support then "mlx-cuda-13" else if cudaSupport then "mlx-cuda-12" else "mlx-cpu";
libmlx_source =
if (builtins.elem "mlx-cuda13" members.exo or [ ]) then "mlx-cuda-13"
else if (builtins.elem "mlx-cuda12" members.exo or [ ]) then "mlx-cuda-12"
else "mlx-cpu";
python = pkgs.python313;
cudaLibs = with cudaPackages; [
cuda_cudart
@@ -113,37 +115,213 @@ let
});
} // lib.optionalAttrs isLinux {
mlx = prev.mlx.overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ lib.optionals cudaSupport [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ lib.optionals cudaSupport cudaLibs;
autoPatchelfIgnoreMissingDeps = lib.optionals cudaSupport [ "libcuda.so.1" ];
postInstall = ''
cp -r "${final.${libmlx_source}}/${final.python.sitePackages}/mlx" "$out/${final.python.sitePackages}/mlx/"
'';
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
});
} // lib.optionalAttrs cudaSupport {
"${libmlx_source}" = prev."${libmlx_source}".overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ cudaLibs;
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
});
nvidia-cufile = prev.nvidia-cufile.overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ [ pkgs.rdma-core ];
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
});
nvidia-cusolver = prev.nvidia-cusolver.overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ cudaLibs;
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
});
nvidia-nvshmem-cu13 = prev.nvidia-nvshmem-cu13.overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ [ pkgs.rdma-core pkgs.pmix pkgs.libfabric pkgs.ucx pkgs.openmpi ];
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
});
nvidia-cusparse = prev.nvidia-cusparse.overrideAttrs (old: {
buildInputs = old.buildInputs ++ [ cudaLibs ];
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ cudaLibs;
});
torch = prev.torch.overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ cudaLibs;
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
});
torchaudio = prev.torchaudio.overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
buildInputs = old.buildInputs ++ [ cudaPackages.cuda_cudart ];
preFixup = "addAutoPatchelfSearchPath '${final.torch}'";
});
torchvision = prev.torchvision.overrideAttrs (old: {
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
preFixup = "addAutoPatchelfSearchPath '${final.torch}'";
});
torch-c-dlpack-ext = prev.torch-c-dlpack-ext.overrideAttrs (old: {
buildInputs = old.buildInputs ++ cudaLibs;
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
preFixup = "addAutoPatchelfSearchPath '${final.torch}'";
});
# Currently treating vllm as a cuda dep. it obviously exists as a non cuda dep
vllm = prev.vllm.overrideAttrs (old:
let
cuda_cccl_compat = pkgs.runCommand "cuda-cccl-compat" { } ''
mkdir -p $out/include
ln -s ${cudaPackages.cuda_cccl}/include $out/include/cccl
'';
cudaRoot = pkgs.symlinkJoin {
name = "cuda-merged-exo";
paths = builtins.concatMap (p: [ (lib.getBin p) (lib.getLib p) (lib.getDev p) ]) (cudaLibs ++ [ cudaPackages.cuda_nvcc cuda_cccl_compat ]);
};
cutlass = pkgs.fetchFromGitHub {
name = "cutlass-source";
owner = "NVIDIA";
repo = "cutlass";
tag = "v4.2.1";
hash = "sha256-iP560D5Vwuj6wX1otJhwbvqe/X4mYVeKTpK533Wr5gY=";
};
triton-kernels = pkgs.fetchFromGitHub {
owner = "triton-lang";
repo = "triton";
tag = "v3.6.0";
hash = "sha256-JFSpQn+WsNnh7CAPlcpOcUp0nyKXNbJEANdXqmkt4Tc=";
};
cutlass-flashmla = pkgs.fetchFromGitHub {
owner = "NVIDIA";
repo = "cutlass";
rev = "147f5673d0c1c3dcf66f78d677fd647e4a020219";
hash = "sha256-dHQto08IwTDOIuFUp9jwm1MWkFi8v2YJ/UESrLuG71g=";
};
flashmla = pkgs.stdenv.mkDerivation {
pname = "flashmla";
version = "1.0.0";
src = pkgs.fetchFromGitHub {
name = "FlashMLA-source";
owner = "vllm-project";
repo = "FlashMLA";
rev = "c2afa9cb93e674d5a9120a170a6da57b89267208";
hash = "sha256-pKlwxV6G9iHag/jbu3bAyvYvnu5TbrQwUMFV0AlGC3s=";
};
dontConfigure = true;
buildPhase = ''
rm -rf csrc/cutlass
ln -sf ${cutlass-flashmla} csrc/cutlass
'';
installPhase = ''
cp -rva . $out
'';
};
qutlass = pkgs.fetchFromGitHub {
name = "qutlass-source";
owner = "IST-DASLab";
repo = "qutlass";
rev = "830d2c4537c7396e14a02a46fbddd18b5d107c65";
hash = "sha256-aG4qd0vlwP+8gudfvHwhtXCFmBOJKQQTvcwahpEqC84=";
};
vllm-flash-attn = pkgs.stdenv.mkDerivation {
pname = "vllm-flash-attn";
version = "2.7.2.post1";
src = pkgs.fetchFromGitHub {
name = "flash-attention-source";
owner = "vllm-project";
repo = "flash-attention";
rev = "188be16520ceefdc625fdf71365585d2ee348fe2";
hash = "sha256-Osec+/IF3+UDtbIhDMBXzUeWJ7hDJNb5FpaVaziPSgM=";
};
patches = [
(pkgs.fetchpatch {
url = "https://github.com/Dao-AILab/flash-attention/commit/dad67c88d4b6122c69d0bed1cebded0cded71cea.patch";
hash = "sha256-JSgXWItOp5KRpFbTQj/cZk+Tqez+4mEz5kmH5EUeQN4=";
})
(pkgs.fetchpatch {
url = "https://github.com/Dao-AILab/flash-attention/commit/e26dd28e487117ee3e6bc4908682f41f31e6f83a.patch";
hash = "sha256-NkCEowXSi+tiWu74Qt+VPKKavx0H9JeteovSJKToK9A=";
})
];
dontConfigure = true;
buildPhase = ''
rm -rf csrc/cutlass
ln -sf ${cutlass} csrc/cutlass
'';
installPhase = ''
cp -rva . $out
'';
};
in
{
patches = (old.patches or [ ]) ++ [ ../nix/vllm-setuppy-cmake.patch ];
nativeBuildInputs = (old.nativeBuildInputs or [ ]) ++ [
pkgs.cmake
pkgs.ninja
pkgs.autoAddDriverRunpath
] ++ lib.optionals cudaSupport [
cudaPackages.cuda_nvcc
];
# TODO: vllm rocm/cpu
VLLM_TARGET_DEVICE = "empty";
preConfigure = ''
export MAX_JOBS="$NIX_BUILD_CORES"
'';
# TODO: vllm non cuda13 support, more arch's, etc.
} // lib.optionalAttrs cudaSupport {
buildInputs = cudaLibs ++ [ cudaRoot ];
VLLM_CUDA_VERSION = cudaPackages.cudaMajorMinorVersion;
CUDA_HOME = "${cudaRoot}";
CUDAToolkit_ROOT = "${cudaRoot}";
CUDACXX = "${cudaRoot}/bin/nvcc";
VLLM_CUTLASS_SRC_DIR = "${lib.getDev cutlass}";
VLLM_TARGET_DEVICE = "cuda";
TORCH_CUDA_ARCH_LIST = "12.0;12.1";
TRITON_KERNELS_SRC_DIR = "${lib.getDev triton-kernels}/python/triton_kernels/triton_kernels";
FLASH_MLA_SRC_DIR = "${lib.getDev flashmla}";
QUTLASS_SRC_DIR = "${lib.getDev qutlass}";
VLLM_FLASH_ATTN_SRC_DIR = "${lib.getDev vllm-flash-attn}";
CAFFE2_USE_CUDNN = "ON";
CAFFE2_USE_CUFILE = "ON";
CUTLASS_ENABLE_CUBLAS = "ON";
CUTLASS_NVCC_ARCHS_ENABLED = "12.0;12.1";
cmakeFlags = [
(lib.cmakeBool "CMAKE_SKIP_INSTALL_RPATH" true)
(lib.cmakeBool "CMAKE_BUILD_WITH_INSTALL_RPATH" true)
(lib.cmakeFeature "CUDA_HOME" "${cudaRoot}")
(lib.cmakeFeature "CUDAToolkit_ROOT" "${cudaRoot}")
(lib.cmakeFeature "CMAKE_CUDA_COMPILER" "${cudaRoot}/bin/nvcc")
(lib.cmakeFeature "CMAKE_PREFIX_PATH" "${cudaRoot}")
(lib.cmakeFeature "FETCHCONTENT_SOURCE_DIR_CUTLASS" "${lib.getDev cutlass}")
(lib.cmakeFeature "FLASH_MLA_SRC_DIR" "${lib.getDev flashmla}")
(lib.cmakeFeature "VLLM_FLASH_ATTN_SRC_DIR" "${lib.getDev vllm-flash-attn}")
(lib.cmakeFeature "QUTLASS_SRC_DIR" "${lib.getDev qutlass}")
(lib.cmakeFeature "TORCH_CUDA_ARCH_LIST" "12.0;12.1")
(lib.cmakeFeature "CUTLASS_NVCC_ARCHS_ENABLED" "${cudaPackages.flags.cmakeCudaArchitecturesString}")
(lib.cmakeFeature "CUDA_TOOLKIT_ROOT_DIR" "${cudaRoot}")
(lib.cmakeFeature "CAFFE2_USE_CUDNN" "ON")
(lib.cmakeFeature "CAFFE2_USE_CUFILE" "ON")
(lib.cmakeFeature "CUTLASS_ENABLE_CUBLAS" "ON")
];
});
} // lib.optionalAttrs (cudaSupport && isx86_64) {
numba = prev.numba.overrideAttrs (old: {
buildInputs = (old.buildInputs or [ ]) ++ [ pkgs.tbb ];
});
};
pyprojectOverlay = workspace.mkPyprojectOverlay {
sourcePreference = "wheel";
@@ -164,24 +342,28 @@ let
buildSystemsOverlay
]
);
venv = name: (pythonSet.mkVirtualEnv "${name}-env" members).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" ]; });
mkApp = cmd: name: pkgs.writeShellApplication {
# mlx and mlx-cuda ship clashing cmake files - we dont need them at runtime anyway
venv = name: (pythonSet.mkVirtualEnv "${name}-venv" members).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" "lib/python${python.pythonVersion}/site-packages/build_backend.py" ]; });
mkApp = text: name: pkgs.writeShellApplication {
inherit name;
text = "exec " + lib.optionalString cudaSupport "nixglhost " + text;
runtimeEnv = {
EXO_DASHBOARD_DIR = self'.packages.dashboard;
EXO_RESOURCES_DIR = inputs.self + /resources;
};
runtimeInputs = [
# mlx and mlx-cuda ship clashing cmake files - we dont need them at runtime anyway
(venv name)
pkgs.nix-gl-host
]
++ lib.optionals isDarwin [ pkgs.macmon ];
text = "exec " + lib.optionalString cudaSupport "${lib.getExe pkgs.nix-gl-host} " + cmd;
passthru = {
venv = venv name;
evenv = ((pythonSet.overrideScope editableOverlay).mkVirtualEnv "${name}-evenv" (members // { exo = (members.exo or [ ]) ++ [ "dev" ]; })).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" "lib/python${python.pythonVersion}/site-packages/build_backend.py" ]; });
};
};
in
{
inherit venv;
editablePythonSet = pythonSet.overrideScope editableOverlay;
mkPythonScript = path: mkApp ''python ${path} "$@"'';
mkExo = mkApp ''exo "$@"'';
};
@@ -191,18 +373,18 @@ in
{ self', pkgs, unfreePkgs, lib, ... }:
let
inherit (pkgs.stdenv.hostPlatform) isLinux;
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "cpu" ]; }; }) editablePythonSet mkExo;
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "mlx-cpu" "vllm-none" ]; }; }) mkExo;
# Virtual environment with dev dependencies for testing
testVenv = (mkPythonSet {
inherit self' pkgs lib; members = {
exo = [ "dev" "cpu" ]; # Include pytest, pytest-asyncio, pytest-env
exo = [ "dev" "mlx-cpu" "vllm-none" ]; # Include pytest, pytest-asyncio, pytest-env
};
}).venv "exo-test";
mkBenchScript = (mkPythonSet {
inherit self' pkgs lib; members = {
exo = [ "cpu" ];
exo = [ "mlx-cpu" "vllm-none" ];
exo-bench = [ ]; # Include pytest, pytest-asyncio, pytest-env
};
}).mkPythonScript;
@@ -212,12 +394,12 @@ in
runtimeInputs = [ pkgs.python313 ];
text = ''exec python ${path} "$@"'';
};
cuda12Set = mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; members = { exo = [ "mlx-cuda12" "vllm-none" ]; }; };
cuda13Set = mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; members = { exo = [ "mlx-cpu" "vllm-cuda13" ]; }; };
in
{
packages = {
exo = mkExo "exo";
editableVenv = editablePythonSet.mkVirtualEnv "exo-dev-env" { exo = [ "dev" ]; };
# for running tests in ci
exo-test-env = testVenv;
exo-bench = mkBenchScript "exo-bench" (inputs.self + /bench/exo_bench.py);
@@ -226,8 +408,8 @@ in
# used by ./tests/run_exo_on.sh
exo-get-all-models-on-cluster = mkSimplePythonScript "exo-get-all-models-on-cluster" (inputs.self + /tests/get_all_models_on_cluster.py);
} // lib.optionalAttrs isLinux {
exo-cuda-12 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; members = { exo = [ "cuda12" ]; }; }).mkExo "exo-cuda-12";
exo-cuda-13 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; members = { exo = [ "cuda13" ]; }; }).mkExo "exo-cuda-13";
exo-cuda-12 = cuda12Set.mkExo "exo-cuda-12";
exo-cuda-13 = cuda13Set.mkExo "exo-cuda-13";
};
checks = {
@@ -0,0 +1,21 @@
model_id = "2imi9/gpt-oss-20B-NVFP4A16-BF16"
n_layers = 24
hidden_size = 2880
num_key_value_heads = 8
supports_tensor = false
tasks = ["TextGeneration"]
family = "gpt-oss"
quantization = "nvfp4"
base_model = "GPT-OSS 20B"
capabilities = ["text", "thinking"]
reasoning_dialect = "channel"
context_length = 131072
requires_vllm = true
[storage_size]
in_bytes = 41829514752
[sampling_defaults]
temperature = 1.0
top_p = 1.0
top_k = 0
@@ -8,6 +8,7 @@ family = "deepseek"
quantization = "4bit"
base_model = "DeepSeek V3.1"
capabilities = ["text", "thinking", "thinking_toggle"]
reasoning_dialect = "post_last_user"
context_length = 131072
@@ -8,6 +8,7 @@ family = "deepseek"
quantization = "8bit"
base_model = "DeepSeek V3.1"
capabilities = ["text", "thinking", "thinking_toggle"]
reasoning_dialect = "post_last_user"
context_length = 131072
@@ -8,6 +8,7 @@ family = "deepseek"
quantization = "4bit"
base_model = "DeepSeek V3.2"
capabilities = ["text", "thinking", "thinking_toggle"]
reasoning_dialect = "tool_conditional"
context_length = 131072
@@ -8,6 +8,7 @@ family = "deepseek"
quantization = "8bit"
base_model = "DeepSeek V3.2"
capabilities = ["text", "thinking", "thinking_toggle"]
reasoning_dialect = "tool_conditional"
context_length = 131072
@@ -8,6 +8,7 @@ family = "deepseek"
quantization = "8bit"
base_model = "DeepSeek V4 Flash"
capabilities = ["text", "thinking", "thinking_toggle"]
reasoning_dialect = "tool_conditional"
context_length = 1048576
@@ -8,6 +8,7 @@ family = "deepseek"
quantization = "8bit"
base_model = "DeepSeek V4 Pro"
capabilities = ["text", "thinking", "thinking_toggle"]
reasoning_dialect = "tool_conditional"
context_length = 1048576
@@ -0,0 +1,27 @@
model_id = "nvidia/Qwen3-30B-A3B-NVFP4"
n_layers = 48
hidden_size = 2048
num_key_value_heads = 4
supports_tensor = false
tasks = ["TextGeneration"]
family = "qwen"
quantization = "nvfp4"
base_model = "Qwen3 30B"
capabilities = ["text", "thinking", "thinking_toggle"]
context_length = 32768
requires_vllm = true
[storage_size]
in_bytes = 18087458688
[sampling_defaults]
temperature = 0.6
top_p = 0.95
top_k = 20
min_p = 0.0
[sampling_defaults.non_thinking]
temperature = 0.7
top_p = 0.8
top_k = 20
min_p = 0.0
@@ -0,0 +1,20 @@
model_id = "openai/gpt-oss-120b"
n_layers = 36
hidden_size = 2880
num_key_value_heads = 8
supports_tensor = false
tasks = ["TextGeneration"]
family = "gpt-oss"
quantization = "mxfp4"
base_model = "GPT-OSS 120B"
capabilities = ["text", "thinking"]
reasoning_dialect = "channel"
context_length = 131072
[storage_size]
in_bytes = 65248815744
[sampling_defaults]
temperature = 1.0
top_p = 1.0
top_k = 0
@@ -0,0 +1,32 @@
model_id = "sakamakismile/Qwen3.6-27B-NVFP4"
n_layers = 64
hidden_size = 5120
num_key_value_heads = 4
supports_tensor = false
tasks = ["TextGeneration"]
family = "qwen"
quantization = "nvfp4"
base_model = "Qwen3.6 27B"
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
reasoning_dialect = "post_last_user"
context_length = 262144
requires_vllm = true
[storage_size]
in_bytes = 16703361232
[sampling_defaults]
temperature = 1.0
top_p = 0.95
top_k = 20
min_p = 0.0
repetition_penalty = 1.0
presence_penalty = 1.5
[sampling_defaults.non_thinking]
temperature = 0.7
top_p = 0.8
top_k = 20
min_p = 0.0
repetition_penalty = 1.0
presence_penalty = 1.5
+222
View File
@@ -0,0 +1,222 @@
#!/usr/bin/env python
"""Standalone smoke test for VllmEngine.serve_prefill.
Loads a real vLLM engine, runs serve_prefill against an in-memory buffer
twice in a row with the same prompt, and verifies both runs produce a
well-formed wire stream (header -> KV chunks -> Done).
The second run is the regression guard: with vLLM APC enabled this would
trip the chunked-prefill + APC + custom kv-connector CUDA assert
(`vectorized_gather_kernel: ind >= ind_dim_size`) and the server would
close the socket before the Done frame.
Usage on the Spark (gx10-de89):
cd /home/larry/exo
/nix/store/2b82iz9ac0pxqafrgxmgdkq8sr2hwlx6-exo-cuda-13-venv/bin/python \\
scripts/check_serve_prefill.py Qwen/Qwen3-0.6B
Exits 0 on success, non-zero with a diagnostic on failure.
"""
from __future__ import annotations
import contextlib
import io
import os
import sys
import traceback
from pathlib import Path
from typing import cast
def _ensure_repo_on_path() -> None:
repo = Path(__file__).resolve().parent.parent
src = repo / "src"
if str(src) not in sys.path:
sys.path.insert(0, str(src))
_ensure_repo_on_path()
from exo.shared.types.common import ModelId # noqa: E402
from exo.worker.disaggregated.protocol import ( # noqa: E402
ArraysState,
Done,
ErrorMessage,
KVChunk,
read_header,
read_message,
)
from exo.worker.disaggregated.server import PrefillRequest # noqa: E402
def _make_token_ids(n: int) -> list[int]:
return [(i * 1009 + 17) % 30000 + 100 for i in range(n)]
def _decode(
payload: bytes,
) -> tuple[list[KVChunk], list[ArraysState], Done | None, ErrorMessage | None]:
buf = io.BytesIO(payload)
_ = read_header(buf)
chunks: list[KVChunk] = []
arrays: list[ArraysState] = []
done: Done | None = None
error: ErrorMessage | None = None
while True:
msg = read_message(buf)
if msg is None:
break
if isinstance(msg, KVChunk):
chunks.append(msg)
elif isinstance(msg, ArraysState):
arrays.append(msg)
elif isinstance(msg, Done):
done = msg
break
elif isinstance(msg, ErrorMessage):
error = msg
break
return chunks, arrays, done, error
def _build_engine(model_id: ModelId) -> object:
from exo.worker.engines.vllm.engine import VllmEngine
from exo.worker.engines.vllm.generator import VllmBatchEngine, load_vllm_engine
from exo.worker.engines.vllm.kv_connector import (
ExoKVProducerConnector,
_patch_gdn_capture,
_patch_vllm_for_connector,
)
_patch_vllm_for_connector(ExoKVProducerConnector)
_patch_gdn_capture()
llm_engine, tool_parser = load_vllm_engine(
model_id=model_id,
trust_remote_code=False,
n_layers=1,
kv_connector_cls=ExoKVProducerConnector,
)
gen = VllmBatchEngine(engine=llm_engine, model_id=model_id)
class _S:
def send(self, _: object) -> None: ...
class _R:
def collect(self) -> list[object]:
return []
return VllmEngine(
tool_parser=tool_parser,
model_id=model_id,
cancel_receiver=cast("object", _R()), # pyright: ignore[reportArgumentType]
event_sender=cast("object", _S()), # pyright: ignore[reportArgumentType]
_gen=gen,
max_concurrent_requests=1,
)
def _run_one(engine: object, n_tokens: int, label: str) -> int:
request = PrefillRequest(
request_id=f"check-{label}-{os.getpid()}",
model_id="ignored",
token_ids=_make_token_ids(n_tokens),
start_pos=0,
use_prefix_cache=True,
)
buf = io.BytesIO()
engine.serve_prefill(request, buf) # pyright: ignore[reportAttributeAccessIssue]
payload = buf.getvalue()
if not payload:
raise AssertionError(f"{label}: server wrote nothing")
chunks, arrays, done, error = _decode(payload)
if error is not None:
raise AssertionError(
f"{label}: server returned ErrorMessage [{error.code}]: {error.message}"
)
if done is None:
raise AssertionError(
f"{label}: stream did not end with Done "
f"({len(chunks)} kv chunks, {len(arrays)} arrays)"
)
if done.total_tokens <= 0:
raise AssertionError(f"{label}: Done reported {done.total_tokens} tokens")
if not chunks:
raise AssertionError(f"{label}: no KV chunks shipped")
expected = max(0, n_tokens - 2)
if done.total_tokens < expected - 64:
raise AssertionError(
f"{label}: got {done.total_tokens} tokens, expected ~{expected}"
)
print(
f" [{label}] OK: tokens={done.total_tokens} "
f"kv_chunks={len(chunks)} arrays={len(arrays)}"
)
return done.total_tokens
def main(argv: list[str]) -> int:
if len(argv) < 2:
print(__doc__)
return 2
model_id = ModelId(argv[1])
from exo.download.download_utils import build_model_path
model_path = build_model_path(model_id)
if not model_path.exists():
print(f"FAIL: model {model_id} not found at {model_path}")
return 1
print(f"Loading vLLM engine for {model_id} ({model_path}) ...")
engine = _build_engine(model_id)
failures: list[str] = []
try:
try:
t1 = _run_one(engine, n_tokens=512, label="run1-fresh")
except AssertionError as e:
failures.append(f"run1: {e}")
t1 = 0
try:
t2 = _run_one(engine, n_tokens=512, label="run2-same-prompt")
except AssertionError as e:
failures.append(f"run2: {e}")
t2 = 0
if t1 and t2 and t1 != t2:
failures.append(
f"run1 returned {t1} tokens but run2 returned {t2} (should match)"
)
try:
ta = _run_one(engine, n_tokens=256, label="run3-shorter")
tb = _run_one(engine, n_tokens=768, label="run4-longer")
if ta and tb and tb <= ta:
failures.append(
f"longer prompt should produce more tokens: 256->{ta} 768->{tb}"
)
except AssertionError as e:
failures.append(f"length-variation: {e}")
finally:
with contextlib.suppress(Exception):
engine.close() # pyright: ignore[reportAttributeAccessIssue]
if failures:
print()
print("FAIL")
for f in failures:
print(f" - {f}")
return 1
print()
print("PASS")
return 0
if __name__ == "__main__":
try:
sys.exit(main(sys.argv))
except Exception:
traceback.print_exc()
sys.exit(1)
+124
View File
@@ -0,0 +1,124 @@
#!/usr/bin/env bash
set -Eeuo pipefail
SELF_IP="169.254.100.1"
PEER_IP="169.254.100.2"
PREFIX="16"
IFACE="enP7s7"
USE_NM="auto"
DRY_RUN=0
usage() {
cat <<EOF
Usage: sudo $(basename "$0") [options]
Configure a Linux Ethernet interface with a static IPv4 for a host-to-host
link to a Mac peer.
Defaults: this host = ${SELF_IP}/${PREFIX}, peer = ${PEER_IP}, iface = ${IFACE}.
Options:
--iface IFACE Default: ${IFACE}
--self-ip IP Default: ${SELF_IP}
--peer-ip IP For verification ping. Default: ${PEER_IP}
--prefix N Default: ${PREFIX}
--no-nm Use 'ip addr' directly (transient, no NetworkManager).
--dry-run Print actions without applying.
-h, --help Show this help.
EOF
}
while (($#)); do
case "$1" in
--iface)
shift
IFACE="${1:?}"
;;
--self-ip)
shift
SELF_IP="${1:?}"
;;
--peer-ip)
shift
PEER_IP="${1:?}"
;;
--prefix)
shift
PREFIX="${1:?}"
;;
--no-nm) USE_NM=no ;;
--dry-run) DRY_RUN=1 ;;
-h | --help)
usage
exit 0
;;
*)
echo "Unknown arg: $1" >&2
usage >&2
exit 1
;;
esac
shift
done
[[ $EUID -eq 0 ]] || {
echo "Run as root." >&2
exit 1
}
run() {
printf '+'
printf ' %q' "$@"
printf '\n'
((DRY_RUN)) || "$@"
}
ip link show "$IFACE" >/dev/null 2>&1 || {
echo "Interface $IFACE does not exist." >&2
exit 1
}
if [[ $USE_NM == "auto" ]]; then
if command -v nmcli >/dev/null 2>&1 && systemctl is-active --quiet NetworkManager 2>/dev/null; then
USE_NM=yes
else
USE_NM=no
fi
fi
if [[ $USE_NM == "yes" ]]; then
CONN="$(nmcli -g GENERAL.CONNECTION device show "$IFACE" 2>/dev/null | head -n1 || true)"
if [[ -z $CONN || $CONN == "--" ]]; then
CONN="static-${IFACE}"
run nmcli connection add type ethernet ifname "$IFACE" con-name "$CONN"
fi
run nmcli connection modify "$CONN" \
connection.interface-name "$IFACE" \
connection.autoconnect yes \
connection.autoconnect-priority 100 \
ipv4.method manual \
ipv4.addresses "${SELF_IP}/${PREFIX}" \
ipv4.gateway "" \
ipv4.dns "" \
ipv4.never-default yes \
ipv6.method link-local \
ipv6.addr-gen-mode stable-privacy
run nmcli connection up "$CONN"
else
run ip link set "$IFACE" up
run ip addr flush dev "$IFACE"
run ip addr add "${SELF_IP}/${PREFIX}" dev "$IFACE"
fi
if ((!DRY_RUN)); then
printf '\n'
ip -br addr show "$IFACE"
printf '\n'
if ping -c2 -W2 "$PEER_IP" >/dev/null 2>&1; then
echo "OK: $PEER_IP reachable on $IFACE."
else
echo "WARN: $PEER_IP not reachable yet."
echo " Verify the peer is configured (run setup_linklocal_mac.sh on the Mac)."
echo " ip neigh show dev $IFACE # check for the peer MAC"
fi
fi
+170
View File
@@ -0,0 +1,170 @@
#!/usr/bin/env bash
set -Eeuo pipefail
SELF_IP="169.254.100.2"
PEER_IP="169.254.100.1"
NETMASK="255.255.0.0"
IFACE=""
DRY_RUN=0
usage() {
cat <<EOF
Usage: sudo $(basename "$0") [options]
Configure a Mac Ethernet interface with a static IPv4 for a host-to-host link
to the DGX/GX10 peer.
Defaults: this Mac = ${SELF_IP}, peer = ${PEER_IP}, mask = ${NETMASK}.
Options:
--iface IFACE Interface (e.g. en12). Default: auto-detect.
--self-ip IP This Mac's address. Default: ${SELF_IP}.
--peer-ip IP Peer for verification ping. Default: ${PEER_IP}.
--netmask MASK Default: ${NETMASK}.
--dry-run Print actions without applying.
-h, --help Show this help.
EOF
}
while (($#)); do
case "$1" in
--iface)
shift
IFACE="${1:?}"
;;
--self-ip)
shift
SELF_IP="${1:?}"
;;
--peer-ip)
shift
PEER_IP="${1:?}"
;;
--netmask)
shift
NETMASK="${1:?}"
;;
--dry-run) DRY_RUN=1 ;;
-h | --help)
usage
exit 0
;;
*)
echo "Unknown arg: $1" >&2
usage >&2
exit 1
;;
esac
shift
done
[[ $EUID -eq 0 ]] || {
echo "Run with sudo." >&2
exit 1
}
run() {
printf '+'
printf ' %q' "$@"
printf '\n'
((DRY_RUN)) || "$@"
}
target_subnet_prefix() {
local ip="$1"
printf '%s.' "${ip%.*}"
}
iface_score() {
local iface="$1" info subnet
info="$(ifconfig "$iface" 2>/dev/null || true)"
[[ -n $info ]] || {
echo 0
return
}
grep -q 'status: active' <<<"$info" || {
echo 0
return
}
subnet="$(target_subnet_prefix "$SELF_IP")"
if grep -qE "inet ${subnet//./\\.}" <<<"$info"; then
echo 100
return
fi
if grep -qE 'inet 169\.254\.' <<<"$info"; then
echo 80
return
fi
if ! grep -qE '^[[:space:]]*inet ' <<<"$info"; then
echo 60
return
fi
echo 10
}
detect_iface() {
local best="" best_score=0 iface score
for iface in $(ifconfig -l); do
[[ $iface =~ ^en[0-9]+$ ]] || continue
score="$(iface_score "$iface")"
if ((score > best_score)); then
best="$iface"
best_score="$score"
fi
done
((best_score >= 60)) || return 1
printf '%s\n' "$best"
}
iface_to_service() {
local iface="$1" line port=""
while IFS= read -r line; do
if [[ $line == "Hardware Port: "* ]]; then
port="${line#Hardware Port: }"
elif [[ $line == "Device: $iface" ]]; then
printf '%s\n' "$port"
return 0
fi
done < <(networksetup -listallhardwareports)
return 1
}
if [[ -z $IFACE ]]; then
IFACE="$(detect_iface || true)"
[[ -n $IFACE ]] || {
echo "Could not auto-detect a wired interface. Pass --iface enX." >&2
echo "Active interfaces:" >&2
ifconfig -l | tr ' ' '\n' | grep -E '^en[0-9]+$' | while read -r i; do
printf ' %-6s %s\n' "$i" "$(ifconfig "$i" | grep -E 'status:|inet ' | tr '\n' ' ')" >&2
done
exit 1
}
echo "Auto-detected interface: $IFACE"
fi
ifconfig "$IFACE" >/dev/null 2>&1 || {
echo "Interface $IFACE does not exist." >&2
exit 1
}
SERVICE="$(iface_to_service "$IFACE" || true)"
[[ -n $SERVICE ]] || {
echo "No network service maps to $IFACE. Check System Settings -> Network." >&2
exit 1
}
echo "Network service: $SERVICE"
run networksetup -setmanual "$SERVICE" "$SELF_IP" "$NETMASK" ""
if ((!DRY_RUN)); then
printf '\n'
ifconfig "$IFACE" | grep -E 'inet |status:'
printf '\n'
if ping -c2 -t3 "$PEER_IP" >/dev/null 2>&1; then
echo "OK: $PEER_IP reachable on $IFACE."
else
echo "WARN: $PEER_IP not reachable yet."
echo " Verify the peer is configured (run setup_linklocal_dgx.sh on the GX10)."
echo " arp -an -i $IFACE # check for the peer MAC"
fi
fi
+9 -3
View File
@@ -131,9 +131,13 @@ async def chat_request_to_text_generation(
multimodal_content.append({"type": "text", "text": part.text})
else:
multimodal_content.append({"type": "image"})
chat_template_messages.append(
{"role": msg.role, "content": multimodal_content}
)
multimodal_msg: dict[str, Any] = {
"role": msg.role,
"content": multimodal_content,
}
if msg.reasoning_content is not None:
multimodal_msg["reasoning_content"] = msg.reasoning_content
chat_template_messages.append(multimodal_msg)
continue
msg_copy = msg.model_copy(update={"content": content})
@@ -168,6 +172,8 @@ async def chat_request_to_text_generation(
min_p=request.min_p,
repetition_penalty=request.repetition_penalty,
repetition_context_size=request.repetition_context_size,
presence_penalty=request.presence_penalty,
frequency_penalty=request.frequency_penalty,
images=images,
)
+91 -4
View File
@@ -79,6 +79,8 @@ from exo.api.types import (
ImageListItem,
ImageListResponse,
ImageSize,
InstanceLinkBody,
InstanceLinkResponse,
ModelList,
ModelListModel,
PlaceInstanceParams,
@@ -122,6 +124,7 @@ from exo.master.placement import place_instance as get_instance_placements
from exo.shared.apply import apply
from exo.shared.constants import (
DASHBOARD_DIR,
ENABLE_DISAGGREGATION,
EXO_CACHE_HOME,
EXO_EVENT_LOG_DIR,
EXO_IMAGE_CACHE_DIR,
@@ -154,6 +157,7 @@ from exo.shared.types.commands import (
DeleteCustomModelCard,
DeleteDownload,
DeleteInstance,
DeleteInstanceLink,
DownloadCommand,
ForwarderCommand,
ForwarderDownloadCommand,
@@ -161,6 +165,7 @@ from exo.shared.types.commands import (
ImageGeneration,
PlaceInstance,
SendInputChunk,
SetInstanceLink,
StartDownload,
TaskCancelled,
TaskFinished,
@@ -174,6 +179,7 @@ from exo.shared.types.events import (
InstanceDeleted,
TracesMerged,
)
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.memory import Memory
from exo.shared.types.state import State
from exo.shared.types.tasks import (
@@ -215,6 +221,17 @@ def _ensure_seed(params: AdvancedImageParams | None) -> AdvancedImageParams:
return params
def _require_disaggregation_enabled() -> None:
if not ENABLE_DISAGGREGATION:
raise HTTPException(
status_code=HTTPStatus.NOT_FOUND,
detail=(
"Prefill/decode disaggregation is disabled. "
"Set ENABLE_DISAGGREGATION=true to enable."
),
)
class API:
def __init__(
self,
@@ -328,6 +345,11 @@ class API:
self.app.get("/instance/previews")(self.get_placement_previews)
self.app.get("/instance/{instance_id}")(self.get_instance)
self.app.delete("/instance/{instance_id}")(self.delete_instance)
self.app.get("/v1/instance-links")(self.list_instance_links)
self.app.post("/v1/instance-links")(self.create_instance_link)
self.app.put("/v1/instance-links/{link_id}")(self.update_instance_link)
self.app.delete("/v1/instance-links/{link_id}")(self.delete_instance_link)
self.app.get("/v1/feature-flags")(self.get_feature_flags)
self.app.get("/models")(self.get_models)
self.app.get("/v1/models")(self.get_models)
self.app.post("/models/add")(self.add_custom_model)
@@ -336,7 +358,9 @@ class API:
self.app.post("/v1/chat/completions", response_model=None)(
self.chat_completions
)
self.app.post("/bench/chat/completions")(self.bench_chat_completions)
self.app.post("/bench/chat/completions", response_model=None)(
self.bench_chat_completions
)
self.app.post("/v1/images/generations", response_model=None)(
self.image_generations
)
@@ -502,8 +526,8 @@ class API:
)
]
)
# TODO: PDD
# instance_combinations.append((Sharding.PrefillDecodeDisaggregation, InstanceMeta.MlxRing, 1))
if any(self.state.node_vllm.values()):
instance_combinations.append((Sharding.Pipeline, InstanceMeta.Vllm, 1))
for sharding, instance_meta, min_nodes in instance_combinations:
try:
@@ -615,6 +639,52 @@ class API:
instance_id=instance_id,
)
async def get_feature_flags(self) -> dict[str, bool]:
return {
"disaggregation": ENABLE_DISAGGREGATION,
"vllm_available": any(self.state.node_vllm.values()),
}
async def list_instance_links(self) -> list[InstanceLink]:
if not ENABLE_DISAGGREGATION:
return []
return list(self.state.instance_links.values())
async def create_instance_link(
self, body: InstanceLinkBody
) -> InstanceLinkResponse:
_require_disaggregation_enabled()
return await self._set_instance_link(InstanceLinkId(), body)
async def update_instance_link(
self, link_id: InstanceLinkId, body: InstanceLinkBody
) -> InstanceLinkResponse:
_require_disaggregation_enabled()
return await self._set_instance_link(link_id, body)
async def _set_instance_link(
self, link_id: InstanceLinkId, body: InstanceLinkBody
) -> InstanceLinkResponse:
command = SetInstanceLink(
link_id=link_id,
prefill_instances=list(body.prefill_instances),
decode_instances=list(body.decode_instances),
)
await self._send(command)
return InstanceLinkResponse(
message="Command received.", command_id=command.command_id
)
async def delete_instance_link(
self, link_id: InstanceLinkId
) -> InstanceLinkResponse:
_require_disaggregation_enabled()
command = DeleteInstanceLink(link_id=link_id)
await self._send(command)
return InstanceLinkResponse(
message="Command received.", command_id=command.command_id
)
async def cancel_command(self, command_id: CommandId) -> CancelCommandResponse:
"""Cancel an active command by closing its stream and notifying workers."""
sender = self._text_generation_queues.get(
@@ -829,7 +899,7 @@ class API:
async def bench_chat_completions(
self, payload: BenchChatCompletionRequest
) -> BenchChatCompletionResponse:
) -> BenchChatCompletionResponse | StreamingResponse:
task_params = await chat_request_to_text_generation(payload)
resolved_model = await self._resolve_and_validate_text_model(
ModelId(task_params.model)
@@ -846,6 +916,22 @@ class API:
command = await self._send_text_generation_with_images(task_params)
if payload.stream:
return StreamingResponse(
with_sse_keepalive(
generate_chat_stream(
command.command_id,
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 self._collect_text_generation_with_stats(command.command_id)
async def _resolve_and_validate_text_model(self, model_id: ModelId) -> ModelId:
@@ -1665,6 +1751,7 @@ class API:
capabilities=card.capabilities,
reasoning_dialect=card.reasoning_dialect,
context_length=card.context_length,
requires_vllm=card.requires_vllm,
)
for card in cards
]
+2
View File
@@ -34,6 +34,8 @@ from .api import ImageGenerationTaskParams as ImageGenerationTaskParams
from .api import ImageListItem as ImageListItem
from .api import ImageListResponse as ImageListResponse
from .api import ImageSize as ImageSize
from .api import InstanceLinkBody as InstanceLinkBody
from .api import InstanceLinkResponse as InstanceLinkResponse
from .api import Logprobs as Logprobs
from .api import LogprobsContentItem as LogprobsContentItem
from .api import ModelList as ModelList
+11
View File
@@ -49,6 +49,7 @@ class ModelListModel(BaseModel):
base_model: str = Field(default="")
capabilities: list[str] = Field(default_factory=list)
reasoning_dialect: ReasoningDialect = "none"
requires_vllm: bool = Field(default=False)
class ModelList(BaseModel):
@@ -296,6 +297,16 @@ class CancelCommandResponse(BaseModel):
command_id: CommandId
class InstanceLinkBody(BaseModel):
prefill_instances: list[InstanceId]
decode_instances: list[InstanceId]
class InstanceLinkResponse(BaseModel):
message: str
command_id: CommandId
ImageSize = Literal[
"auto",
"512x512",
+109 -7
View File
@@ -10,6 +10,7 @@ from exo.master.placement import (
get_transition_events,
place_instance,
)
from exo.master.placement_utils import find_ip_prioritised
from exo.shared.apply import apply
from exo.shared.constants import EXO_EVENT_LOG_DIR, EXO_TRACING_ENABLED
from exo.shared.types.commands import (
@@ -17,6 +18,7 @@ from exo.shared.types.commands import (
CreateInstance,
DeleteCustomModelCard,
DeleteInstance,
DeleteInstanceLink,
ForwarderCommand,
ForwarderDownloadCommand,
ImageEdits,
@@ -24,6 +26,7 @@ from exo.shared.types.commands import (
PlaceInstance,
RequestEventLog,
SendInputChunk,
SetInstanceLink,
TaskCancelled,
TaskFinished,
TestCommand,
@@ -38,6 +41,8 @@ from exo.shared.types.events import (
IndexedEvent,
InputChunkReceived,
InstanceDeleted,
InstanceLinkCreated,
InstanceLinkDeleted,
LocalForwarderEvent,
NodeGatheredInfo,
NodeTimedOut,
@@ -48,6 +53,7 @@ from exo.shared.types.events import (
TracesCollected,
TracesMerged,
)
from exo.shared.types.instance_link import InstanceLink
from exo.shared.types.state import State
from exo.shared.types.tasks import (
ImageEdits as ImageEditsTask,
@@ -69,6 +75,46 @@ from exo.utils.event_buffer import MultiSourceBuffer
from exo.utils.task_group import TaskGroup
def _prefill_endpoint_for(state: State, decode_instance_id: InstanceId) -> str | None:
decode = state.instances.get(decode_instance_id)
if decode is None:
return None
decode_node = next(iter(decode.shard_assignments.node_to_runner.keys()), None)
if decode_node is None:
return None
sources: set[InstanceId] = set()
for link in state.instance_links.values():
if decode_instance_id in link.decode_instances:
sources.update(link.prefill_instances)
sources.discard(decode_instance_id)
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_counts: dict[InstanceId, int] = {
src_id: sum(
1
for task in state.tasks.values()
if task.instance_id == src_id and task.task_status in in_flight
)
for src_id in sources
}
for src_id in sorted(sources, key=lambda sid: task_counts[sid]):
instance = state.instances.get(src_id)
if instance is None:
continue
for node_id, runner_id in instance.shard_assignments.node_to_runner.items():
port = state.prefill_server_ports.get(runner_id)
if port is None:
continue
ip = find_ip_prioritised(
decode_node, node_id, state.topology, state.node_network, ring=True
)
if ip is None:
continue
return f"{ip}:{port}"
return None
class Master:
def __init__(
self,
@@ -128,15 +174,45 @@ class Master:
case TestCommand():
pass
case TextGeneration():
for instance in self.state.instances.values():
if (
instance.shard_assignments.model_id
== command.task_params.model
prefill_only: set[InstanceId] = set()
for link in self.state.instance_links.values():
prefill_only.update(link.prefill_instances)
for link in self.state.instance_links.values():
prefill_only.difference_update(link.decode_instances)
# If the user typed a prefill-only model id (e.g.
# the vLLM-side producer of a P/D pair), the
# candidate decode side is whatever it's linked
# to. Expand the requested model id to also
# include those linked decode instances.
requested_model = command.task_params.model
linked_decode_ids: set[InstanceId] = set()
for link in self.state.instance_links.values():
if any(
self.state.instances.get(pid) is not None
and self.state.instances[
pid
].shard_assignments.model_id
== requested_model
for pid in link.prefill_instances
):
linked_decode_ids.update(link.decode_instances)
for instance in self.state.instances.values():
model_match = (
instance.shard_assignments.model_id
== requested_model
) or (instance.instance_id in linked_decode_ids)
if (
model_match
and instance.instance_id not in prefill_only
):
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = (
task_count
@@ -154,20 +230,27 @@ class Master:
],
)
decode_instance_id = available_instance_ids[0]
task_id = TaskId()
params = command.task_params.model_copy(
update={
"prefill_endpoint": _prefill_endpoint_for(
self.state, decode_instance_id
),
}
)
generated_events.append(
TaskCreated(
task_id=task_id,
task=TextGenerationTask(
task_id=task_id,
command_id=command.command_id,
instance_id=available_instance_ids[0],
instance_id=decode_instance_id,
task_status=TaskStatus.Pending,
task_params=command.task_params,
task_params=params,
),
)
)
self.command_task_mapping[command.command_id] = task_id
case ImageGeneration():
for instance in self.state.instances.values():
@@ -175,10 +258,12 @@ class Master:
instance.shard_assignments.model_id
== command.task_params.model
):
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = (
task_count
@@ -229,10 +314,12 @@ class Master:
instance.shard_assignments.model_id
== command.task_params.model
):
in_flight = {TaskStatus.Pending, TaskStatus.Running}
task_count = sum(
1
for task in self.state.tasks.values()
if task.instance_id == instance.instance_id
and task.task_status in in_flight
)
instance_task_counts[instance.instance_id] = (
task_count
@@ -357,6 +444,21 @@ class Master:
generated_events.append(
CustomModelCardDeleted(model_id=command.model_id)
)
case SetInstanceLink():
link = InstanceLink(
link_id=command.link_id,
prefill_instances=list(
dict.fromkeys(command.prefill_instances)
),
decode_instances=list(
dict.fromkeys(command.decode_instances)
),
)
generated_events.append(InstanceLinkCreated(link=link))
case DeleteInstanceLink():
generated_events.append(
InstanceLinkDeleted(link_id=command.link_id)
)
case RequestEventLog():
# We should just be able to send everything, since other buffers will ignore old messages
# rate limit to 1000 at a time
+8 -7
View File
@@ -1,4 +1,3 @@
import random
from collections.abc import Mapping
from copy import deepcopy
from typing import Sequence
@@ -44,13 +43,10 @@ from exo.shared.types.worker.instances import (
InstanceMeta,
MlxJacclInstance,
MlxRingInstance,
VllmInstance,
)
from exo.shared.types.worker.shards import Sharding
def random_ephemeral_port() -> int:
port = random.randint(49153, 65535)
return port - 1 if port <= 52415 else port
from exo.utils.ports import random_ephemeral_port
def add_instance_to_placements(
@@ -207,7 +203,7 @@ def place_instance(
)
# Single-node: force Pipeline/Ring (Tensor and Jaccl require multi-node)
if len(selected_cycle) == 1:
if len(selected_cycle) == 1 and command.instance_meta != InstanceMeta.Vllm:
command = command.model_copy(
update={
"instance_meta": InstanceMeta.MlxRing,
@@ -271,6 +267,11 @@ def place_instance(
hosts_by_node=hosts_by_node,
ephemeral_port=ephemeral_port,
)
case InstanceMeta.Vllm:
target_instances[instance_id] = VllmInstance(
instance_id=instance_id,
shard_assignments=shard_assignments,
)
return target_instances
+10 -4
View File
@@ -336,7 +336,7 @@ def _find_connection_ip(
yield connection.sink_multiaddr.ip_address
def _find_ip_prioritised(
def find_ip_prioritised(
node_id: NodeId,
other_node_id: NodeId,
cycle_digraph: Topology,
@@ -375,7 +375,13 @@ def _find_ip_prioritised(
"maybe_ethernet": 3,
"thunderbolt": 4,
}
return min(ips, key=lambda ip: priority.get(ip_to_type.get(ip, "unknown"), 2))
def _key(ip: str) -> tuple[int, int]:
link_local = 0 if ip.startswith("169.254.") else 1
type_pri = priority.get(ip_to_type.get(ip, "unknown"), 2)
return (link_local, type_pri)
return min(ips, key=_key)
def get_mlx_ring_hosts_by_node(
@@ -413,7 +419,7 @@ def get_mlx_ring_hosts_by_node(
hosts_for_node.append(Host(ip="198.51.100.1", port=0))
continue
connection_ip = _find_ip_prioritised(
connection_ip = find_ip_prioritised(
node_id, other_node_id, cycle_digraph, node_network, ring=True
)
if connection_ip is None:
@@ -445,7 +451,7 @@ def get_mlx_jaccl_coordinators(
if n == coordinator:
return "0.0.0.0"
ip = _find_ip_prioritised(
ip = find_ip_prioritised(
n, coordinator, cycle_digraph, node_network, ring=False
)
if ip is not None:
-33
View File
@@ -1,33 +0,0 @@
from typing import cast
def as_str(value: object) -> str | None:
return value if isinstance(value, str) else None
def as_list(value: object) -> list[object] | None:
if isinstance(value, list):
return cast(list[object], value)
return None
def as_dict(value: object) -> dict[str, object] | None:
if isinstance(value, dict):
return cast(dict[str, object], value)
return None
def as_int(value: object, default: int = 0) -> int:
return value if isinstance(value, int) and not isinstance(value, bool) else default
def dict_get_str(d: dict[str, object], key: str) -> str | None:
return as_str(d.get(key))
def dict_get_list(d: dict[str, object], key: str) -> list[object] | None:
return as_list(d.get(key))
def dict_get_dict(d: dict[str, object], key: str) -> dict[str, object] | None:
return as_dict(d.get(key))
-261
View File
@@ -1,261 +0,0 @@
"""Accumulators for capturing the emitted assistant shape from a streaming response.
Both accumulators are fed raw SSE chunks (bytes) as they pass through. At stream
end, they expose a canonical assistant-message shape suitable for hashing, plus
the reasoning text that should be cached against that hash.
"""
import json
import logging
from typing import cast
from exo.reasoning_proxy._helpers import (
as_dict,
as_str,
dict_get_dict,
dict_get_list,
dict_get_str,
)
logger = logging.getLogger(__name__)
class OpenAIAccumulator:
"""Captures content, tool_calls, and reasoning_content from OpenAI SSE chunks.
OpenAI can emit multiple choices per chunk; we only track choice index 0
(the common case for chat completions; n>1 is uncommon and re-hash misses
there degrade gracefully to no-op cache insert).
"""
def __init__(self) -> None:
self._content_parts: list[str] = []
self._reasoning_parts: list[str] = []
self._tool_calls_by_index: dict[int, dict[str, object]] = {}
self._buffer = ""
def feed_bytes(self, chunk: bytes) -> None:
self._buffer += chunk.decode("utf-8", errors="replace")
while "\n" in self._buffer:
line, self._buffer = self._buffer.split("\n", 1)
line = line.strip()
if not line.startswith("data:"):
continue
payload = line[len("data:") :].strip()
if payload == "[DONE]" or not payload:
continue
try:
parsed = cast(object, json.loads(payload))
except json.JSONDecodeError:
continue
event = as_dict(parsed)
if event is None:
continue
self._consume_event(event)
def _consume_event(self, event: dict[str, object]) -> None:
choices = dict_get_list(event, "choices")
if not choices:
return
for choice_raw in choices:
choice = as_dict(choice_raw)
if choice is None:
continue
index_val = choice.get("index", 0)
if (
not (isinstance(index_val, int) and not isinstance(index_val, bool))
or index_val != 0
):
continue
delta = dict_get_dict(choice, "delta")
if delta is None:
continue
content = dict_get_str(delta, "content")
if content is not None:
self._content_parts.append(content)
reasoning = dict_get_str(delta, "reasoning_content")
if reasoning is not None:
self._reasoning_parts.append(reasoning)
tool_calls = dict_get_list(delta, "tool_calls")
if tool_calls is not None:
self._merge_tool_calls(tool_calls)
def _merge_tool_calls(self, deltas: list[object]) -> None:
for raw in deltas:
d = as_dict(raw)
if d is None:
continue
index_val = d.get("index", 0)
if not (isinstance(index_val, int) and not isinstance(index_val, bool)):
continue
entry = self._tool_calls_by_index.setdefault(
index_val,
{
"id": "",
"type": "function",
"function": {"name": "", "arguments": ""},
},
)
tc_id = dict_get_str(d, "id")
if tc_id is not None:
entry["id"] = tc_id
tc_type = dict_get_str(d, "type")
if tc_type is not None:
entry["type"] = tc_type
fn = dict_get_dict(d, "function")
if fn is not None:
entry_fn = entry.get("function")
if not isinstance(entry_fn, dict):
entry_fn = {"name": "", "arguments": ""}
entry["function"] = entry_fn
entry_fn_typed = cast(dict[str, object], entry_fn)
name = dict_get_str(fn, "name")
if name is not None:
prev_name = as_str(entry_fn_typed.get("name")) or ""
entry_fn_typed["name"] = prev_name + name
args = dict_get_str(fn, "arguments")
if args is not None:
prev_args = as_str(entry_fn_typed.get("arguments")) or ""
entry_fn_typed["arguments"] = prev_args + args
@property
def content(self) -> str | None:
joined = "".join(self._content_parts)
return joined if joined else None
@property
def tool_calls(self) -> list[dict[str, object]] | None:
if not self._tool_calls_by_index:
return None
ordered = [
self._tool_calls_by_index[i] for i in sorted(self._tool_calls_by_index)
]
return ordered
@property
def reasoning(self) -> str:
return "".join(self._reasoning_parts)
class ClaudeAccumulator:
"""Captures Claude streaming content blocks.
Tracks per-index content blocks. At end, exposes the final `content_blocks`
list (excluding thinking blocks those go into `reasoning` as joined text)
in a shape suitable for hashing.
"""
def __init__(self) -> None:
self._blocks_by_index: dict[int, dict[str, object]] = {}
self._buffer = ""
self._current_event: str | None = None
def feed_bytes(self, chunk: bytes) -> None:
self._buffer += chunk.decode("utf-8", errors="replace")
while "\n" in self._buffer:
line, self._buffer = self._buffer.split("\n", 1)
line = line.rstrip("\r")
if not line:
self._current_event = None
continue
if line.startswith("event:"):
self._current_event = line[len("event:") :].strip()
continue
if line.startswith("data:"):
payload = line[len("data:") :].strip()
try:
parsed = cast(object, json.loads(payload))
except json.JSONDecodeError:
continue
event = as_dict(parsed)
if event is None:
continue
self._consume_event(event)
def _consume_event(self, event: dict[str, object]) -> None:
event_type = dict_get_str(event, "type") or self._current_event
if event_type == "content_block_start":
index_val = event.get("index", 0)
if not (isinstance(index_val, int) and not isinstance(index_val, bool)):
return
block = dict_get_dict(event, "content_block")
if block is None:
return
btype = dict_get_str(block, "type")
if btype == "text":
self._blocks_by_index[index_val] = {"type": "text", "text": ""}
elif btype == "thinking":
self._blocks_by_index[index_val] = {
"type": "thinking",
"thinking": "",
}
elif btype == "tool_use":
self._blocks_by_index[index_val] = {
"type": "tool_use",
"id": dict_get_str(block, "id") or "",
"name": dict_get_str(block, "name") or "",
"input_json": "",
}
elif event_type == "content_block_delta":
index_val = event.get("index", 0)
if not (isinstance(index_val, int) and not isinstance(index_val, bool)):
return
delta = dict_get_dict(event, "delta")
if delta is None:
return
block = self._blocks_by_index.get(index_val)
if block is None:
return
dtype = dict_get_str(delta, "type")
if dtype == "text_delta":
text = dict_get_str(delta, "text")
if text is not None:
prev = as_str(block.get("text")) or ""
block["text"] = prev + text
elif dtype == "thinking_delta":
thinking = dict_get_str(delta, "thinking")
if thinking is not None:
prev = as_str(block.get("thinking")) or ""
block["thinking"] = prev + thinking
elif dtype == "input_json_delta":
partial = dict_get_str(delta, "partial_json")
if partial is not None:
prev = as_str(block.get("input_json")) or ""
block["input_json"] = prev + partial
@property
def content_blocks(self) -> list[dict[str, object]]:
"""Public blocks (excludes thinking), with tool_use input parsed from JSON."""
public: list[dict[str, object]] = []
for index in sorted(self._blocks_by_index):
block = self._blocks_by_index[index]
if block.get("type") == "thinking":
continue
if block.get("type") == "tool_use":
input_json = as_str(block.get("input_json")) or "{}"
parsed_input_raw: object
try:
parsed_input_raw = cast(object, json.loads(input_json))
except json.JSONDecodeError:
parsed_input_raw = {}
parsed_input = as_dict(parsed_input_raw) or {}
public.append(
{
"type": "tool_use",
"id": as_str(block.get("id")) or "",
"name": as_str(block.get("name")) or "",
"input": parsed_input,
}
)
else:
public.append({k: v for k, v in block.items() if k != "input_json"})
return public
@property
def reasoning(self) -> str:
parts: list[str] = []
for index in sorted(self._blocks_by_index):
block = self._blocks_by_index[index]
if block.get("type") == "thinking":
parts.append(as_str(block.get("thinking")) or "")
return "".join(parts)
-21
View File
@@ -1,21 +0,0 @@
import threading
class ReasoningCache:
def __init__(self) -> None:
self._store: dict[str, str] = {}
self._lock = threading.Lock()
def get(self, content_hash: str) -> str | None:
with self._lock:
return self._store.get(content_hash)
def put(self, content_hash: str, reasoning: str) -> None:
if not reasoning:
return
with self._lock:
self._store[content_hash] = reasoning
def size(self) -> int:
with self._lock:
return len(self._store)
-170
View File
@@ -1,170 +0,0 @@
"""Dialect strategies for deciding which assistant history indices should receive
cached reasoning on inbound requests.
Each dialect inspects the message list and returns the set of indices where
reasoning_content (OpenAI) or a thinking block (Claude) should be reattached if
the cache has it. The dialect does not mutate messages the caller does.
Dialect selection is driven by the `reasoning_dialect` field on each model card,
surfaced through /v1/models.
"""
from typing import Protocol
from exo.reasoning_proxy._helpers import as_dict, as_list, dict_get_list, dict_get_str
from exo.shared.types.text_generation import ReasoningDialect
class Dialect(Protocol):
def select_attach_indices(
self, messages: list[dict[str, object]], has_tools: bool
) -> set[int]: ...
def _is_assistant(msg: dict[str, object]) -> bool:
return msg.get("role") == "assistant"
def _is_user(msg: dict[str, object]) -> bool:
return msg.get("role") == "user"
class NoneDialect:
def select_attach_indices(
self, messages: list[dict[str, object]], has_tools: bool
) -> set[int]:
return set()
class PostLastUserDialect:
"""MiniMax / GLM / Qwen-thinking / V4-with-tools.
Preserve reasoning on every assistant message appearing after the last
non-tool-response user message. Tool-response user messages (role=tool, or
role=user with tool_call_id set, or Claude's tool_result block) don't count
as "real" user turns they're part of the assistant's tool-calling chain.
"""
def select_attach_indices(
self, messages: list[dict[str, object]], has_tools: bool
) -> set[int]:
last_user_index = -1
for i, msg in enumerate(messages):
if _is_user(msg) and not _is_tool_response(msg):
last_user_index = i
return {
i
for i, msg in enumerate(messages)
if i > last_user_index and _is_assistant(msg)
}
class SuffixDialect:
"""Kimi K2 Thinking / K2.6.
Preserve reasoning only on the tail run of tool-call-carrying assistant
messages (the current, unresolved tool-call chain). Walk backward: include
every assistant with tool_calls until we hit an assistant without tool_calls
or a non-assistant message.
"""
def select_attach_indices(
self, messages: list[dict[str, object]], has_tools: bool
) -> set[int]:
indices: set[int] = set()
for i in range(len(messages) - 1, -1, -1):
msg = messages[i]
if not _is_assistant(msg):
if _is_tool_response(msg):
continue
break
if not _has_tool_calls(msg):
break
indices.add(i)
return indices
class ChannelDialect:
"""GPT-OSS Harmony format.
Preserve analysis-channel content on assistant turns that follow the most
recent assistant message tagged with a "final" channel marker. If no prior
final exists, the whole conversation is one unresolved chain.
"""
def select_attach_indices(
self, messages: list[dict[str, object]], has_tools: bool
) -> set[int]:
last_final_index = -1
for i, msg in enumerate(messages):
if _is_assistant(msg) and _has_final_channel(msg):
last_final_index = i
return {
i
for i, msg in enumerate(messages)
if i > last_final_index and _is_assistant(msg)
}
class ToolConditionalDialect:
"""DeepSeek V4 Flash.
If the request has tools, behave as PostLastUserDialect; otherwise passthrough.
"""
def __init__(self) -> None:
self._inner = PostLastUserDialect()
def select_attach_indices(
self, messages: list[dict[str, object]], has_tools: bool
) -> set[int]:
if not has_tools:
return set()
return self._inner.select_attach_indices(messages, has_tools)
def _is_tool_response(msg: dict[str, object]) -> bool:
if msg.get("role") == "tool":
return True
if msg.get("role") == "user" and msg.get("tool_call_id"):
return True
content = as_list(msg.get("content"))
if content is not None:
for raw in content:
block = as_dict(raw)
if block is not None and block.get("type") == "tool_result":
return True
return False
def _has_tool_calls(msg: dict[str, object]) -> bool:
tc = dict_get_list(msg, "tool_calls")
if tc:
return True
content = as_list(msg.get("content"))
if content is not None:
for raw in content:
block = as_dict(raw)
if block is not None and block.get("type") == "tool_use":
return True
return False
def _has_final_channel(msg: dict[str, object]) -> bool:
if msg.get("channel") == "final":
return True
content = dict_get_str(msg, "content")
return bool(content and content.strip())
_DIALECTS: dict[ReasoningDialect, Dialect] = {
"none": NoneDialect(),
"post_last_user": PostLastUserDialect(),
"suffix": SuffixDialect(),
"channel": ChannelDialect(),
"tool_conditional": ToolConditionalDialect(),
}
def get_dialect(name: ReasoningDialect) -> Dialect:
return _DIALECTS[name]
-75
View File
@@ -1,75 +0,0 @@
import hashlib
import json
from exo.reasoning_proxy._helpers import as_dict, dict_get_str
def _canonical_tool_calls(
tool_calls: list[dict[str, object]] | None,
) -> list[dict[str, object]]:
if not tool_calls:
return []
result: list[dict[str, object]] = []
for tc in tool_calls:
entry: dict[str, object] = {}
if "id" in tc:
entry["id"] = tc["id"]
fn = as_dict(tc.get("function"))
if fn is not None:
entry["function"] = {
"name": dict_get_str(fn, "name") or "",
"arguments": dict_get_str(fn, "arguments") or "",
}
if "type" in tc:
entry["type"] = tc["type"]
result.append(entry)
return result
def hash_openai_assistant(
content: str | list[object] | None,
tool_calls: list[dict[str, object]] | None,
) -> str:
"""Deterministic hash of an OpenAI assistant message's observable surface.
Canonicalizes None content to "" and tool_calls to a minimal id/function shape
so trivial shape differences between client render and our re-emit don't miss.
"""
shape: dict[str, object] = {
"content": content if content is not None else "",
"tool_calls": _canonical_tool_calls(tool_calls),
}
payload = json.dumps(
shape, sort_keys=True, ensure_ascii=False, separators=(",", ":")
)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
def hash_claude_assistant(content_blocks: list[dict[str, object]]) -> str:
"""Deterministic hash of a Claude assistant message's observable surface.
Skips thinking blocks (we're hashing what the *client sends back*, which typically
omits thinking) and normalizes tool_use blocks to id/name/input.
"""
normalized: list[dict[str, object]] = []
for block in content_blocks:
btype = block.get("type")
if btype == "text":
normalized.append(
{"type": "text", "text": dict_get_str(block, "text") or ""}
)
elif btype == "tool_use":
normalized.append(
{
"type": "tool_use",
"id": dict_get_str(block, "id") or "",
"name": dict_get_str(block, "name") or "",
"input": block.get("input")
if block.get("input") is not None
else {},
}
)
payload = json.dumps(
normalized, sort_keys=True, ensure_ascii=False, separators=(",", ":")
)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
-70
View File
@@ -1,70 +0,0 @@
import argparse
import logging
from contextlib import asynccontextmanager
from typing import cast
import httpx
import uvicorn
from fastapi import FastAPI
from exo.reasoning_proxy.cache import ReasoningCache
from exo.reasoning_proxy.registry import DialectRegistry
from exo.reasoning_proxy.routes import register_routes
logger = logging.getLogger("exo.reasoning_proxy")
def build_app(upstream: str) -> FastAPI:
client = httpx.AsyncClient(timeout=httpx.Timeout(None, connect=10.0))
cache = ReasoningCache()
registry = DialectRegistry(upstream=upstream, client=client)
@asynccontextmanager
async def lifespan(_: FastAPI):
await registry.refresh()
try:
yield
finally:
await client.aclose()
app = FastAPI(lifespan=lifespan, title="exo-reasoning-proxy")
register_routes(
app=app,
client=client,
upstream=upstream.rstrip("/"),
cache=cache,
registry=registry,
)
return app
def main() -> None:
parser = argparse.ArgumentParser(prog="exo-reasoning-proxy")
_ = parser.add_argument("--upstream", default="http://localhost:52415")
_ = parser.add_argument("--host", default="127.0.0.1")
_ = parser.add_argument("--port", type=int, default=52416)
_ = parser.add_argument("-v", "--verbose", action="count", default=0)
args = parser.parse_args()
verbose = cast(int, args.verbose)
upstream = cast(str, args.upstream)
host = cast(str, args.host)
port = cast(int, args.port)
level = logging.WARNING
if verbose == 1:
level = logging.INFO
elif verbose >= 2:
level = logging.DEBUG
logging.basicConfig(
level=level, format="%(asctime)s %(levelname)s %(name)s: %(message)s"
)
logger.info("Starting exo-reasoning-proxy on %s:%d%s", host, port, upstream)
app = build_app(upstream=upstream)
uvicorn.run(app, host=host, port=port, log_level=level)
if __name__ == "__main__":
main()
-75
View File
@@ -1,75 +0,0 @@
import asyncio
import logging
from typing import cast, get_args
import httpx
from exo.reasoning_proxy._helpers import as_dict, as_list, dict_get_str
from exo.shared.types.text_generation import ReasoningDialect
logger = logging.getLogger(__name__)
class DialectRegistry:
def __init__(self, upstream: str, client: httpx.AsyncClient) -> None:
self._upstream = upstream.rstrip("/")
self._client = client
self._by_model: dict[str, ReasoningDialect] = {}
self._unknown_logged: set[str] = set()
self._lock = asyncio.Lock()
self._initialized = False
async def refresh(self) -> None:
await self._fetch()
async def _fetch(self) -> None:
try:
resp = await self._client.get(f"{self._upstream}/v1/models", timeout=10.0)
resp.raise_for_status()
body = as_dict(cast(object, resp.json()))
if body is None:
return
data = as_list(body.get("data")) or []
updated: dict[str, ReasoningDialect] = {}
for entry_raw in data:
entry = as_dict(entry_raw)
if entry is None:
continue
model_id = dict_get_str(entry, "id")
dialect_raw = entry.get("reasoning_dialect", "none")
if model_id is not None:
updated[model_id] = _coerce_dialect(dialect_raw)
self._by_model = updated
self._initialized = True
logger.info(
"Loaded %d model dialects from %s", len(updated), self._upstream
)
except Exception as exc:
logger.warning(
"Failed to fetch /v1/models from %s: %s", self._upstream, exc
)
async def resolve(self, model_id: str) -> ReasoningDialect:
async with self._lock:
if not self._initialized:
await self._fetch()
if model_id in self._by_model:
return self._by_model[model_id]
await self._fetch()
if model_id in self._by_model:
return self._by_model[model_id]
if model_id not in self._unknown_logged:
logger.info(
"No dialect declared for model %s; passing through", model_id
)
self._unknown_logged.add(model_id)
return "none"
_VALID_DIALECTS: frozenset[str] = frozenset(get_args(ReasoningDialect))
def _coerce_dialect(value: object) -> ReasoningDialect:
if isinstance(value, str) and value in _VALID_DIALECTS:
return cast(ReasoningDialect, value)
return "none"
-384
View File
@@ -1,384 +0,0 @@
"""FastAPI handlers for the reasoning proxy.
Two handlers, one shape: read body resolve dialect reattach cached
reasoning to designated history indices forward tee the response stream
capture emitted reasoning cache under the emitted assistant's hash.
"""
import json
import logging
from collections.abc import AsyncIterator
from typing import cast
import httpx
from fastapi import FastAPI, Request
from fastapi.responses import JSONResponse, Response, StreamingResponse
from starlette.datastructures import Headers as StarletteHeaders
from exo.reasoning_proxy._helpers import (
as_dict,
as_list,
dict_get_dict,
dict_get_list,
dict_get_str,
)
from exo.reasoning_proxy.accumulator import ClaudeAccumulator, OpenAIAccumulator
from exo.reasoning_proxy.cache import ReasoningCache
from exo.reasoning_proxy.dialects import get_dialect
from exo.reasoning_proxy.hashing import hash_claude_assistant, hash_openai_assistant
from exo.reasoning_proxy.registry import DialectRegistry
logger = logging.getLogger(__name__)
def _parse_body(raw_body: bytes) -> dict[str, object] | None:
try:
parsed = cast(object, json.loads(raw_body))
except json.JSONDecodeError:
return None
return as_dict(parsed)
def _content_for_hash(value: object) -> str | list[object] | None:
if isinstance(value, str):
return value
as_listed = as_list(value)
if as_listed is not None:
return as_listed
return None
def _attach_openai_reasoning(
messages: list[dict[str, object]],
indices: set[int],
cache: ReasoningCache,
) -> None:
for i in indices:
msg = messages[i]
existing = dict_get_str(msg, "reasoning_content")
if existing:
continue
content = _content_for_hash(msg.get("content"))
tool_calls_raw = dict_get_list(msg, "tool_calls") or []
tool_calls: list[dict[str, object]] = [
d for d in (as_dict(t) for t in tool_calls_raw) if d is not None
]
h = hash_openai_assistant(content, tool_calls or None)
cached = cache.get(h)
if cached is not None:
msg["reasoning_content"] = cached
def _attach_claude_reasoning(
messages: list[dict[str, object]],
indices: set[int],
cache: ReasoningCache,
) -> None:
for i in indices:
msg = messages[i]
content = as_list(msg.get("content"))
if content is None:
continue
has_thinking = False
normalized: list[dict[str, object]] = []
for raw in content:
block = as_dict(raw)
if block is None:
continue
if block.get("type") == "thinking":
has_thinking = True
normalized.append(block)
if has_thinking:
continue
h = hash_claude_assistant(normalized)
cached = cache.get(h)
if cached is None:
continue
new_content: list[dict[str, object]] = [
{"type": "thinking", "thinking": cached}
]
new_content.extend(normalized)
msg["content"] = new_content
async def _stream_and_capture_openai(
upstream_resp: httpx.Response,
cache: ReasoningCache,
) -> AsyncIterator[bytes]:
accumulator = OpenAIAccumulator()
try:
async for chunk in upstream_resp.aiter_raw():
accumulator.feed_bytes(chunk)
yield chunk
finally:
await upstream_resp.aclose()
reasoning = accumulator.reasoning
if not reasoning:
return
h = hash_openai_assistant(accumulator.content, accumulator.tool_calls)
cache.put(h, reasoning)
async def _stream_and_capture_claude(
upstream_resp: httpx.Response,
cache: ReasoningCache,
) -> AsyncIterator[bytes]:
accumulator = ClaudeAccumulator()
try:
async for chunk in upstream_resp.aiter_raw():
accumulator.feed_bytes(chunk)
yield chunk
finally:
await upstream_resp.aclose()
reasoning = accumulator.reasoning
if not reasoning:
return
h = hash_claude_assistant(accumulator.content_blocks)
cache.put(h, reasoning)
def _capture_openai_nonstream(body_text: str, cache: ReasoningCache) -> None:
body = _parse_body(body_text.encode("utf-8"))
if body is None:
return
choices = dict_get_list(body, "choices")
if not choices:
return
first = as_dict(choices[0])
if first is None:
return
message = dict_get_dict(first, "message")
if message is None:
return
reasoning = dict_get_str(message, "reasoning_content")
if not reasoning:
return
content = _content_for_hash(message.get("content"))
tool_calls_raw = dict_get_list(message, "tool_calls") or []
tool_calls: list[dict[str, object]] = [
d for d in (as_dict(t) for t in tool_calls_raw) if d is not None
]
h = hash_openai_assistant(content, tool_calls or None)
cache.put(h, reasoning)
def _capture_claude_nonstream(body_text: str, cache: ReasoningCache) -> None:
body = _parse_body(body_text.encode("utf-8"))
if body is None:
return
content = as_list(body.get("content"))
if content is None:
return
reasoning_parts: list[str] = []
public_blocks: list[dict[str, object]] = []
for raw in content:
block = as_dict(raw)
if block is None:
continue
if block.get("type") == "thinking":
thinking = dict_get_str(block, "thinking")
if thinking is not None:
reasoning_parts.append(thinking)
else:
public_blocks.append(block)
reasoning = "".join(reasoning_parts)
if not reasoning:
return
h = hash_claude_assistant(public_blocks)
cache.put(h, reasoning)
def _messages_from_body(body: dict[str, object]) -> list[dict[str, object]] | None:
raw = as_list(body.get("messages"))
if raw is None:
return None
result: list[dict[str, object]] = []
for item in raw:
m = as_dict(item)
if m is None:
return None
result.append(m)
return result
def register_routes(
app: FastAPI,
client: httpx.AsyncClient,
upstream: str,
cache: ReasoningCache,
registry: DialectRegistry,
) -> None:
async def handle_chat_completions(request: Request) -> Response:
raw_body = await request.body()
body = _parse_body(raw_body)
if body is None:
return _bad_request("invalid JSON body")
model_id = dict_get_str(body, "model")
if model_id is None:
return _bad_request("missing or invalid 'model' field")
dialect_name = await registry.resolve(model_id)
dialect = get_dialect(dialect_name)
messages = _messages_from_body(body)
if messages is not None:
has_tools = bool(body.get("tools"))
indices = dialect.select_attach_indices(messages, has_tools=has_tools)
if indices:
_attach_openai_reasoning(messages, indices, cache)
body["messages"] = messages
forward_body = json.dumps(body).encode("utf-8")
forward_headers = _copy_headers(request.headers)
forward_headers["content-length"] = str(len(forward_body))
is_stream = bool(body.get("stream"))
try:
if is_stream:
req = client.build_request(
"POST",
f"{upstream}/v1/chat/completions",
content=forward_body,
headers=forward_headers,
)
upstream_resp = await client.send(req, stream=True)
return StreamingResponse(
_stream_and_capture_openai(upstream_resp, cache),
status_code=upstream_resp.status_code,
media_type=_media_type(upstream_resp.headers, "text/event-stream"),
headers=_response_headers(upstream_resp.headers),
)
upstream_resp = await client.post(
f"{upstream}/v1/chat/completions",
content=forward_body,
headers=forward_headers,
)
text = upstream_resp.text
if upstream_resp.status_code == 200:
_capture_openai_nonstream(text, cache)
return Response(
content=text,
status_code=upstream_resp.status_code,
media_type=_media_type(upstream_resp.headers, "application/json"),
headers=_response_headers(upstream_resp.headers),
)
except httpx.RequestError as exc:
logger.warning("Upstream request failed: %s", exc)
return _bad_gateway(str(exc))
async def handle_claude_messages(request: Request) -> Response:
raw_body = await request.body()
body = _parse_body(raw_body)
if body is None:
return _bad_request("invalid JSON body")
model_id = dict_get_str(body, "model")
if model_id is None:
return _bad_request("missing or invalid 'model' field")
dialect_name = await registry.resolve(model_id)
dialect = get_dialect(dialect_name)
messages = _messages_from_body(body)
if messages is not None:
has_tools = bool(body.get("tools"))
indices = dialect.select_attach_indices(messages, has_tools=has_tools)
if indices:
_attach_claude_reasoning(messages, indices, cache)
body["messages"] = messages
forward_body = json.dumps(body).encode("utf-8")
forward_headers = _copy_headers(request.headers)
forward_headers["content-length"] = str(len(forward_body))
is_stream = bool(body.get("stream"))
try:
if is_stream:
req = client.build_request(
"POST",
f"{upstream}/v1/messages",
content=forward_body,
headers=forward_headers,
)
upstream_resp = await client.send(req, stream=True)
return StreamingResponse(
_stream_and_capture_claude(upstream_resp, cache),
status_code=upstream_resp.status_code,
media_type=_media_type(upstream_resp.headers, "text/event-stream"),
headers=_response_headers(upstream_resp.headers),
)
upstream_resp = await client.post(
f"{upstream}/v1/messages",
content=forward_body,
headers=forward_headers,
)
text = upstream_resp.text
if upstream_resp.status_code == 200:
_capture_claude_nonstream(text, cache)
return Response(
content=text,
status_code=upstream_resp.status_code,
media_type=_media_type(upstream_resp.headers, "application/json"),
headers=_response_headers(upstream_resp.headers),
)
except httpx.RequestError as exc:
logger.warning("Upstream request failed: %s", exc)
return _bad_gateway(str(exc))
async def health() -> dict[str, object]:
return {"status": "ok", "cache_entries": cache.size()}
_ = app.post("/v1/chat/completions")(handle_chat_completions)
_ = app.post("/v1/messages")(handle_claude_messages)
_ = app.get("/health")(health)
_HOP_BY_HOP = {
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailers",
"transfer-encoding",
"upgrade",
"content-length",
"host",
}
def _copy_headers(headers: httpx.Headers | StarletteHeaders) -> dict[str, str]:
out: dict[str, str] = {}
for k, v in headers.items():
if k.lower() in _HOP_BY_HOP:
continue
out[k] = v
return out
def _response_headers(headers: httpx.Headers) -> dict[str, str]:
return _copy_headers(headers)
def _media_type(headers: httpx.Headers, default: str) -> str:
value = cast(object, headers.get("content-type", default))
return value if isinstance(value, str) else default
def _bad_request(msg: str) -> JSONResponse:
return JSONResponse(
status_code=400,
content={"error": {"message": msg, "type": "invalid_request_error"}},
)
def _bad_gateway(msg: str) -> JSONResponse:
return JSONResponse(
status_code=502,
content={
"error": {"message": f"upstream unreachable: {msg}", "type": "bad_gateway"}
},
)
+79 -4
View File
@@ -14,6 +14,8 @@ from exo.shared.types.events import (
InputChunkReceived,
InstanceCreated,
InstanceDeleted,
InstanceLinkCreated,
InstanceLinkDeleted,
NodeDownloadProgress,
NodeGatheredInfo,
NodeTimedOut,
@@ -29,6 +31,7 @@ from exo.shared.types.events import (
TracesCollected,
TracesMerged,
)
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.profiling import (
NodeIdentity,
NodeNetworkInfo,
@@ -41,7 +44,12 @@ from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.topology import Connection, RDMAConnection
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
from exo.shared.types.worker.runners import RunnerId, RunnerShutdown, RunnerStatus
from exo.shared.types.worker.runners import (
RunnerId,
RunnerReady,
RunnerShutdown,
RunnerStatus,
)
from exo.utils.info_gatherer.info_gatherer import (
MacmonMetrics,
MacThunderboltConnections,
@@ -51,9 +59,11 @@ from exo.utils.info_gatherer.info_gatherer import (
NodeConfig,
NodeDiskUsage,
NodeNetworkInterfaces,
NvmlMetrics,
RdmaCtlStatus,
StaticNodeInformation,
ThunderboltBridgeInfo,
VllmCapability,
)
@@ -95,6 +105,10 @@ def event_apply(event: Event, state: State) -> State:
return apply_topology_edge_created(event, state)
case TopologyEdgeDeleted():
return apply_topology_edge_deleted(event, state)
case InstanceLinkCreated():
return apply_instance_link_created(event, state)
case InstanceLinkDeleted():
return apply_instance_link_deleted(event, state)
def apply(state: State, event: IndexedEvent) -> State:
@@ -194,7 +208,38 @@ def apply_instance_deleted(event: InstanceDeleted, state: State) -> State:
new_instances: Mapping[InstanceId, Instance] = {
iid: inst for iid, inst in state.instances.items() if iid != event.instance_id
}
return state.model_copy(update={"instances": new_instances})
new_links: dict[InstanceLinkId, InstanceLink] = {}
for link_id, link in state.instance_links.items():
prefill = [i for i in link.prefill_instances if i != event.instance_id]
decode = [i for i in link.decode_instances if i != event.instance_id]
if not prefill or not decode:
continue
if prefill == list(link.prefill_instances) and decode == list(
link.decode_instances
):
new_links[link_id] = link
else:
new_links[link_id] = link.model_copy(
update={"prefill_instances": prefill, "decode_instances": decode}
)
return state.model_copy(
update={"instances": new_instances, "instance_links": new_links}
)
def apply_instance_link_created(event: InstanceLinkCreated, state: State) -> State:
new_links: Mapping[InstanceLinkId, InstanceLink] = {
**state.instance_links,
event.link.link_id: event.link,
}
return state.model_copy(update={"instance_links": new_links})
def apply_instance_link_deleted(event: InstanceLinkDeleted, state: State) -> State:
new_links: Mapping[InstanceLinkId, InstanceLink] = {
lid: link for lid, link in state.instance_links.items() if lid != event.link_id
}
return state.model_copy(update={"instance_links": new_links})
def apply_runner_status_updated(event: RunnerStatusUpdated, state: State) -> State:
@@ -202,12 +247,28 @@ def apply_runner_status_updated(event: RunnerStatusUpdated, state: State) -> Sta
new_runners: Mapping[RunnerId, RunnerStatus] = {
rid: rs for rid, rs in state.runners.items() if rid != event.runner_id
}
return state.model_copy(update={"runners": new_runners})
new_ports: Mapping[RunnerId, int] = {
rid: p
for rid, p in state.prefill_server_ports.items()
if rid != event.runner_id
}
return state.model_copy(
update={"runners": new_runners, "prefill_server_ports": new_ports}
)
new_runners = {
**state.runners,
event.runner_id: event.runner_status,
}
return state.model_copy(update={"runners": new_runners})
update: dict[str, object] = {"runners": new_runners}
if (
isinstance(event.runner_status, RunnerReady)
and event.runner_status.prefill_server_port is not None
):
update["prefill_server_ports"] = {
**state.prefill_server_ports,
event.runner_id: event.runner_status.prefill_server_port,
}
return state.model_copy(update=update)
def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
@@ -245,6 +306,9 @@ def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
node_rdma_ctl = {
key: value for key, value in state.node_rdma_ctl.items() if key != event.node_id
}
node_vllm = {
key: value for key, value in state.node_vllm.items() if key != event.node_id
}
# Only recompute cycles if the leaving node had TB bridge enabled
leaving_node_status = state.node_thunderbolt_bridge.get(event.node_id)
leaving_node_had_tb_enabled = (
@@ -267,6 +331,7 @@ def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
"node_thunderbolt": node_thunderbolt,
"node_thunderbolt_bridge": node_thunderbolt_bridge,
"node_rdma_ctl": node_rdma_ctl,
"node_vllm": node_vllm,
"thunderbolt_bridge_cycles": thunderbolt_bridge_cycles,
}
)
@@ -293,6 +358,11 @@ def apply_node_gathered_info(event: NodeGatheredInfo, state: State) -> State:
event.node_id: info.system_profile,
}
update["node_memory"] = {**state.node_memory, event.node_id: info.memory}
case NvmlMetrics():
update["node_system"] = {
**state.node_system,
event.node_id: info.system_profile,
}
case MemoryUsage():
update["node_memory"] = {**state.node_memory, event.node_id: info}
case NodeDiskUsage():
@@ -373,6 +443,11 @@ def apply_node_gathered_info(event: NodeGatheredInfo, state: State) -> State:
**state.node_rdma_ctl,
event.node_id: NodeRdmaCtlStatus(enabled=info.enabled),
}
case VllmCapability():
update["node_vllm"] = {
**state.node_vllm,
event.node_id: info.available,
}
return state.model_copy(update=update)
+2
View File
@@ -96,6 +96,8 @@ EXO_OFFLINE = os.getenv("EXO_OFFLINE", "false").lower() == "true"
EXO_TRACING_ENABLED = os.getenv("EXO_TRACING_ENABLED", "false").lower() == "true"
ENABLE_DISAGGREGATION = os.getenv("ENABLE_DISAGGREGATION", "false").lower() == "true"
EXO_MAX_CONCURRENT_REQUESTS = int(os.getenv("EXO_MAX_CONCURRENT_REQUESTS", "8"))
EXO_MAX_INSTANCE_RETRIES = 5
+24 -15
View File
@@ -150,6 +150,7 @@ class ModelCard(FrozenModel):
context_length: int = 0
uses_cfg: bool = False
trust_remote_code: bool = True
requires_vllm: bool = False
is_custom: bool = False
vision: VisionCardConfig | None = None
sampling_defaults: SamplingDefaults = Field(default_factory=SamplingDefaults)
@@ -349,7 +350,11 @@ async def fetch_config_data(model_id: ModelId) -> ConfigData:
async def fetch_safetensors_size(model_id: ModelId) -> Memory:
"""Gets model size from safetensors index or falls back to HF API."""
"""Gets model size from safetensors index or falls back to HF API.
Single-shard repos don't have a `model.safetensors.index.json`; fall back
to the HF API for those.
"""
from exo.download.download_utils import (
download_file_with_retry,
resolve_model_dir,
@@ -357,21 +362,25 @@ async def fetch_safetensors_size(model_id: ModelId) -> Memory:
from exo.shared.types.worker.downloads import ModelSafetensorsIndex
target_dir = await resolve_model_dir(model_id)
index_path = await download_file_with_retry(
model_id,
"main",
"model.safetensors.index.json",
target_dir,
lambda curr_bytes, total_bytes, is_renamed: logger.debug(
f"Downloading model.safetensors.index.json for {model_id}: {curr_bytes}/{total_bytes} ({is_renamed=})"
),
)
async with aiofiles.open(index_path, "r") as f:
index_data = ModelSafetensorsIndex.model_validate_json(await f.read())
try:
index_path = await download_file_with_retry(
model_id,
"main",
"model.safetensors.index.json",
target_dir,
lambda curr_bytes, total_bytes, is_renamed: logger.debug(
f"Downloading model.safetensors.index.json for {model_id}: {curr_bytes}/{total_bytes} ({is_renamed=})"
),
)
except FileNotFoundError:
index_path = None
metadata = index_data.metadata
if metadata is not None and metadata.total_size is not None:
return Memory.from_bytes(metadata.total_size)
if index_path is not None:
async with aiofiles.open(index_path, "r") as f:
index_data = ModelSafetensorsIndex.model_validate_json(await f.read())
metadata = index_data.metadata
if metadata is not None and metadata.total_size is not None:
return Memory.from_bytes(metadata.total_size)
info = model_info(model_id)
if info.safetensors is None:
@@ -0,0 +1,72 @@
from exo.shared.apply import (
apply_instance_deleted,
apply_instance_link_created,
apply_instance_link_deleted,
)
from exo.shared.types.events import (
InstanceDeleted,
InstanceLinkCreated,
InstanceLinkDeleted,
)
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.state import State
from exo.shared.types.worker.instances import InstanceId
def _link(
prefill: list[InstanceId],
decode: list[InstanceId],
link_id: InstanceLinkId | None = None,
) -> InstanceLink:
return InstanceLink(
link_id=link_id or InstanceLinkId(),
prefill_instances=prefill,
decode_instances=decode,
)
def test_create_link() -> None:
state = State()
link = _link([InstanceId("a")], [InstanceId("b")])
new_state = apply_instance_link_created(InstanceLinkCreated(link=link), state)
assert new_state.instance_links == {link.link_id: link}
def test_update_replaces_existing_link() -> None:
a, b, c = InstanceId("a"), InstanceId("b"), InstanceId("c")
link = _link([a], [b])
state = State(instance_links={link.link_id: link})
updated = link.model_copy(update={"decode_instances": [b, c]})
new_state = apply_instance_link_created(InstanceLinkCreated(link=updated), state)
assert set(new_state.instance_links[link.link_id].decode_instances) == {b, c}
def test_delete_link() -> None:
link = _link([InstanceId("a")], [InstanceId("b")])
state = State(instance_links={link.link_id: link})
new_state = apply_instance_link_deleted(
InstanceLinkDeleted(link_id=link.link_id), state
)
assert new_state.instance_links == {}
def test_instance_deleted_strips_from_links() -> None:
a, b, c = InstanceId("a"), InstanceId("b"), InstanceId("c")
link = _link([a, c], [b])
state = State(instance_links={link.link_id: link})
new_state = apply_instance_deleted(InstanceDeleted(instance_id=a), state)
remaining = new_state.instance_links[link.link_id]
assert remaining.prefill_instances == [c]
assert remaining.decode_instances == [b]
def test_instance_deleted_drops_link_when_role_empties() -> None:
a, b = InstanceId("a"), InstanceId("b")
link = _link([a], [b])
state = State(instance_links={link.link_id: link})
new_state = apply_instance_deleted(InstanceDeleted(instance_id=a), state)
assert link.link_id not in new_state.instance_links
+3 -3
View File
@@ -85,6 +85,6 @@ class PrefillProgressChunk(BaseChunk):
total_tokens: int
GenerationChunk = (
TokenChunk | ImageChunk | ToolCallChunk | ErrorChunk | PrefillProgressChunk
)
StatusChunk = PrefillProgressChunk
GenerationChunk = TokenChunk | ImageChunk | ToolCallChunk | ErrorChunk
Chunk = StatusChunk | GenerationChunk
+13
View File
@@ -7,6 +7,7 @@ from exo.api.types import (
from exo.shared.models.model_cards import ModelCard, ModelId
from exo.shared.types.chunks import InputImageChunk
from exo.shared.types.common import CommandId, NodeId, SystemId
from exo.shared.types.instance_link import InstanceLinkId
from exo.shared.types.text_generation import TextGenerationTaskParams
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
from exo.shared.types.worker.shards import Sharding, ShardMetadata
@@ -89,6 +90,16 @@ class DeleteCustomModelCard(BaseCommand):
model_id: ModelId
class SetInstanceLink(BaseCommand):
link_id: InstanceLinkId
prefill_instances: list[InstanceId]
decode_instances: list[InstanceId]
class DeleteInstanceLink(BaseCommand):
link_id: InstanceLinkId
DownloadCommand = StartDownload | DeleteDownload | CancelDownload
@@ -106,6 +117,8 @@ Command = (
| SendInputChunk
| AddCustomModelCard
| DeleteCustomModelCard
| SetInstanceLink
| DeleteInstanceLink
)
+13 -2
View File
@@ -5,8 +5,9 @@ from pydantic import Field
from exo.shared.models.model_cards import ModelCard
from exo.shared.topology import Connection
from exo.shared.types.chunks import GenerationChunk, InputImageChunk
from exo.shared.types.chunks import Chunk, InputImageChunk
from exo.shared.types.common import CommandId, Id, ModelId, NodeId, SessionId, SystemId
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.tasks import Task, TaskId, TaskStatus
from exo.shared.types.worker.downloads import DownloadProgress
from exo.shared.types.worker.instances import Instance, InstanceId
@@ -91,7 +92,7 @@ class NodeDownloadProgress(BaseEvent):
class ChunkGenerated(BaseEvent):
command_id: CommandId
chunk: GenerationChunk
chunk: Chunk
class InputChunkReceived(BaseEvent):
@@ -137,6 +138,14 @@ class TracesMerged(BaseEvent):
traces: list[TraceEventData]
class InstanceLinkCreated(BaseEvent):
link: InstanceLink
class InstanceLinkDeleted(BaseEvent):
link_id: InstanceLinkId
Event = (
TestEvent
| TaskCreated
@@ -158,6 +167,8 @@ Event = (
| TracesMerged
| CustomModelCardAdded
| CustomModelCardDeleted
| InstanceLinkCreated
| InstanceLinkDeleted
)
+13
View File
@@ -0,0 +1,13 @@
from exo.shared.types.common import Id
from exo.shared.types.worker.instances import InstanceId
from exo.utils.pydantic_ext import FrozenModel
class InstanceLinkId(Id):
pass
class InstanceLink(FrozenModel):
link_id: InstanceLinkId
prefill_instances: list[InstanceId]
decode_instances: list[InstanceId]
+5
View File
@@ -7,6 +7,7 @@ from pydantic.alias_generators import to_camel
from exo.shared.topology import Topology, TopologySnapshot
from exo.shared.types.common import NodeId
from exo.shared.types.instance_link import InstanceLink, InstanceLinkId
from exo.shared.types.profiling import (
DiskUsage,
MemoryUsage,
@@ -57,10 +58,14 @@ class State(FrozenModel):
node_thunderbolt: Mapping[NodeId, NodeThunderboltInfo] = {}
node_thunderbolt_bridge: Mapping[NodeId, ThunderboltBridgeStatus] = {}
node_rdma_ctl: Mapping[NodeId, NodeRdmaCtlStatus] = {}
node_vllm: Mapping[NodeId, bool] = {}
# Detected cycles where all nodes have Thunderbolt bridge enabled (>2 nodes)
thunderbolt_bridge_cycles: Sequence[Sequence[NodeId]] = []
instance_links: Mapping[InstanceLinkId, InstanceLink] = {}
prefill_server_ports: Mapping[RunnerId, int] = {}
@field_serializer("topology", mode="plain")
def _encode_topology(self, value: Topology) -> TopologySnapshot:
return value.to_snapshot()
+3
View File
@@ -101,3 +101,6 @@ Task = (
| ImageEdits
| Shutdown
)
TextTask = TextGeneration
ImageTask = ImageGeneration | ImageEdits
GenerationTask = TextTask | ImageTask
+16
View File
@@ -13,6 +13,20 @@ from exo.shared.types.common import ModelId, TruncatingString
MessageRole = Literal["user", "assistant", "system", "developer", "tool"]
ReasoningEffort = Literal["none", "minimal", "low", "medium", "high", "xhigh"]
# How a model wants prior-turn reasoning content handled. Drives both the
# server-side encoder (drop vs keep) and the integration configs we emit
# (e.g. opencode's per-model `interleaved` flag).
# - "none": model has no reasoning channel.
# - "post_last_user": reasoning is only meaningful for the latest assistant
# turn; older turns can drop it (drop_thinking=True).
# - "suffix": reasoning is embedded in the assistant content as a
# suffix/prefix; round-tripping content already covers
# it (no separate `reasoning_content` round-trip).
# - "channel": reasoning lives on a dedicated channel (Harmony, etc.)
# and must be sent back verbatim every turn.
# - "tool_conditional": always round-trip when the conversation has tools;
# the model relies on prior reasoning to chain tool
# calls (DeepSeek V3.2 / V4).
ReasoningDialect = Literal[
"none", "post_last_user", "suffix", "channel", "tool_conditional"
]
@@ -118,6 +132,8 @@ class TextGenerationTaskParams(BaseModel, frozen=True):
images: list[Base64Image] = Field(default_factory=list)
image_hashes: dict[int, Base64ImageHash] = Field(default_factory=dict)
prefill_endpoint: str | None = None
def with_card_sampling_defaults(self) -> "TextGenerationTaskParams":
from exo.shared.models.model_cards import get_card
File renamed without changes.
+6 -1
View File
@@ -15,6 +15,7 @@ class InstanceId(Id):
class InstanceMeta(str, Enum):
MlxRing = "MlxRing"
MlxJaccl = "MlxJaccl"
Vllm = "Vllm"
class BaseInstance(TaggedModel):
@@ -35,8 +36,12 @@ class MlxJacclInstance(BaseInstance):
jaccl_coordinators: dict[NodeId, str]
class VllmInstance(BaseInstance):
pass
# TODO: Single node instance
Instance = MlxRingInstance | MlxJacclInstance
Instance = MlxRingInstance | MlxJacclInstance | VllmInstance
class BoundInstance(FrozenModel):
@@ -16,10 +16,6 @@ class BaseRunnerResponse(TaggedModel):
pass
class TokenizedResponse(BaseRunnerResponse):
prompt_tokens: int
class GenerationResponse(BaseRunnerResponse):
text: str
token: int
@@ -75,6 +71,10 @@ class ModelLoadingResponse(BaseRunnerResponse):
total: int
class CancelledResponse(BaseRunnerResponse):
pass
class PrefillProgressResponse(BaseRunnerResponse):
processed_tokens: int
total_tokens: int
+1 -1
View File
@@ -47,7 +47,7 @@ class RunnerWarmingUp(BaseRunnerStatus):
class RunnerReady(BaseRunnerStatus):
pass
prefill_server_port: int | None = None
class RunnerRunning(BaseRunnerStatus):
+6 -6
View File
@@ -25,12 +25,12 @@ def print_startup_banner(port: int) -> None:
banner = f"""
Distributed AI Inference Cluster
@@ -31,6 +31,7 @@ from exo.utils.pydantic_ext import TaggedModel
from exo.utils.task_group import TaskGroup
from .macmon import MacmonMetrics
from .nvml import NvmlMetrics, gather_nvidia_metrics, has_nvml
from .system_info import (
get_friendly_name,
get_model_and_chip,
@@ -353,6 +354,24 @@ async def _gather_iface_map() -> dict[str, str] | None:
return ports
class VllmCapability(TaggedModel):
available: bool
version: str | None = None
@classmethod
async def gather(cls) -> Self:
try:
import importlib
vllm = importlib.import_module("vllm")
return cls(
available=True,
version=cast(str | None, getattr(vllm, "__version__", None)),
)
except ImportError:
return cls(available=False)
GatheredInfo = (
MacmonMetrics
| MemoryUsage
@@ -361,6 +380,8 @@ GatheredInfo = (
| MacThunderboltConnections
| RdmaCtlStatus
| ThunderboltBridgeInfo
| NvmlMetrics
| VllmCapability
| NodeConfig
| MiscData
| StaticNodeInformation
@@ -419,6 +440,8 @@ class InfoGatherer:
tg.start_soon(self._monitor_rdma_ctl_status, 10)
if not IS_DARWIN:
tg.start_soon(self._monitor_memory_usage, 1)
if has_nvml():
tg.start_soon(self._monitor_nvml_metrics, 1)
tg.start_soon(self._watch_system_info, 10)
tg.start_soon(self._monitor_misc, 60)
tg.start_soon(self._monitor_static_info, 60)
@@ -427,6 +450,10 @@ class InfoGatherer:
nc = await NodeConfig.gather()
if nc is not None:
await self.info_sender.send(nc)
try:
await self.info_sender.send(await VllmCapability.gather())
except Exception as e:
logger.warning(f"Error gathering vLLM capability: {e}")
def shutdown(self):
self._tg.cancel_tasks()
@@ -475,6 +502,16 @@ class InfoGatherer:
logger.opt(exception=e).warning("Error gathering Thunderbolt data")
await anyio.sleep(system_profiler_interval)
async def _monitor_nvml_metrics(self, nvml_poll_rate: float):
while True:
try:
metrics = gather_nvidia_metrics()
if metrics is not None:
await self.info_sender.send(metrics)
except Exception as e:
logger.opt(exception=e).warning("Error gathering NVML metrics")
await anyio.sleep(nvml_poll_rate)
async def _monitor_memory_usage(self, memory_poll_rate: float):
if self._psutil_enabled:
return
+70
View File
@@ -0,0 +1,70 @@
from exo.shared.types.profiling import SystemPerformanceProfile
from exo.utils.pydantic_ext import TaggedModel
try:
import pynvml as nvml
except ImportError:
nvml = None
_CPU_POWER_IDLE = 20.0
_CPU_POWER_MAX = 100.0
_GPU_POWER_MAX = 120.0
class NvmlMetrics(TaggedModel):
system_profile: SystemPerformanceProfile
def has_nvml() -> bool:
if nvml is None:
return False
try:
nvml.nvmlInit()
count = nvml.nvmlDeviceGetCount()
nvml.nvmlShutdown()
return count > 0
except Exception:
return False
def gather_nvidia_metrics() -> NvmlMetrics | None:
if nvml is None:
return None
is_init = False
try:
nvml.nvmlInit()
is_init = True
count = nvml.nvmlDeviceGetCount()
if count == 0:
return None
total_gpu_util = 0.0
total_temp = 0.0
total_gpu_power = 0.0
for i in range(count):
handle = nvml.nvmlDeviceGetHandleByIndex(i)
util = nvml.nvmlDeviceGetUtilizationRates(handle)
total_gpu_util += float(util.gpu)
total_temp += float(
nvml.nvmlDeviceGetTemperatureV(handle, nvml.NVML_TEMPERATURE_GPU)
)
total_gpu_power += float(nvml.nvmlDeviceGetPowerUsage(handle)) / 1000.0
gpu_load_fraction = min(total_gpu_power / _GPU_POWER_MAX, 1.0)
estimated_cpu_power = (
_CPU_POWER_IDLE + (_CPU_POWER_MAX - _CPU_POWER_IDLE) * gpu_load_fraction
)
return NvmlMetrics(
system_profile=SystemPerformanceProfile(
gpu_usage=total_gpu_util / count / 100.0,
temp=total_temp / count,
sys_power=total_gpu_power + estimated_cpu_power,
),
)
except Exception:
return None
finally:
if is_init:
nvml.nvmlShutdown()
+81 -2
View File
@@ -1,6 +1,7 @@
import platform
import socket
import sys
from pathlib import Path
from subprocess import CalledProcessError
import psutil
@@ -117,12 +118,90 @@ async def get_network_interfaces() -> list[NetworkInterfaceInfo]:
return interfaces_info
def _read_dmi_field(name: str) -> str | None:
try:
path = Path(f"/sys/class/dmi/id/{name}")
if path.exists():
return path.read_text().strip()
except (OSError, PermissionError):
pass
return None
async def _get_linux_model_and_chip() -> tuple[str, str]:
model = "Linux"
chip = "Unknown Chip"
product_name = _read_dmi_field("product_name")
sys_vendor = _read_dmi_field("sys_vendor")
# DGX Spark: DMI product_name may be "DGX_Spark" or "gx10" variant
product_lower = (product_name or "").lower()
if product_name and ("dgx" in product_lower or "gx10" in product_lower):
model = "DGX Spark"
try:
process = await run_process(
["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"]
)
gpu_name = process.stdout.decode().strip().split("\n")[0]
chip = gpu_name if gpu_name and gpu_name != "[N/A]" else "NVIDIA GB10"
except (CalledProcessError, FileNotFoundError):
chip = "NVIDIA GB10"
return (model, chip)
# Other NVIDIA systems (sys_vendor contains "NVIDIA")
if sys_vendor and "NVIDIA" in sys_vendor:
model = product_name.replace("_", " ") if product_name else "NVIDIA System"
try:
process = await run_process(
["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"]
)
gpu_name = process.stdout.decode().strip().split("\n")[0]
if gpu_name and gpu_name != "[N/A]":
chip = gpu_name
except (CalledProcessError, FileNotFoundError):
pass
return (model, chip)
# Generic Linux — detect laptop vs desktop via chassis_type
# SMBIOS chassis types: 8,9,10,14,31,32 = portable/laptop
chassis_type = _read_dmi_field("chassis_type")
laptop_chassis_types = {"8", "9", "10", "14", "31", "32"}
if chassis_type in laptop_chassis_types:
model = "Linux Laptop"
elif chassis_type is not None:
model = "Linux Desktop"
# Also check for battery as a fallback laptop indicator
if model == "Linux" and Path("/sys/class/power_supply/BAT0").exists():
model = "Linux Laptop"
# Use /proc/cpuinfo for chip
cpuinfo_path = Path("/proc/cpuinfo")
if cpuinfo_path.exists():
try:
for line in cpuinfo_path.read_text().splitlines():
if line.startswith("model name"):
chip = line.split(":", 1)[1].strip()
break
except OSError:
pass
return (model, chip)
async def get_model_and_chip() -> tuple[str, str]:
"""Get Mac system information using system_profiler."""
"""Get system model and chip information.
On macOS, uses ``system_profiler``. On Linux, reads DMI data from
sysfs and CPU info from ``/proc/cpuinfo``.
"""
model = "Unknown Model"
chip = "Unknown Chip"
# TODO: better non mac support
if sys.platform == "linux":
return await _get_linux_model_and_chip()
if sys.platform != "darwin":
return (model, chip)
+6
View File
@@ -0,0 +1,6 @@
import random
def random_ephemeral_port() -> int:
port = random.randint(49153, 65535)
return port - 1 if port <= 52415 else port
Whitespace-only changes.
+204
View File
@@ -0,0 +1,204 @@
from dataclasses import dataclass
from typing import BinaryIO, Literal
import msgspec
DType = Literal["bfloat16", "float16", "float32"]
class ProtocolError(Exception):
pass
class Header(msgspec.Struct):
request_id: str = ""
model_id: str = ""
num_layers: int = 0
dtype: DType = "bfloat16"
start_pos: int = 0
class TensorBlob(msgspec.Struct):
dtype: DType
shape: tuple[int, ...]
data: bytes
class _KVChunkHeader(msgspec.Struct, tag="kv_chunk"):
"""Wire-side KV chunk metadata. Raw `keys` then `values` bytes follow on
the stream, lengths given by `keys_len` / `values_len`. Splitting them out
of the msgpack frame lets the producer pass tensor buffers via the buffer
protocol straight into the socket (one host-side memcpy total).
"""
layer_idx: int
num_tokens: int
n_heads: int
head_dim: int
dtype: DType
keys_len: int
values_len: int
@dataclass(frozen=True)
class KVChunk:
"""In-memory KV chunk reconstructed by `read_message` from
`_KVChunkHeader` + the raw bytes that follow on the wire.
"""
layer_idx: int
num_tokens: int
n_heads: int
head_dim: int
dtype: DType
keys: bytes
values: bytes
@property
def shape(self) -> tuple[int, int, int]:
return (self.num_tokens, self.n_heads, self.head_dim)
class ArraysState(msgspec.Struct, tag="arrays_state"):
layer_idx: int
arrays: list[TensorBlob] = []
class Done(msgspec.Struct, tag="done"):
total_tokens: int
class ErrorMessage(msgspec.Struct, tag="error"):
code: int
message: str
_WireMessage = _KVChunkHeader | ArraysState | Done | ErrorMessage
Message = KVChunk | ArraysState | Done | ErrorMessage
_msg_encoder = msgspec.msgpack.Encoder()
_msg_decoder: msgspec.msgpack.Decoder[_WireMessage] = msgspec.msgpack.Decoder(
_WireMessage
)
_header_encoder = msgspec.msgpack.Encoder()
_header_decoder: msgspec.msgpack.Decoder[Header] = msgspec.msgpack.Decoder(Header)
def _read_exactly(stream: BinaryIO, n: int) -> bytes:
buf = bytearray()
while len(buf) < n:
chunk = stream.read(n - len(buf))
if not chunk:
if len(buf) == 0:
return b""
raise ConnectionError(f"Connection closed after {len(buf)}/{n} bytes")
buf.extend(chunk)
return bytes(buf)
def write_frame(stream: BinaryIO, payload: bytes) -> None:
stream.write(len(payload).to_bytes(4, "big"))
stream.write(payload)
stream.flush()
def read_frame(stream: BinaryIO) -> bytes:
raw = _read_exactly(stream, 4)
if not raw:
return b""
length = int.from_bytes(raw, "big")
return _read_exactly(stream, length)
def write_header(stream: BinaryIO, header: Header) -> None:
write_frame(stream, _header_encoder.encode(header))
def read_header(stream: BinaryIO) -> Header:
payload = read_frame(stream)
if not payload:
raise ConnectionError("No header received")
try:
return _header_decoder.decode(payload)
except msgspec.DecodeError as exc:
raise ProtocolError(f"Bad header: {exc}") from exc
def write_message(stream: BinaryIO, msg: _WireMessage) -> None:
write_frame(stream, _msg_encoder.encode(msg))
def read_message(stream: BinaryIO) -> Message | None:
payload = read_frame(stream)
if not payload:
return None
try:
msg = _msg_decoder.decode(payload)
except msgspec.DecodeError as exc:
raise ProtocolError(f"Bad message: {exc}") from exc
if isinstance(msg, _KVChunkHeader):
keys = _read_exactly(stream, msg.keys_len)
values = _read_exactly(stream, msg.values_len)
return KVChunk(
layer_idx=msg.layer_idx,
num_tokens=msg.num_tokens,
n_heads=msg.n_heads,
head_dim=msg.head_dim,
dtype=msg.dtype,
keys=keys,
values=values,
)
return msg
def write_kv_chunk(
stream: BinaryIO,
*,
layer_idx: int,
num_tokens: int,
n_heads: int,
head_dim: int,
dtype: DType,
keys: "bytes | memoryview",
values: "bytes | memoryview",
) -> None:
"""Stream KV chunk metadata + raw key/value bytes to the wire.
`keys` / `values` may be bytes-like (bytes, bytearray, memoryview) the
raw payload is written directly to the buffered stream after the
msgpack-framed header, avoiding a memcpy through the msgpack encoder.
"""
keys_len = len(keys)
values_len = len(values)
header_payload = _msg_encoder.encode(
_KVChunkHeader(
layer_idx=layer_idx,
num_tokens=num_tokens,
n_heads=n_heads,
head_dim=head_dim,
dtype=dtype,
keys_len=keys_len,
values_len=values_len,
)
)
stream.write(len(header_payload).to_bytes(4, "big"))
stream.write(header_payload)
stream.write(keys)
stream.write(values)
# No per-chunk flush: the K/V payload is far larger than the
# BufferedWriter's internal buffer so it bypasses to the socket directly.
# The trailing `Done` frame's `write_frame` flushes once at the end.
def write_arrays_state(
stream: BinaryIO, layer_idx: int, arrays: list[TensorBlob]
) -> None:
write_message(stream, ArraysState(layer_idx=layer_idx, arrays=arrays))
def write_done(stream: BinaryIO, total_tokens: int) -> None:
write_message(stream, Done(total_tokens=total_tokens))
def write_error(stream: BinaryIO, code: int, message: str) -> None:
write_message(stream, ErrorMessage(code=code, message=message))
+109
View File
@@ -0,0 +1,109 @@
import socket
import socketserver
import threading
from collections.abc import Callable
from typing import BinaryIO, cast
import msgspec
from loguru import logger
from exo.worker.disaggregated.protocol import (
Header,
read_frame,
write_error,
write_frame,
write_header,
)
class PrefillRequest(msgspec.Struct):
request_id: str = ""
model_id: str = ""
token_ids: list[int] = msgspec.field(default_factory=list)
start_pos: int = 0
use_prefix_cache: bool = True
_request_encoder = msgspec.msgpack.Encoder()
_request_decoder: msgspec.msgpack.Decoder[PrefillRequest] = msgspec.msgpack.Decoder(
PrefillRequest
)
def write_request(stream: BinaryIO, job: PrefillRequest) -> None:
write_frame(stream, _request_encoder.encode(job))
def read_request(stream: BinaryIO) -> PrefillRequest:
payload = read_frame(stream)
if not payload:
raise ConnectionError("No request received")
return _request_decoder.decode(payload)
ResolveHandler = Callable[[PrefillRequest, BinaryIO], bool]
def _send_error(wfile: BinaryIO, code: int, message: str) -> None:
try:
write_header(wfile, Header(num_layers=0, dtype="float32"))
write_error(wfile, code=code, message=message)
except Exception:
pass
class _PrefillHandler(socketserver.StreamRequestHandler):
def setup(self) -> None:
super().setup()
sock = cast(socket.socket, self.request)
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
# 64MB send buffer: K/V chunks are ~33MB each; a small SNDBUF
# back-pressures the writer thread between chunks and serializes
# network with compute.
sock.setsockopt(socket.SOL_SOCKET, socket.SO_SNDBUF, 64 * 1024 * 1024)
def handle(self) -> None:
server = cast(PrefillServer, self.server)
wfile: BinaryIO = cast(BinaryIO, cast(object, self.wfile))
rfile: BinaryIO = cast(BinaryIO, cast(object, self.rfile))
try:
job = read_request(rfile)
except ConnectionError:
return
except (msgspec.DecodeError, ValueError) as exc:
_send_error(wfile, 400, f"Bad request: {exc}")
return
try:
picked_up = server.resolve(job, wfile)
except Exception as e:
logger.opt(exception=e).warning(
f"Prefill resolve error for request_id={job.request_id}"
)
_send_error(wfile, 500, str(e))
return
if not picked_up:
_send_error(
wfile, 503, f"Prefill not picked up for request_id={job.request_id!r}"
)
class PrefillServer(socketserver.ThreadingTCPServer):
allow_reuse_address = True
daemon_threads = True
resolve: ResolveHandler
def __init__(self, resolve: ResolveHandler, host: str, port: int) -> None:
super().__init__((host, port), _PrefillHandler)
self.resolve = resolve
self._thread = threading.Thread(
target=self.serve_forever, name="prefill-server"
)
self._thread.start()
logger.info(f"Prefill server listening on {host}:{port}")
def stop(self) -> None:
self.shutdown()
self.server_close()
if self._thread is not None:
self._thread.join(timeout=5)
self._thread = None
+60
View File
@@ -0,0 +1,60 @@
from abc import ABC, abstractmethod
from collections.abc import Generator, Iterable
from typing import BinaryIO
from exo.shared.types.chunks import Chunk
from exo.shared.types.tasks import CANCEL_ALL_TASKS, GenerationTask, TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import (
CancelledResponse,
FinishedResponse,
ModelLoadingResponse,
)
from exo.worker.disaggregated.server import PrefillRequest
class Engine(ABC):
_cancelled_tasks: set[TaskId]
def should_cancel(self, task_id: TaskId) -> bool:
return (
task_id in self._cancelled_tasks
or CANCEL_ALL_TASKS in self._cancelled_tasks
)
@abstractmethod
def warmup(self) -> None: ...
@abstractmethod
def submit(
self,
task: GenerationTask,
) -> None: ...
@abstractmethod
def step(
self,
) -> Iterable[tuple[TaskId, Chunk | CancelledResponse | FinishedResponse]]: ...
@abstractmethod
def close(self) -> None: ...
@abstractmethod
def serve_prefill(self, request: PrefillRequest, wfile: BinaryIO) -> None: ...
class Builder(ABC):
@abstractmethod
def connect(self, bound_instance: BoundInstance) -> None: ...
@abstractmethod
def load(
self,
bound_instance: BoundInstance,
) -> Generator[ModelLoadingResponse]: ...
@abstractmethod
def build(self) -> Engine: ...
@abstractmethod
def close(self) -> None: ...
+6 -2
View File
@@ -1,12 +1,16 @@
from exo.worker.engines.image.builder import (
ImageEngine,
MfluxBuilder,
)
from exo.worker.engines.image.distributed_model import (
DistributedImageModel,
initialize_image_model,
)
from exo.worker.engines.image.generate import generate_image, warmup_image_generator
__all__ = [
"MfluxBuilder",
"ImageEngine",
"DistributedImageModel",
"generate_image",
"initialize_image_model",
"warmup_image_generator",
]
+219
View File
@@ -0,0 +1,219 @@
import contextlib
from collections import deque
from collections.abc import Generator, Iterable
from dataclasses import dataclass, field
from typing import BinaryIO
import mlx.core as mx
from loguru import logger
from exo.api.types import ImageEditsTaskParams, ImageGenerationTaskParams
from exo.shared.constants import EXO_TRACING_ENABLED
from exo.shared.tracing import clear_trace_buffer, get_trace_buffer
from exo.shared.types.chunks import Chunk, ErrorChunk
from exo.shared.types.events import (
Event,
TraceEventData,
TracesCollected,
)
from exo.shared.types.tasks import (
GenerationTask,
ImageEdits,
ImageGeneration,
ImageTask,
TaskId,
)
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import (
CancelledResponse,
FinishedResponse,
ModelLoadingResponse,
)
from exo.shared.types.worker.shards import (
CfgShardMetadata,
PipelineShardMetadata,
ShardMetadata,
)
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.disaggregated.server import PrefillRequest
from exo.worker.engines.base import Builder, Engine
from exo.worker.engines.image.distributed_model import (
DistributedImageModel,
)
from exo.worker.engines.image.generate import (
generate_image,
warmup_image_generator,
)
from exo.worker.engines.mlx.utils_mlx import (
initialize_mlx,
)
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 _send_traces_if_enabled(
event_sender: MpSender[Event],
task_id: TaskId,
rank: int,
) -> None:
if not EXO_TRACING_ENABLED:
return
traces = get_trace_buffer()
if traces:
trace_data = [
TraceEventData(
name=t.name,
start_us=t.start_us,
duration_us=t.duration_us,
rank=t.rank,
category=t.category,
)
for t in traces
]
event_sender.send(
TracesCollected(
task_id=task_id,
rank=rank,
traces=trace_data,
)
)
clear_trace_buffer()
@dataclass
class MfluxBuilder(Builder):
event_sender: MpSender[Event]
cancel_receiver: MpReceiver[TaskId]
shard_metadata: ShardMetadata | None = None
image_model: DistributedImageModel | None = None
group: mx.distributed.Group | None = None
def connect(self, bound_instance: BoundInstance) -> None:
self.group = initialize_mlx(bound_instance)
def load(self, bound_instance: BoundInstance) -> Generator[ModelLoadingResponse]:
self.shard_metadata = bound_instance.bound_shard
self.image_model = DistributedImageModel.from_shard_metadata(
bound_instance.bound_shard, self.group
)
return
# very important!
yield
def close(self) -> None:
with contextlib.suppress(NameError, AttributeError):
del self.image_model, self.group
def build(
self,
) -> Engine:
assert self.image_model
assert self.shard_metadata
return ImageEngine(
self.image_model,
self.shard_metadata,
self.event_sender,
self.cancel_receiver,
)
@dataclass
class ImageEngine(Engine):
image_model: DistributedImageModel
shard_metadata: ShardMetadata
event_sender: MpSender[Event]
cancel_receiver: MpReceiver[TaskId]
current_gen: (
Generator[tuple[TaskId, Chunk | FinishedResponse | CancelledResponse]] | None
) = field(init=False, default=None)
queue: deque[ImageTask] = field(init=False, default_factory=deque)
def warmup(self) -> None:
image = warmup_image_generator(model=self.image_model)
if image is not None:
logger.info(f"warmed up by generating {image.size} image")
else:
logger.info("warmup completed (non-primary node)")
def submit(
self,
task: GenerationTask,
) -> None:
assert isinstance(task, (ImageGeneration, ImageEdits))
self.queue.append(task)
def step(
self,
) -> Iterable[tuple[TaskId, Chunk | CancelledResponse | FinishedResponse]]:
resp = None
if self.current_gen is not None:
resp = next(self.current_gen, None)
if resp is None and len(self.queue) > 0:
task = self.queue.popleft()
self.current_gen = self._run_image_task(task.task_id, task.task_params)
resp = next(self.current_gen, None)
return (resp,) if resp is not None else ()
def close(self) -> None:
with contextlib.suppress(NameError, AttributeError):
del self.image_model
def serve_prefill(self, request: PrefillRequest, wfile: BinaryIO) -> None:
raise NotImplementedError() from None
def _run_image_task(
self,
task_id: TaskId,
task_params: ImageGenerationTaskParams | ImageEditsTaskParams,
) -> Generator[tuple[TaskId, Chunk | FinishedResponse | CancelledResponse]]:
assert self.image_model
logger.info(f"received image task: {str(task_params)[:500]}")
def cancel_checker() -> bool:
for cancel_id in self.cancel_receiver.collect():
self._cancelled_tasks.add(cancel_id)
return self.should_cancel(task_id)
try:
# todo: yield CancelledResponse properly
for response in generate_image(
model=self.image_model,
task=task_params,
cancel_checker=cancel_checker,
):
if _is_primary_output_node(self.shard_metadata):
yield (task_id, response)
except Exception as e:
if _is_primary_output_node(self.shard_metadata):
yield (
task_id,
ErrorChunk(
model=self.shard_metadata.model_card.model_id,
finish_reason="error",
error_message=str(e),
),
)
raise
finally:
_send_traces_if_enabled(
self.event_sender, task_id, self.shard_metadata.device_rank
)
yield (task_id, FinishedResponse())
return
@@ -1,6 +1,6 @@
from collections.abc import Callable, Generator
from pathlib import Path
from typing import Any, Literal, Optional
from typing import Any, Literal
import mlx.core as mx
from mflux.models.common.config.config import Config
@@ -9,8 +9,11 @@ from PIL import Image
from exo.api.types import AdvancedImageParams
from exo.download.download_utils import build_model_path
from exo.shared.types.common import ModelId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata
from exo.shared.types.worker.shards import (
CfgShardMetadata,
PipelineShardMetadata,
ShardMetadata,
)
from exo.worker.engines.image.config import ImageModelConfig
from exo.worker.engines.image.models import (
create_adapter_for_model,
@@ -18,7 +21,7 @@ from exo.worker.engines.image.models import (
)
from exo.worker.engines.image.models.base import ModelAdapter
from exo.worker.engines.image.pipeline import DiffusionRunner
from exo.worker.engines.mlx.utils_mlx import mlx_distributed_init, mx_barrier
from exo.worker.engines.mlx.utils_mlx import mx_barrier
from exo.worker.runner.bootstrap import logger
@@ -33,7 +36,7 @@ class DistributedImageModel:
model_id: ModelId,
local_path: Path,
shard_metadata: PipelineShardMetadata | CfgShardMetadata,
group: Optional[mx.distributed.Group] = None,
group: mx.distributed.Group | None,
quantize: int | None = None,
):
config = get_config_for_model(model_id)
@@ -76,32 +79,21 @@ class DistributedImageModel:
self._runner = runner
@classmethod
def from_bound_instance(
cls, bound_instance: BoundInstance
def from_shard_metadata(
cls, shard: ShardMetadata, group: mx.distributed.Group | None
) -> "DistributedImageModel":
model_id = bound_instance.bound_shard.model_card.model_id
model_id = shard.model_card.model_id
model_path = build_model_path(model_id)
shard_metadata = bound_instance.bound_shard
if not isinstance(shard_metadata, (PipelineShardMetadata, CfgShardMetadata)):
if not isinstance(shard, (PipelineShardMetadata, CfgShardMetadata)):
raise ValueError(
"Expected PipelineShardMetadata or CfgShardMetadata for image generation"
)
is_distributed = (
len(bound_instance.instance.shard_assignments.node_to_runner) > 1
)
if is_distributed:
logger.info("Starting distributed init for image model")
group = mlx_distributed_init(bound_instance)
else:
group = None
return cls(
model_id=model_id,
local_path=model_path,
shard_metadata=shard_metadata,
shard_metadata=shard,
group=group,
)
@@ -176,7 +168,3 @@ class DistributedImageModel:
else:
logger.info("generated image")
yield result
def initialize_image_model(bound_instance: BoundInstance) -> DistributedImageModel:
return DistributedImageModel.from_bound_instance(bound_instance)
+2 -2
View File
@@ -918,7 +918,7 @@ class DeepseekV4ShardingStrategy(TensorParallelShardingStrategy):
# Head-parallel attention with interleaved-per-group sharding.
_shard_v4_attention_heads(layer.attn, self.N, self.group.rank())
self.sharded_to_all_linear_in_place(layer.attn.wo_a)
layer.attn.wo_b = _AllSumLinear(layer.attn.wo_b, self.group) # type: ignore[assignment]
layer.attn.wo_b = _AllSumLinear(layer.attn.wo_b, self.group) # type: ignore
ffn = layer.ffn
if getattr(ffn, "shared_experts", None) is not None:
@@ -930,7 +930,7 @@ class DeepseekV4ShardingStrategy(TensorParallelShardingStrategy):
self.all_to_sharded_linear_in_place(ffn.switch_mlp.up_proj)
wrapped = ShardedMoEV4(ffn)
wrapped.sharding_group = self.group
layer.ffn = wrapped # type: ignore[assignment]
layer.ffn = wrapped # type: ignore
mx.eval(layer)
mx.clear_cache()
+113
View File
@@ -0,0 +1,113 @@
import contextlib
import os
from collections.abc import Generator
from dataclasses import dataclass
import mlx.core as mx
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.types.common import ModelId
from exo.shared.types.events import Event
from exo.shared.types.tasks import TaskId
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runner_response import ModelLoadingResponse
from exo.utils.channels import MpReceiver, MpSender
from exo.worker.engines.base import Builder, Engine
from exo.worker.runner.bootstrap import logger
from exo.worker.runner.llm_inference.batch_generator import (
BatchGenerator,
SequentialGenerator,
)
from exo.worker.runner.llm_inference.tool_parsers import make_mlx_parser
from .cache import KVPrefixCache
from .types import Model
from .utils_mlx import (
initialize_mlx,
load_mlx_items,
)
from .vision import VisionProcessor
@dataclass
class MlxBuilder(Builder):
model_id: ModelId
event_sender: MpSender[Event]
cancel_receiver: MpReceiver[TaskId]
model: Model | None = None
tokenizer: TokenizerWrapper | None = None
group: mx.distributed.Group | None = None
vision_processor: VisionProcessor | None = None
def connect(self, bound_instance: BoundInstance) -> None:
self.group = initialize_mlx(bound_instance)
def load(self, bound_instance: BoundInstance) -> Generator[ModelLoadingResponse]:
(
self.model,
self.tokenizer,
self.vision_processor,
) = yield from load_mlx_items(bound_instance, self.group)
def close(self) -> None:
with contextlib.suppress(NameError, AttributeError):
del self.model
with contextlib.suppress(NameError, AttributeError):
del self.tokenizer
with contextlib.suppress(NameError, AttributeError):
del self.group
def build(
self,
) -> Engine:
assert self.model
assert self.tokenizer
vision_processor = self.vision_processor
tool_parser = None
logger.info(
f"model has_tool_calling={self.tokenizer.has_tool_calling} using tokens {self.tokenizer.tool_call_start}, {self.tokenizer.tool_call_end}"
)
if (
self.tokenizer.tool_call_start
and self.tokenizer.tool_call_end
and self.tokenizer.tool_parser # type: ignore
):
tool_parser = make_mlx_parser(
self.tokenizer.tool_call_start,
self.tokenizer.tool_call_end,
self.tokenizer.tool_parser, # type: ignore
)
kv_prefix_cache = KVPrefixCache(self.group)
device_rank = 0 if self.group is None else self.group.rank()
if os.environ.get("EXO_NO_BATCH"):
logger.info("using SequentialGenerator (batching disabled)")
return SequentialGenerator(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
tool_parser=tool_parser,
kv_prefix_cache=kv_prefix_cache,
model_id=self.model_id,
device_rank=device_rank,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
vision_processor=vision_processor,
)
else:
logger.info("using BatchGenerator")
return BatchGenerator(
model=self.model,
tokenizer=self.tokenizer,
group=self.group,
tool_parser=tool_parser,
kv_prefix_cache=kv_prefix_cache,
model_id=self.model_id,
device_rank=device_rank,
cancel_receiver=self.cancel_receiver,
event_sender=self.event_sender,
vision_processor=vision_processor,
)
+102 -21
View File
@@ -13,11 +13,17 @@ from mlx_lm.models.cache import (
QuantizedKVCache,
RotatingKVCache,
)
from mlx_lm.models.deepseek_v4 import (
DeepseekV4Cache,
)
from mlx_lm.models.deepseek_v4 import (
_CompressorBranch as CompressorBranch, # type: ignore
)
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.shared.types.memory import Memory
from exo.shared.types.mlx import KVCacheType, Model
from exo.worker.engines.mlx.constants import CACHE_GROUP_SIZE, KV_CACHE_BITS
from exo.worker.engines.mlx.types import KVCacheType, Model
from exo.worker.runner.bootstrap import logger
if TYPE_CHECKING:
@@ -47,7 +53,9 @@ class CacheSnapshot:
def __init__(
self,
states: list[RotatingKVCache | ArraysCache | CacheList | None],
states: list[
RotatingKVCache | ArraysCache | CacheList | DeepseekV4Cache | None
],
token_count: int,
):
self.states = states
@@ -112,21 +120,71 @@ def _copy_cache_list(cl: CacheList) -> CacheList:
return CacheList(*copied)
def restore_snapshot_entry(
entry: ArraysCache | RotatingKVCache | CacheList | None,
) -> ArraysCache | RotatingKVCache | CacheList | None:
if entry is None:
def _detached_copy_or_none(a: mx.array | None) -> mx.array | None:
if a is None:
return None
if isinstance(entry, RotatingKVCache):
snap = copy_rotating_kv_cache(entry)
return snap if snap is not None else deepcopy(entry)
if isinstance(entry, ArraysCache):
return _copy_arrays_cache(entry)
return _copy_cache_list(entry)
out = _detached_copy(a)
mx.eval(out)
return out
def _copy_compressor_branch(b: CompressorBranch) -> CompressorBranch:
out = CompressorBranch.__new__(CompressorBranch)
out.buffer_kv = _detached_copy_or_none(b.buffer_kv)
out.buffer_gate = _detached_copy_or_none(b.buffer_gate)
out.prev_kv = _detached_copy_or_none(b.prev_kv)
out.prev_gate = _detached_copy_or_none(b.prev_gate)
out.pool = _detached_copy_or_none(b.pool)
out.buffer_lengths = deepcopy(b.buffer_lengths)
out.pool_lengths = deepcopy(b.pool_lengths)
out.buffer_count = deepcopy(b.buffer_count)
out._new_pool_lengths = deepcopy(b._new_pool_lengths)
return out
def _copy_v4_cache(c: DeepseekV4Cache) -> DeepseekV4Cache:
snap = DeepseekV4Cache.__new__(DeepseekV4Cache)
local: RotatingKVCache = c.local
local_snap = copy_rotating_kv_cache(local)
if local_snap is None:
local_snap = RotatingKVCache.__new__(RotatingKVCache)
local_snap.keys = None
local_snap.values = None
local_snap.offset = local.offset
local_snap._idx = 0
local_snap.keep = local.keep
local_snap.max_size = local.max_size
snap.local = local_snap
snap._branches = {
key: _copy_compressor_branch(branch) for key, branch in c._branches.items()
}
snap._pending_lengths = deepcopy(c._pending_lengths)
return snap
def copy_snapshot_entry(
entry: ArraysCache | RotatingKVCache | CacheList | DeepseekV4Cache | None,
) -> ArraysCache | RotatingKVCache | CacheList | DeepseekV4Cache | None:
match entry:
case None:
return None
case RotatingKVCache():
snap = copy_rotating_kv_cache(entry)
return snap if snap is not None else deepcopy(entry)
case ArraysCache():
return _copy_arrays_cache(entry)
case CacheList():
return _copy_cache_list(entry)
case DeepseekV4Cache():
return _copy_v4_cache(entry)
def snapshot_ssm_states(cache: KVCacheType) -> CacheSnapshot:
states: list[ArraysCache | RotatingKVCache | CacheList | None] = []
states: list[
RotatingKVCache | ArraysCache | CacheList | DeepseekV4Cache | None
] = []
for c in cache:
if isinstance(c, ArraysCache):
states.append(_copy_arrays_cache(c))
@@ -134,6 +192,8 @@ def snapshot_ssm_states(cache: KVCacheType) -> CacheSnapshot:
states.append(copy_rotating_kv_cache(c))
elif isinstance(c, CacheList) and not bool(c.is_trimmable()): # type: ignore[reportUnknownMemberType]
states.append(_copy_cache_list(c))
elif isinstance(c, DeepseekV4Cache):
states.append(_copy_v4_cache(c))
else:
states.append(None)
token_count = cache_length(cache)
@@ -153,14 +213,20 @@ def _find_nearest_snapshot(
return best
def is_non_trimmable_cache_entry(c: object) -> bool:
"""A cache entry is non-trimmable if `trim(n)` can't roll back its full
state meaning the prefill +2 rollback must snapshot+restore it instead.
"""
if isinstance(c, (ArraysCache, RotatingKVCache)):
return True
if isinstance(c, CacheList):
return not bool(c.is_trimmable()) # type: ignore[reportUnknownMemberType]
return isinstance(c, DeepseekV4Cache)
def has_non_kv_caches(cache: KVCacheType) -> bool:
"""Check if a cache contains any ArraysCache (SSM) entries."""
for c in cache:
if isinstance(c, CacheList):
return any(isinstance(_c, (ArraysCache, RotatingKVCache)) for _c in c) # type: ignore[reportUnknownVariableType]
elif isinstance(c, (ArraysCache, RotatingKVCache)):
return True
return False
return any(is_non_trimmable_cache_entry(c) for c in cache)
class KVPrefixCache:
@@ -316,6 +382,10 @@ class KVPrefixCache:
trim_cache(prompt_cache, tokens_to_trim, restore_snap)
# Reset cache offset to match trimmed length
for c in prompt_cache:
if isinstance(c, (ArraysCache, RotatingKVCache)):
continue
if isinstance(c, DeepseekV4Cache):
continue
if hasattr(c, "offset"):
c.offset = restore_pos
@@ -411,16 +481,22 @@ def trim_cache(
)
if non_trimmable:
if snapshot is not None and snapshot.states[i] is not None:
restored = restore_snapshot_entry(snapshot.states[i])
restored = copy_snapshot_entry(snapshot.states[i])
if restored is not None:
cache[i] = restored # type: ignore
elif isinstance(c, (ArraysCache, RotatingKVCache)):
c.state = [None] * len(c.state)
if isinstance(c, RotatingKVCache):
c.offset = 0
c._idx = 0
else:
# CacheList without a snapshot — zero each inner cache's state
for inner in c: # type: ignore[reportUnknownVariableType]
if isinstance(inner, (ArraysCache, RotatingKVCache)):
inner.state = [None] * len(inner.state)
if isinstance(inner, RotatingKVCache):
inner.offset = 0
inner._idx = 0
else:
c.trim(num_tokens)
@@ -438,7 +514,12 @@ def encode_prompt(tokenizer: TokenizerWrapper, prompt: str) -> mx.array:
def _entry_length(
c: KVCache | RotatingKVCache | QuantizedKVCache | ArraysCache | CacheList,
c: KVCache
| RotatingKVCache
| QuantizedKVCache
| ArraysCache
| CacheList
| DeepseekV4Cache,
) -> int:
# Use .offset attribute which KVCache types have (len() not implemented in older QuantizedKVCache).
if hasattr(c, "offset"):
Whitespace-only changes.
@@ -0,0 +1,272 @@
from typing import BinaryIO
import mlx.core as mx
import numpy as np
from mlx_lm.models.cache import (
ArraysCache,
CacheList,
KVCache,
QuantizedKVCache,
RotatingKVCache,
)
from mlx_lm.models.deepseek_v4 import DeepseekV4Cache
from exo.worker.disaggregated.protocol import (
DType,
Header,
KVChunk,
TensorBlob,
write_arrays_state,
write_done,
write_header,
write_kv_chunk,
)
from exo.worker.engines.mlx.types import KVCacheType
_STR_TO_MX: dict[DType, mx.Dtype] = {
"bfloat16": mx.bfloat16,
"float16": mx.float16,
"float32": mx.float32,
}
_MX_TO_STR: dict[mx.Dtype, DType] = {v: k for k, v in _STR_TO_MX.items()}
def mx_dtype_to_str(dtype: mx.Dtype) -> DType:
if dtype not in _MX_TO_STR:
raise ValueError(f"Unsupported mlx dtype on wire: {dtype}")
return _MX_TO_STR[dtype]
def wire_dtype_from_cache(caches: KVCacheType) -> DType:
for c in caches:
keys: mx.array | None = getattr(c, "keys", None)
if keys is None:
continue
if keys.dtype in _MX_TO_STR:
return _MX_TO_STR[keys.dtype]
break
return "bfloat16"
def str_to_mx_dtype(dtype: DType) -> mx.Dtype:
if dtype not in _STR_TO_MX:
raise ValueError(f"Unsupported wire dtype: {dtype!r}")
return _STR_TO_MX[dtype]
def array_to_bytes(t: mx.array) -> bytes:
# bf16 has no native numpy dtype; bitcast through uint16.
if t.dtype == mx.bfloat16:
return np.asarray(t.view(mx.uint16)).tobytes()
if t.dtype in (mx.float16, mx.float32):
return np.asarray(t).tobytes()
raise ValueError(f"Unsupported mlx dtype for wire: {t.dtype}")
def bytes_to_array(data: bytes, shape: tuple[int, ...], dtype: DType) -> mx.array:
match dtype:
case "bfloat16":
arr = np.frombuffer(data, dtype=np.uint16).reshape(shape).copy()
return mx.array(arr).view(mx.bfloat16)
case "float16":
arr = np.frombuffer(data, dtype=np.float16).reshape(shape).copy()
return mx.array(arr)
case "float32":
arr = np.frombuffer(data, dtype=np.float32).reshape(shape).copy()
return mx.array(arr)
def bhsd_to_nhd(t: mx.array) -> mx.array:
if t.ndim != 4 or int(t.shape[0]) != 1:
raise ValueError(f"Expected BHSD with B=1, got shape={tuple(t.shape)}")
return mx.transpose(t[0], (1, 0, 2))
def nhd_to_bhsd(t: mx.array) -> mx.array:
if t.ndim != 3:
raise ValueError(f"Expected NHD (3D), got shape={tuple(t.shape)}")
return mx.expand_dims(mx.transpose(t, (1, 0, 2)), 0)
def _rotating_to_temporal(buf: mx.array, idx: int, offset: int, keep: int) -> mx.array:
seq = int(buf.shape[2])
if idx == seq:
return buf
if idx < offset:
return mx.concatenate(
[buf[..., :keep, :], buf[..., idx:, :], buf[..., keep:idx, :]],
axis=2,
)
return buf[..., :idx, :]
def send_mlx_kv_cache(
stream: BinaryIO,
caches: KVCacheType,
*,
dtype: DType,
start_pos: int = 0,
max_tokens: int | None = None,
) -> int:
tokens_sent = 0
for layer_idx, c in enumerate(caches):
match c:
case QuantizedKVCache() | CacheList() | DeepseekV4Cache():
raise NotImplementedError
case KVCache():
keys = c.keys
values = c.values
if keys is None or values is None:
continue
offset = int(c.offset)
if max_tokens is not None:
offset = min(offset, max_tokens)
if offset <= start_pos:
continue
with mx.stream(mx.Device(mx.cpu)):
k = mx.array(keys[:, :, start_pos:offset, :])
v = mx.array(values[:, :, start_pos:offset, :])
k_nhd = bhsd_to_nhd(k)
v_nhd = bhsd_to_nhd(v)
mx.eval(k_nhd, v_nhd)
num_tokens = int(k_nhd.shape[0])
n_heads = int(k_nhd.shape[1])
head_dim = int(k_nhd.shape[2])
write_kv_chunk(
stream,
layer_idx=layer_idx,
num_tokens=num_tokens,
n_heads=n_heads,
head_dim=head_dim,
dtype=dtype,
keys=array_to_bytes(k_nhd),
values=array_to_bytes(v_nhd),
)
tokens_sent = max(tokens_sent, num_tokens)
case RotatingKVCache():
keys = c.keys
values = c.values
if keys is None or values is None:
continue
offset = int(c.offset)
if offset <= 0:
continue
idx = int(c._idx)
keep = int(c.keep)
with mx.stream(mx.Device(mx.cpu)):
k_temporal = _rotating_to_temporal(keys, idx, offset, keep)
v_temporal = _rotating_to_temporal(values, idx, offset, keep)
k = mx.array(k_temporal)
v = mx.array(v_temporal)
k_nhd = bhsd_to_nhd(k)
v_nhd = bhsd_to_nhd(v)
mx.eval(k_nhd, v_nhd)
num_tokens = int(k_nhd.shape[0])
n_heads = int(k_nhd.shape[1])
head_dim = int(k_nhd.shape[2])
write_kv_chunk(
stream,
layer_idx=layer_idx,
num_tokens=num_tokens,
n_heads=n_heads,
head_dim=head_dim,
dtype=dtype,
keys=array_to_bytes(k_nhd),
values=array_to_bytes(v_nhd),
)
tokens_sent = max(tokens_sent, offset)
case ArraysCache():
blobs: list[TensorBlob] = []
for a in c.state:
if a is None:
continue
with mx.stream(mx.Device(mx.cpu)):
a_cpu = mx.array(a)
mx.eval(a_cpu)
blobs.append(
TensorBlob(
dtype=mx_dtype_to_str(a_cpu.dtype),
shape=tuple(int(d) for d in a_cpu.shape),
data=array_to_bytes(a_cpu),
)
)
if blobs:
write_arrays_state(stream, layer_idx, blobs)
return tokens_sent
def chunk_to_mlx_nhd(chunk: KVChunk) -> tuple[mx.array, mx.array]:
shape = chunk.shape
return (
bytes_to_array(chunk.keys, shape, chunk.dtype),
bytes_to_array(chunk.values, shape, chunk.dtype),
)
def blob_to_mlx(blob: TensorBlob) -> mx.array:
return bytes_to_array(blob.data, blob.shape, blob.dtype)
def inject_kv_chunk(
cache: KVCache,
keys_nhd: mx.array,
values_nhd: mx.array,
offset: int,
*,
start_pos: int = 0,
existing_k: mx.array | None = None,
existing_v: mx.array | None = None,
) -> None:
k_bhsd = nhd_to_bhsd(keys_nhd)
v_bhsd = nhd_to_bhsd(values_nhd)
if start_pos > 0 and existing_k is not None and existing_v is not None:
cache.keys = mx.concatenate([existing_k[:, :, :start_pos, :], k_bhsd], axis=2)
cache.values = mx.concatenate([existing_v[:, :, :start_pos, :], v_bhsd], axis=2)
else:
cache.keys = k_bhsd
cache.values = v_bhsd
cache.offset = offset
def inject_rotating_kv_chunk(
cache: RotatingKVCache,
keys_nhd: mx.array,
values_nhd: mx.array,
offset: int,
) -> None:
k_bhsd = nhd_to_bhsd(keys_nhd)
v_bhsd = nhd_to_bhsd(values_nhd)
cache.keys = k_bhsd
cache.values = v_bhsd
cache.offset = offset
cache._idx = int(k_bhsd.shape[2])
def inject_arrays_cache(cache: ArraysCache, blobs: list[TensorBlob]) -> None:
cache.state = [blob_to_mlx(b) for b in blobs]
def write_cache_to_wire(
wfile: BinaryIO,
cache: KVCacheType,
*,
request_id: str = "",
model_id: str = "",
start_pos: int = 0,
) -> int:
dtype = wire_dtype_from_cache(cache)
write_header(
wfile,
Header(
request_id=request_id,
model_id=model_id,
num_layers=len(cache),
dtype=dtype,
start_pos=start_pos,
),
)
tokens_sent = send_mlx_kv_cache(wfile, cache, dtype=dtype, start_pos=start_pos)
write_done(wfile, tokens_sent)
wfile.flush()
return tokens_sent
@@ -0,0 +1,155 @@
import socket
from collections import defaultdict
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import BinaryIO, cast
import mlx.core as mx
from loguru import logger
from mlx_lm.models.cache import ArraysCache, KVCache, RotatingKVCache
from exo.worker.disaggregated.protocol import (
ArraysState,
Done,
Header,
KVChunk,
TensorBlob,
read_header,
read_message,
)
from exo.worker.disaggregated.server import PrefillRequest, write_request
from exo.worker.engines.mlx.disaggregated.adapter import (
chunk_to_mlx_nhd,
inject_arrays_cache,
inject_kv_chunk,
inject_rotating_kv_chunk,
)
_SOCKET_TIMEOUT_SECS = 60
_RECV_BUFFER_BYTES = 4 * 1024 * 1024
@dataclass
class PrefillResult:
header: Header
kv_chunks: dict[int, list[KVChunk]] = field(
default_factory=dict[int, list[KVChunk]]
)
arrays: dict[int, list[TensorBlob]] = field(
default_factory=dict[int, list[TensorBlob]]
)
total_tokens: int = 0
def _parse_endpoint(endpoint: str) -> tuple[str, int]:
if ":" in endpoint:
host, port_str = endpoint.rsplit(":", 1)
return host, int(port_str)
raise ValueError(f"Invalid endpoint {endpoint}")
def remote_prefill_fetch(
endpoint: str,
request: PrefillRequest,
on_header: Callable[[Header], None] | None = None,
on_kv_chunk: Callable[[KVChunk, int], None] | None = None,
timeout_secs: float = _SOCKET_TIMEOUT_SECS,
) -> PrefillResult:
host, port = _parse_endpoint(endpoint)
logger.info(
f"Connecting to prefill server at {host}:{port} "
f"({len(request.token_ids)} tokens, start_pos={request.start_pos})"
)
sock = socket.create_connection((host, port), timeout=timeout_secs)
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
sock.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, _RECV_BUFFER_BYTES)
try:
wfile = sock.makefile("wb", buffering=256 * 1024)
wstream: BinaryIO = cast(BinaryIO, cast(object, wfile))
write_request(wstream, request)
raw_stream = sock.makefile("rb", buffering=256 * 1024)
stream: BinaryIO = cast(BinaryIO, cast(object, raw_stream))
header = read_header(stream)
if on_header is not None:
on_header(header)
result = PrefillResult(header=header)
kv_by_layer: dict[int, list[KVChunk]] = defaultdict(list)
chunks_received = 0
done_seen = False
while True:
msg = read_message(stream)
if msg is None:
break
if isinstance(msg, KVChunk):
kv_by_layer[msg.layer_idx].append(msg)
chunks_received += 1
if on_kv_chunk is not None:
on_kv_chunk(msg, chunks_received)
elif isinstance(msg, ArraysState):
result.arrays[msg.layer_idx] = msg.arrays
elif isinstance(msg, Done):
result.total_tokens = msg.total_tokens
done_seen = True
break
else:
raise RuntimeError(f"Prefill server error [{msg.code}]: {msg.message}")
if not done_seen:
raise ConnectionError(
"Prefill server closed before Done frame "
f"(received {chunks_received} kv chunks, {len(result.arrays)} arrays)"
)
result.kv_chunks = dict(kv_by_layer)
return result
finally:
sock.close()
def ingest_into_mlx_cache(
result: PrefillResult,
caches: list[KVCache | RotatingKVCache | ArraysCache],
*,
start_pos: int = 0,
) -> int:
max_received = max(
(sum(c.num_tokens for c in chunks) for chunks in result.kv_chunks.values()),
default=0,
)
final_offset = start_pos + max_received
for i, cache in enumerate(caches):
if i in result.kv_chunks:
chunks = result.kv_chunks[i]
if len(chunks) == 1:
k_nhd, v_nhd = chunk_to_mlx_nhd(chunks[0])
else:
decoded = [chunk_to_mlx_nhd(c) for c in chunks]
k_nhd = mx.concatenate([k for k, _ in decoded], axis=0)
v_nhd = mx.concatenate([v for _, v in decoded], axis=0)
if isinstance(cache, RotatingKVCache):
inject_rotating_kv_chunk(cache, k_nhd, v_nhd, final_offset)
elif isinstance(cache, KVCache):
if start_pos > 0:
inject_kv_chunk(
cache,
k_nhd,
v_nhd,
final_offset,
start_pos=start_pos,
existing_k=cache.keys,
existing_v=cache.values,
)
else:
inject_kv_chunk(cache, k_nhd, v_nhd, final_offset)
if i in result.arrays and isinstance(cache, ArraysCache):
inject_arrays_cache(cache, result.arrays[i])
return final_offset
@@ -0,0 +1,86 @@
import time
import mlx.core as mx
from mlx_lm.sample_utils import make_sampler
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.worker.disaggregated.server import PrefillRequest
from exo.worker.engines.mlx.cache import (
KVPrefixCache,
cache_length,
make_kv_cache,
snapshot_ssm_states,
)
from exo.worker.engines.mlx.generator.generate import prefill as mlx_prefill
from exo.worker.engines.mlx.types import KVCacheType, Model
from exo.worker.engines.mlx.utils_mlx import fix_unmatched_think_end_tokens
from exo.worker.runner.bootstrap import logger
def run_prefill_for_request(
*,
model: Model,
tokenizer: TokenizerWrapper,
group: mx.distributed.Group | None,
kv_prefix_cache: KVPrefixCache | None,
request: PrefillRequest,
) -> KVCacheType:
prompt_tokens = mx.array(request.token_ids)
prompt_tokens = fix_unmatched_think_end_tokens(prompt_tokens, tokenizer)
n_tokens = int(prompt_tokens.shape[0])
t0 = time.perf_counter()
matched_index: int | None = None
prefix_hit_length = 0
if kv_prefix_cache is not None:
cache, remaining, matched_index, _ = kv_prefix_cache.get_kv_cache(
model, prompt_tokens
)
prefix_hit_length = n_tokens - int(remaining.shape[0])
else:
cache = make_kv_cache(model)
remaining = prompt_tokens
target_offset = max(0, n_tokens - 2)
new_tokens = max(0, target_offset - prefix_hit_length)
prefill_input = remaining[:new_tokens]
if int(prefill_input.shape[0]) > 0:
sampler = make_sampler(temp=1.0)
_ = mlx_prefill(
model=model,
tokenizer=tokenizer,
sampler=sampler,
prompt_tokens=prefill_input,
cache=cache,
group=group,
on_prefill_progress=None,
distributed_prompt_progress_callback=None,
)
if kv_prefix_cache is not None:
try:
cache_snapshots = [snapshot_ssm_states(cache)]
hit_ratio = prefix_hit_length / n_tokens if n_tokens > 0 else 0.0
if matched_index is not None and hit_ratio >= 0.5:
kv_prefix_cache.update_kv_cache(
matched_index,
prompt_tokens,
cache,
cache_snapshots,
restore_pos=prefix_hit_length,
)
else:
kv_prefix_cache.add_kv_cache(prompt_tokens, cache, cache_snapshots)
except Exception:
logger.opt(exception=True).warning(
"Failed to save prefix cache on prefill server"
)
elapsed = time.perf_counter() - t0
final_offset = cache_length(cache)
logger.info(
f"Prefill: request_id={request.request_id} "
f"{n_tokens} tokens (prefix_hit={prefix_hit_length}, "
f"final_offset={final_offset}) in {elapsed * 1000:.0f}ms"
)
return cache
Whitespace-only changes.
@@ -0,0 +1,167 @@
from typing import BinaryIO
import mlx.core as mx
import numpy as np
import pytest
from mlx_lm.models.cache import KVCache
from exo.worker.disaggregated.protocol import Header, write_done, write_header
from exo.worker.disaggregated.server import PrefillRequest, PrefillServer
from exo.worker.engines.mlx.disaggregated.adapter import (
send_mlx_kv_cache,
wire_dtype_from_cache,
)
from exo.worker.engines.mlx.disaggregated.client import (
ingest_into_mlx_cache,
remote_prefill_fetch,
)
def _equal(a: mx.array, b: mx.array) -> bool:
if a.dtype != b.dtype or tuple(a.shape) != tuple(b.shape):
return False
if a.dtype == mx.bfloat16:
return bool(
np.array_equal(np.asarray(a.view(mx.uint16)), np.asarray(b.view(mx.uint16)))
)
return bool(np.array_equal(np.asarray(a), np.asarray(b)))
def _make_cache(seq_len: int, n_heads: int, head_dim: int) -> KVCache:
mx.random.seed(0)
cache = KVCache()
with mx.stream(mx.Device(mx.cpu)):
cache.keys = (
mx.random.uniform(shape=(1, n_heads, seq_len, head_dim)) * 10
).astype(mx.bfloat16)
cache.values = (
mx.random.uniform(shape=(1, n_heads, seq_len, head_dim)) * 10
).astype(mx.bfloat16)
mx.eval(cache.keys, cache.values)
cache.offset = seq_len
return cache
def _stream_cache(
wfile: BinaryIO, cache: KVCache, *, request_id: str, start_pos: int = 0
) -> None:
dtype = wire_dtype_from_cache([cache])
write_header(
wfile,
Header(
request_id=request_id,
model_id="test-model",
num_layers=1,
dtype=dtype,
start_pos=start_pos,
),
)
tokens_sent = send_mlx_kv_cache(wfile, [cache], dtype=dtype, start_pos=start_pos)
write_done(wfile, tokens_sent)
wfile.flush()
@pytest.mark.slow
def test_server_client_roundtrip() -> None:
seq_len = 5
n_heads = 2
head_dim = 4
gold = _make_cache(seq_len, n_heads, head_dim)
def resolve(job: PrefillRequest, wfile: BinaryIO) -> bool:
_stream_cache(wfile, gold, request_id=job.request_id)
return True
server = PrefillServer(resolve=resolve, host="127.0.0.1", port=52417)
try:
result = remote_prefill_fetch(
endpoint="127.0.0.1:52417",
request=PrefillRequest(
model_id="test-model",
token_ids=list(range(seq_len)),
request_id="req-1",
),
)
assert result.total_tokens == seq_len
assert 0 in result.kv_chunks
dst = KVCache()
final_offset = ingest_into_mlx_cache(result, [dst])
assert final_offset == seq_len
assert dst.offset == seq_len
dst_k = dst.keys
dst_v = dst.values
gold_k = gold.keys
gold_v = gold.values
assert dst_k is not None and dst_v is not None
assert gold_k is not None and gold_v is not None
assert _equal(dst_k, gold_k)
assert _equal(dst_v, gold_v)
finally:
server.stop()
@pytest.mark.slow
def test_server_reports_pickup_failure() -> None:
def resolve(_job: PrefillRequest, _wfile: BinaryIO) -> bool:
return False
server = PrefillServer(resolve=resolve, host="127.0.0.1", port=52418)
try:
with pytest.raises(RuntimeError, match="not picked up"):
_ = remote_prefill_fetch(
endpoint="127.0.0.1:52418",
request=PrefillRequest(
model_id="test-model",
token_ids=[1, 2, 3],
request_id="never-registered",
),
)
finally:
server.stop()
@pytest.mark.slow
def test_server_client_roundtrip_with_start_pos() -> None:
seq_len = 8
start_pos = 5
n_heads = 2
head_dim = 4
gold = _make_cache(seq_len, n_heads, head_dim)
def resolve(job: PrefillRequest, wfile: BinaryIO) -> bool:
_stream_cache(wfile, gold, request_id=job.request_id, start_pos=start_pos)
return True
server = PrefillServer(resolve=resolve, host="127.0.0.1", port=52419)
try:
result = remote_prefill_fetch(
endpoint="127.0.0.1:52419",
request=PrefillRequest(
model_id="test-model",
token_ids=list(range(seq_len)),
request_id="req-1",
start_pos=start_pos,
),
)
assert result.total_tokens == seq_len - start_pos
assert result.header.start_pos == start_pos
dst = KVCache()
gold_k = gold.keys
gold_v = gold.values
assert gold_k is not None and gold_v is not None
dst.keys = mx.array(gold_k[:, :, :start_pos, :])
dst.values = mx.array(gold_v[:, :, :start_pos, :])
dst.offset = start_pos
final_offset = ingest_into_mlx_cache(result, [dst], start_pos=start_pos)
assert final_offset == seq_len
assert dst.offset == seq_len
dst_k = dst.keys
dst_v = dst.values
assert dst_k is not None and dst_v is not None
assert _equal(dst_k, gold_k)
assert _equal(dst_v, gold_v)
finally:
server.stop()
@@ -0,0 +1,270 @@
import io
import mlx.core as mx
import numpy as np
from mlx_lm.models.cache import ArraysCache, KVCache, RotatingKVCache
from exo.worker.disaggregated.protocol import (
ArraysState,
Done,
Header,
KVChunk,
TensorBlob,
read_header,
read_message,
write_done,
write_header,
)
from exo.worker.engines.mlx.disaggregated.adapter import (
array_to_bytes,
bhsd_to_nhd,
bytes_to_array,
chunk_to_mlx_nhd,
inject_arrays_cache,
inject_kv_chunk,
inject_rotating_kv_chunk,
nhd_to_bhsd,
send_mlx_kv_cache,
wire_dtype_from_cache,
)
from exo.worker.engines.mlx.disaggregated.client import (
PrefillResult,
ingest_into_mlx_cache,
)
def _equal(a: mx.array, b: mx.array) -> bool:
if a.dtype != b.dtype or tuple(a.shape) != tuple(b.shape):
return False
if a.dtype == mx.bfloat16:
return bool(
np.array_equal(np.asarray(a.view(mx.uint16)), np.asarray(b.view(mx.uint16)))
)
return bool(np.array_equal(np.asarray(a), np.asarray(b)))
def _rand(shape: tuple[int, ...], dtype: mx.Dtype) -> mx.array:
mx.random.seed(0)
return (mx.random.uniform(shape=shape) * 10).astype(dtype)
def _make_kv_cache(seq_len: int, n_heads: int, head_dim: int) -> KVCache:
cache = KVCache()
cache.keys = _rand((1, n_heads, seq_len, head_dim), mx.bfloat16)
cache.values = _rand((1, n_heads, seq_len, head_dim), mx.bfloat16)
cache.offset = seq_len
return cache
def test_bytes_roundtrip_bf16() -> None:
x = _rand((2, 3, 4), mx.bfloat16)
y = bytes_to_array(array_to_bytes(x), (2, 3, 4), "bfloat16")
assert _equal(x, y)
def test_bytes_roundtrip_f16() -> None:
x = _rand((5,), mx.float16)
y = bytes_to_array(array_to_bytes(x), (5,), "float16")
assert _equal(x, y)
def test_bytes_roundtrip_f32() -> None:
x = _rand((2, 2), mx.float32)
y = bytes_to_array(array_to_bytes(x), (2, 2), "float32")
assert _equal(x, y)
def test_bhsd_nhd_roundtrip() -> None:
bhsd = _rand((1, 4, 7, 8), mx.float32)
nhd = bhsd_to_nhd(bhsd)
assert tuple(nhd.shape) == (7, 4, 8)
back = nhd_to_bhsd(nhd)
assert _equal(bhsd, back)
def test_kv_cache_inject_roundtrip() -> None:
n_heads, seq_len, head_dim = 3, 5, 4
k_bhsd = _rand((1, n_heads, seq_len, head_dim), mx.float32)
v_bhsd = _rand((1, n_heads, seq_len, head_dim), mx.float32)
k_nhd = bhsd_to_nhd(k_bhsd)
v_nhd = bhsd_to_nhd(v_bhsd)
cache = KVCache()
inject_kv_chunk(cache, k_nhd, v_nhd, offset=seq_len)
assert cache.offset == seq_len
assert cache.keys is not None and cache.values is not None
assert _equal(cache.keys, k_bhsd)
assert _equal(cache.values, v_bhsd)
def test_arrays_cache_inject() -> None:
a = _rand((3,), mx.float32)
b = _rand((2, 2), mx.bfloat16)
blobs = [
TensorBlob(dtype="float32", shape=(3,), data=array_to_bytes(a)),
TensorBlob(dtype="bfloat16", shape=(2, 2), data=array_to_bytes(b)),
]
cache = ArraysCache(size=2)
inject_arrays_cache(cache, blobs)
s0 = cache.state[0]
s1 = cache.state[1]
assert s0 is not None and s1 is not None
assert _equal(s0, a)
assert _equal(s1, b)
def test_send_mlx_cache_end_to_end() -> None:
n_heads, head_dim = 2, 4
seq_len = 3
src = _make_kv_cache(seq_len, n_heads, head_dim)
k_bhsd, v_bhsd = src.keys, src.values
assert k_bhsd is not None and v_bhsd is not None
buf = io.BytesIO()
write_header(buf, Header(num_layers=1, dtype="bfloat16"))
tokens = send_mlx_kv_cache(buf, [src], dtype="bfloat16")
write_done(buf, tokens)
buf.seek(0)
got_hdr = read_header(buf)
assert got_hdr.num_layers == 1
msg = read_message(buf)
assert isinstance(msg, KVChunk)
assert msg.num_tokens == seq_len
k_nhd, v_nhd = chunk_to_mlx_nhd(msg)
dst = KVCache()
inject_kv_chunk(dst, k_nhd, v_nhd, offset=msg.num_tokens)
done = read_message(buf)
assert isinstance(done, Done)
assert done.total_tokens == seq_len
assert dst.offset == seq_len
assert dst.keys is not None and dst.values is not None
assert _equal(dst.keys, k_bhsd)
assert _equal(dst.values, v_bhsd)
_ = ArraysState
def test_send_with_start_pos_only_ships_suffix() -> None:
n_heads, head_dim = 2, 4
seq_len, start_pos = 6, 4
src = _make_kv_cache(seq_len, n_heads, head_dim)
buf = io.BytesIO()
write_header(buf, Header(num_layers=1, dtype="bfloat16", start_pos=start_pos))
tokens = send_mlx_kv_cache(buf, [src], dtype="bfloat16", start_pos=start_pos)
write_done(buf, tokens)
buf.seek(0)
_ = read_header(buf)
msg = read_message(buf)
assert isinstance(msg, KVChunk)
assert msg.num_tokens == seq_len - start_pos
def test_send_skips_layer_when_offset_below_start_pos() -> None:
n_heads, head_dim = 2, 4
seq_len, start_pos = 3, 5
src = _make_kv_cache(seq_len, n_heads, head_dim)
buf = io.BytesIO()
write_header(buf, Header(num_layers=1, dtype="bfloat16", start_pos=start_pos))
tokens = send_mlx_kv_cache(buf, [src], dtype="bfloat16", start_pos=start_pos)
write_done(buf, tokens)
buf.seek(0)
_ = read_header(buf)
msg = read_message(buf)
assert isinstance(msg, Done)
assert msg.total_tokens == 0
assert tokens == 0
def test_wire_dtype_from_cache() -> None:
src = _make_kv_cache(3, 2, 4)
assert wire_dtype_from_cache([src]) == "bfloat16"
f32 = KVCache()
f32.keys = _rand((1, 2, 3, 4), mx.float32)
f32.values = _rand((1, 2, 3, 4), mx.float32)
f32.offset = 3
assert wire_dtype_from_cache([f32]) == "float32"
def _decode_payload(payload: bytes) -> PrefillResult:
buf = io.BytesIO(payload)
hdr = read_header(buf)
result = PrefillResult(header=hdr)
while True:
msg = read_message(buf)
if msg is None:
break
if isinstance(msg, KVChunk):
result.kv_chunks.setdefault(msg.layer_idx, []).append(msg)
elif isinstance(msg, ArraysState):
result.arrays[msg.layer_idx] = msg.arrays
elif isinstance(msg, Done):
result.total_tokens = msg.total_tokens
break
return result
def test_mixed_cache_roundtrip() -> None:
n_heads, head_dim, seq_len = 2, 4, 6
src_kv = _make_kv_cache(seq_len, n_heads, head_dim)
src_rot = RotatingKVCache(max_size=16, keep=0)
src_rot.keys = _rand((1, n_heads, seq_len, head_dim), mx.bfloat16)
src_rot.values = _rand((1, n_heads, seq_len, head_dim), mx.bfloat16)
src_rot.offset = seq_len
src_rot._idx = seq_len
src_arr = ArraysCache(size=2)
arr_a = _rand((3,), mx.bfloat16)
arr_b = _rand((2, 4), mx.bfloat16)
src_arr.state = [arr_a, arr_b]
buf = io.BytesIO()
write_header(
buf,
Header(request_id="req", model_id="m", num_layers=3, dtype="bfloat16"),
)
tokens_sent = send_mlx_kv_cache(buf, [src_kv, src_rot, src_arr], dtype="bfloat16")
write_done(buf, tokens_sent)
result = _decode_payload(buf.getvalue())
assert result.header.num_layers == 3
assert result.total_tokens == seq_len
dst_kv = KVCache()
dst_rot = RotatingKVCache(max_size=16, keep=0)
dst_arr = ArraysCache(size=2)
final_offset = ingest_into_mlx_cache(result, [dst_kv, dst_rot, dst_arr])
assert final_offset == seq_len
assert dst_kv.offset == seq_len
assert dst_kv.keys is not None and dst_kv.values is not None
src_kv_k, src_kv_v = src_kv.keys, src_kv.values
assert src_kv_k is not None and src_kv_v is not None
assert _equal(dst_kv.keys, src_kv_k)
assert _equal(dst_kv.values, src_kv_v)
assert dst_rot.offset == seq_len
assert dst_rot.keys is not None and dst_rot.values is not None
src_rot_k, src_rot_v = src_rot.keys, src_rot.values
assert src_rot_k is not None and src_rot_v is not None
assert _equal(dst_rot.keys, src_rot_k)
assert _equal(dst_rot.values, src_rot_v)
assert dst_rot._idx == seq_len
assert len(dst_arr.state) == 2
s0, s1 = dst_arr.state[0], dst_arr.state[1]
assert s0 is not None and s1 is not None
assert _equal(s0, arr_a)
assert _equal(s1, arr_b)
_ = inject_rotating_kv_chunk
_ = nhd_to_bhsd
@@ -0,0 +1,154 @@
import io
import pytest
from exo.worker.disaggregated.protocol import (
ArraysState,
Done,
ErrorMessage,
Header,
KVChunk,
ProtocolError,
TensorBlob,
read_header,
read_message,
write_arrays_state,
write_done,
write_error,
write_header,
write_kv_chunk,
)
def _mk_bytes(n: int) -> bytes:
return bytes(i & 0xFF for i in range(n))
def test_header_roundtrip() -> None:
hdr = Header(
request_id="r",
model_id="m",
num_layers=32,
dtype="bfloat16",
start_pos=42,
)
buf = io.BytesIO()
write_header(buf, hdr)
buf.seek(0)
got = read_header(buf)
assert got == hdr
assert got.dtype == "bfloat16"
assert got.num_layers == 32
assert got.start_pos == 42
def test_kv_chunk_roundtrip() -> None:
num_tokens, n_heads, head_dim = 7, 4, 8
n_bytes = num_tokens * n_heads * head_dim * 2
keys = _mk_bytes(n_bytes)
values = _mk_bytes(n_bytes)[::-1]
buf = io.BytesIO()
write_kv_chunk(
buf,
layer_idx=3,
num_tokens=num_tokens,
n_heads=n_heads,
head_dim=head_dim,
dtype="bfloat16",
keys=keys,
values=values,
)
buf.seek(0)
msg = read_message(buf)
assert isinstance(msg, KVChunk)
assert msg.layer_idx == 3
assert msg.shape == (num_tokens, n_heads, head_dim)
assert msg.dtype == "bfloat16"
assert msg.keys == keys
assert msg.values == values
def test_arrays_state_roundtrip() -> None:
arrs = [
TensorBlob(dtype="float32", shape=(2, 3), data=_mk_bytes(2 * 3 * 4)),
TensorBlob(dtype="bfloat16", shape=(5,), data=_mk_bytes(5 * 2)),
]
buf = io.BytesIO()
write_arrays_state(buf, layer_idx=9, arrays=arrs)
buf.seek(0)
msg = read_message(buf)
assert isinstance(msg, ArraysState)
assert msg.layer_idx == 9
assert len(msg.arrays) == 2
assert msg.arrays[0].dtype == "float32"
assert msg.arrays[0].shape == (2, 3)
assert msg.arrays[0].data == arrs[0].data
assert msg.arrays[1].dtype == "bfloat16"
assert msg.arrays[1].shape == (5,)
assert msg.arrays[1].data == arrs[1].data
def test_done_roundtrip() -> None:
buf = io.BytesIO()
write_done(buf, 1234)
buf.seek(0)
msg = read_message(buf)
assert isinstance(msg, Done)
assert msg.total_tokens == 1234
def test_error_roundtrip() -> None:
buf = io.BytesIO()
write_error(buf, code=42, message="boom")
buf.seek(0)
msg = read_message(buf)
assert isinstance(msg, ErrorMessage)
assert msg.code == 42
assert msg.message == "boom"
def test_stream_of_messages() -> None:
hdr = Header(num_layers=2, dtype="float32")
buf = io.BytesIO()
write_header(buf, hdr)
write_kv_chunk(
buf,
layer_idx=0,
num_tokens=1,
n_heads=1,
head_dim=2,
dtype="float32",
keys=_mk_bytes(1 * 1 * 2 * 4),
values=_mk_bytes(1 * 1 * 2 * 4),
)
write_arrays_state(
buf,
layer_idx=1,
arrays=[TensorBlob(dtype="float32", shape=(1,), data=_mk_bytes(4))],
)
write_done(buf, total_tokens=1)
buf.seek(0)
got_hdr = read_header(buf)
assert got_hdr == hdr
m1 = read_message(buf)
m2 = read_message(buf)
m3 = read_message(buf)
m4 = read_message(buf)
assert isinstance(m1, KVChunk)
assert isinstance(m2, ArraysState)
assert isinstance(m3, Done)
assert m4 is None
def test_corrupt_message_raises() -> None:
buf = io.BytesIO()
write_header(buf, Header(num_layers=1, dtype="float32"))
buf.write((5).to_bytes(4, "big"))
buf.write(b"\xff\xff\xff\xff\xff")
buf.seek(0)
_ = read_header(buf)
with pytest.raises(ProtocolError):
_ = read_message(buf)
@@ -0,0 +1,120 @@
"""Server thread receives request, runs resolve in another thread (mimicking
runner main thread + work queue), streams cache bytes."""
import queue
import threading
from typing import BinaryIO
import mlx.core as mx
import numpy as np
import pytest
from mlx_lm.models.cache import KVCache
from exo.utils.ports import random_ephemeral_port
from exo.worker.disaggregated.protocol import Header, write_done, write_header
from exo.worker.disaggregated.server import PrefillRequest, PrefillServer
from exo.worker.engines.mlx.disaggregated.adapter import (
send_mlx_kv_cache,
wire_dtype_from_cache,
)
from exo.worker.engines.mlx.disaggregated.client import (
PrefillResult,
ingest_into_mlx_cache,
remote_prefill_fetch,
)
def _equal(a: mx.array, b: mx.array) -> bool:
if a.dtype != b.dtype or tuple(a.shape) != tuple(b.shape):
return False
if a.dtype == mx.bfloat16:
return bool(
np.array_equal(np.asarray(a.view(mx.uint16)), np.asarray(b.view(mx.uint16)))
)
return bool(np.array_equal(np.asarray(a), np.asarray(b)))
def _make_cache(seq_len: int, n_heads: int, head_dim: int) -> KVCache:
mx.random.seed(0)
cache = KVCache()
cache.keys = (mx.random.uniform(shape=(1, n_heads, seq_len, head_dim)) * 10).astype(
mx.bfloat16
)
cache.values = (
mx.random.uniform(shape=(1, n_heads, seq_len, head_dim)) * 10
).astype(mx.bfloat16)
cache.offset = seq_len
return cache
@pytest.mark.slow
def test_server_drains_via_main_thread() -> None:
seq_len = 4
n_heads = 2
head_dim = 4
gold = _make_cache(seq_len, n_heads, head_dim)
request_queue: queue.Queue[tuple[PrefillRequest, BinaryIO, threading.Event]] = (
queue.Queue()
)
def resolve(job: PrefillRequest, wfile: BinaryIO) -> bool:
done = threading.Event()
request_queue.put((job, wfile, done))
return done.wait(timeout=5)
server = PrefillServer(
resolve=resolve, host="127.0.0.1", port=(port := random_ephemeral_port())
)
def serve_one(wfile: BinaryIO) -> None:
dtype = wire_dtype_from_cache([gold])
write_header(
wfile,
Header(request_id="req-1", model_id="m", num_layers=1, dtype=dtype),
)
tokens = send_mlx_kv_cache(wfile, [gold], dtype=dtype)
write_done(wfile, tokens)
wfile.flush()
drained_job: list[PrefillRequest] = []
fetch_result: list[PrefillResult] = []
def fetcher() -> None:
fetch_result.append(
remote_prefill_fetch(
endpoint=f"127.0.0.1:{port}",
request=PrefillRequest(
model_id="m", token_ids=list(range(seq_len)), request_id="req-1"
),
)
)
fetch = threading.Thread(target=fetcher, daemon=True)
fetch.start()
try:
job, wfile, done = request_queue.get(timeout=5)
drained_job.append(job)
try:
serve_one(wfile)
finally:
done.set()
fetch.join(timeout=5)
assert fetch_result, "fetcher did not return"
result = fetch_result[0]
assert drained_job[0].request_id == "req-1"
assert result.total_tokens == seq_len
dst = KVCache()
ingest_into_mlx_cache(result, [dst])
assert dst.offset == seq_len
dst_k = dst.keys
dst_v = dst.values
gold_k = gold.keys
gold_v = gold.values
assert dst_k is not None and dst_v is not None
assert gold_k is not None and gold_v is not None
assert _equal(dst_k, gold_k)
assert _equal(dst_v, gold_v)
finally:
server.stop()
@@ -1,5 +1,6 @@
import contextlib
import time
import uuid
from dataclasses import dataclass, field
from typing import Callable, Literal, cast
@@ -23,7 +24,6 @@ from exo.api.types import (
Usage,
)
from exo.shared.types.memory import Memory
from exo.shared.types.mlx import KVCacheType, Model
from exo.shared.types.text_generation import TextGenerationTaskParams
from exo.shared.types.worker.runner_response import GenerationResponse
from exo.worker.engines.mlx.cache import (
@@ -40,10 +40,12 @@ from exo.worker.engines.mlx.generator.generate import (
patch_embed_tokens,
prefill,
)
from exo.worker.engines.mlx.generator.remote_prefill import remote_prefill
from exo.worker.engines.mlx.patches.opt_batch_gen import (
set_needs_topk,
take_ready_topk,
)
from exo.worker.engines.mlx.types import KVCacheType, Model
from exo.worker.engines.mlx.utils_mlx import (
fix_unmatched_think_end_tokens,
system_prompt_token_count,
@@ -57,6 +59,7 @@ from exo.worker.engines.mlx.vision import (
from exo.worker.runner.bootstrap import logger
_MIN_PREFIX_HIT_RATIO_TO_UPDATE = 0.5
REMOTE_PREFILL_MIN_TOKENS = 1000
def _stop_sequences(task_params: TextGenerationTaskParams) -> list[str]:
@@ -199,17 +202,49 @@ class ExoBatchGenerator:
if vision is not None
else contextlib.nullcontext()
)
uncached_count = len(prompt_tokens)
use_remote = (
uncached_count > REMOTE_PREFILL_MIN_TOKENS
and task_params.prefill_endpoint is not None
)
_prefill_tps: float = 0.0
_prefill_tokens: int = 0
cache_snapshots: list[CacheSnapshot] = []
remote_prefilled = False
with vision_ctx:
_prefill_tps, _prefill_tokens, cache_snapshots = prefill(
self.model,
self.tokenizer,
sampler,
prompt_tokens[:-1],
cache,
self.group,
on_prefill_progress,
distributed_prompt_progress_callback,
)
if use_remote and task_params.prefill_endpoint is not None:
try:
# Send full prompt; producer's vLLM APC handles the prefix
# match. `start_pos` aligns the writer's skip_tokens with
# the consumer's locally-cached prefix.
_prefill_tps, _prefill_tokens, cache_snapshots = remote_prefill(
all_prompt_tokens[:-1],
cache,
on_prefill_progress,
endpoint=task_params.prefill_endpoint,
request_id=str(uuid.uuid4()),
model_id=str(task_params.model),
start_pos=prefix_hit_length,
use_prefix_cache=not is_bench or task_params.use_prefix_cache,
)
remote_prefilled = True
except Exception:
logger.opt(exception=True).warning(
"Remote prefill failed, falling back to local prefill"
)
if not remote_prefilled:
_prefill_tps, _prefill_tokens, cache_snapshots = prefill(
self.model,
self.tokenizer,
sampler,
prompt_tokens[:-1],
cache,
self.group,
on_prefill_progress,
distributed_prompt_progress_callback,
)
prefix_cache_hit: Literal["none", "partial", "exact"] = "none"
if matched_index is not None and prefix_hit_length > 0:
@@ -457,6 +492,7 @@ class ExoBatchGenerator:
def close(self) -> None:
self._mlx_gen.close()
mx.clear_cache()
def _save_prefix_cache(
self,
@@ -2,6 +2,7 @@ import contextlib
import functools
import math
import time
import uuid
from typing import Callable, Generator, cast, get_args
import mlx.core as mx
@@ -9,7 +10,6 @@ from mlx_lm.generate import (
maybe_quantize_kv_cache,
stream_generate,
)
from mlx_lm.models.cache import ArraysCache, CacheList, RotatingKVCache
from mlx_lm.sample_utils import make_logits_processors, make_sampler
from mlx_lm.tokenizer_utils import TokenizerWrapper
@@ -23,7 +23,6 @@ from exo.api.types import (
)
from exo.shared.types.common import ModelId
from exo.shared.types.memory import Memory
from exo.shared.types.mlx import KVCacheType, Model
from exo.shared.types.text_generation import (
InputMessage,
InputMessageContent,
@@ -43,10 +42,11 @@ from exo.worker.engines.mlx.auto_parallel import (
from exo.worker.engines.mlx.cache import (
CacheSnapshot,
KVPrefixCache,
copy_snapshot_entry,
encode_prompt,
has_non_kv_caches,
is_non_trimmable_cache_entry,
make_kv_cache,
restore_snapshot_entry,
snapshot_ssm_states,
)
from exo.worker.engines.mlx.constants import (
@@ -55,6 +55,8 @@ from exo.worker.engines.mlx.constants import (
KV_GROUP_SIZE,
MAX_TOKENS,
)
from exo.worker.engines.mlx.generator.remote_prefill import remote_prefill
from exo.worker.engines.mlx.types import KVCacheType, Model
from exo.worker.engines.mlx.utils_mlx import (
apply_chat_template,
fix_unmatched_think_end_tokens,
@@ -70,6 +72,8 @@ from exo.worker.engines.mlx.vision import (
)
from exo.worker.runner.bootstrap import logger
REMOTE_PREFILL_MIN_TOKENS = 1000
generation_stream = mx.new_stream(mx.default_device())
_MIN_PREFIX_HIT_RATIO_TO_UPDATE = 0.5
@@ -372,12 +376,10 @@ def prefill(
# Because of needing to roll back arrays cache, we will generate on 2 tokens so trim 1 more.
pre_gen = snapshots[-2] if has_ssm else None
for i, c in enumerate(cache):
non_trimmable = isinstance(c, (ArraysCache, RotatingKVCache)) or (
isinstance(c, CacheList) and not bool(c.is_trimmable()) # type: ignore[reportUnknownMemberType]
)
non_trimmable = is_non_trimmable_cache_entry(c)
if has_ssm and non_trimmable:
assert pre_gen is not None
restored = restore_snapshot_entry(pre_gen.states[i])
restored = copy_snapshot_entry(pre_gen.states[i])
if restored is not None:
cache[i] = restored # type: ignore
else:
@@ -635,17 +637,48 @@ def mlx_generate(
if vision is not None
else contextlib.nullcontext()
)
use_remote = (
len(prompt_tokens) > REMOTE_PREFILL_MIN_TOKENS
and task.prefill_endpoint is not None
)
remote_prefilled = False
prefill_tps = 0.0
prefill_tokens = 0
ssm_snapshots_list: list[CacheSnapshot] = []
with maybe_vision_ctx:
prefill_tps, prefill_tokens, ssm_snapshots_list = prefill(
model,
tokenizer,
sampler,
prompt_tokens[:-1],
caches,
group,
on_prefill_progress,
distributed_prompt_progress_callback,
)
if use_remote and task.prefill_endpoint is not None:
try:
# Send the FULL prompt to the producer (not the cache-stripped
# suffix). vLLM's APC handles the prefix match internally;
# `start_pos` tells our extractor / wire writer how much of the
# producer-side capture corresponds to tokens the consumer
# already has, so the writer's skip_tokens math aligns.
prefill_tps, prefill_tokens, ssm_snapshots_list = remote_prefill(
all_prompt_tokens[:-1],
caches,
on_prefill_progress,
endpoint=task.prefill_endpoint,
request_id=str(uuid.uuid4()),
model_id=str(task.model),
start_pos=prefix_hit_length,
use_prefix_cache=not is_bench or task.use_prefix_cache,
)
remote_prefilled = True
except Exception:
logger.opt(exception=True).warning(
"Remote prefill failed, falling back to local prefill"
)
if not remote_prefilled:
prefill_tps, prefill_tokens, ssm_snapshots_list = prefill(
model,
tokenizer,
sampler,
prompt_tokens[:-1],
caches,
group,
on_prefill_progress,
distributed_prompt_progress_callback,
)
cache_snapshots: list[CacheSnapshot] | None = ssm_snapshots_list or None
if kv_prefix_cache is not None and matched_index is not None and is_exact_hit:
@@ -0,0 +1,97 @@
import time
from collections.abc import Callable
from typing import cast
import mlx.core as mx
from mlx_lm.models.cache import ArraysCache, KVCache, RotatingKVCache
from exo.worker.disaggregated.protocol import Header, KVChunk
from exo.worker.disaggregated.server import PrefillRequest
from exo.worker.engines.mlx.cache import CacheSnapshot, snapshot_ssm_states
from exo.worker.engines.mlx.disaggregated.client import (
ingest_into_mlx_cache,
remote_prefill_fetch,
)
from exo.worker.engines.mlx.types import KVCacheType
from exo.worker.runner.bootstrap import logger
def remote_prefill(
prompt_tokens: mx.array,
cache: KVCacheType,
on_prefill_progress: Callable[[int, int], None] | None,
*,
endpoint: str,
request_id: str,
model_id: str,
start_pos: int = 0,
use_prefix_cache: bool = True,
) -> tuple[float, int, list[CacheSnapshot]]:
t0 = time.perf_counter()
total_prompt_tokens = int(prompt_tokens.shape[0])
num_layers: int = 0
tokens_received_total: int = 0
def _on_header(header: Header) -> None:
nonlocal num_layers
num_layers = header.num_layers
def _on_chunk(chunk: KVChunk, chunks_received: int) -> None:
nonlocal num_layers, tokens_received_total
tokens_received_total += chunk.num_tokens
if on_prefill_progress is None:
return
if num_layers > 0 and chunks_received % num_layers == 0:
on_prefill_progress(
min(tokens_received_total // num_layers, total_prompt_tokens),
total_prompt_tokens,
)
request = PrefillRequest(
model_id=model_id,
token_ids=cast(list[int], prompt_tokens.tolist()),
start_pos=start_pos,
request_id=request_id,
use_prefix_cache=use_prefix_cache,
)
result = remote_prefill_fetch(
endpoint, request, on_header=_on_header, on_kv_chunk=_on_chunk
)
t_received = time.perf_counter()
caches = cast(list[KVCache | RotatingKVCache | ArraysCache], list(cache))
final_offset = ingest_into_mlx_cache(result, caches, start_pos=start_pos)
t_done = time.perf_counter()
num_tokens = final_offset - start_pos
# The producer strips the last 2 tokens of the prompt (consumer warm-starts
# decode from those locally). Anything within `producer_strip` of the full
# suffix is the expected outcome, not a bug.
producer_strip = 2 if total_prompt_tokens > 2 else 0
expected_min = max(0, total_prompt_tokens - start_pos - producer_strip)
expected_max = max(0, total_prompt_tokens - start_pos)
if num_tokens <= 0:
raise RuntimeError(
f"Remote prefill returned no KV (start_pos={start_pos}, "
f"final_offset={final_offset}, expected={expected_min}, "
f"transfer={(t_received - t0) * 1000:.0f}ms)"
)
if num_tokens < expected_min:
logger.warning(
f"Remote prefill returned {num_tokens} tokens, expected at least "
f"{expected_min} (start_pos={start_pos}, final_offset={final_offset})"
)
elif num_tokens > expected_max:
logger.warning(
f"Remote prefill returned {num_tokens} tokens, expected at most "
f"{expected_max} (start_pos={start_pos}, final_offset={final_offset})"
)
tps = num_tokens / max(t_done - t0, 0.001)
logger.info(
f"Remote prefill: {num_tokens} tokens (start_pos={start_pos}, "
f"final_offset={final_offset}) at {tps:.0f} tok/s, "
f"transfer={(t_received - t0) * 1000:.0f}ms, "
f"inject={(t_done - t_received) * 1000:.0f}ms"
)
return tps, num_tokens, [snapshot_ssm_states(cache)]
Loaded 100 of 137 files, more files were not shown because too many files have changed in this diff. Show more