mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-09 03:51:22 -04:00
Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8bc74d2cea | ||
|
|
5ece607264 | ||
|
|
8106d8a7e3 | ||
|
|
2ef3a4c707 | ||
|
|
bba012f15b | ||
|
|
3babf9d070 | ||
|
|
5d7ea4c6c0 | ||
|
|
e116097f64 | ||
|
|
c6467094b1 |
No files matched your search
@@ -1 +1,8 @@
|
||||
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
|
||||
@@ -191,10 +191,13 @@ 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: tuple[str, ...]) -> None: ...
|
||||
def is_trimmable(self) -> bool: ...
|
||||
def meta_state(self, v): # -> None:
|
||||
...
|
||||
def is_trimmable(self): # -> bool:
|
||||
...
|
||||
def trim(self, n: int) -> int: ...
|
||||
def to_quantized(
|
||||
self, group_size: int = ..., bits: int = ...
|
||||
|
||||
@@ -108,10 +108,12 @@ class Compressor(nn.Module):
|
||||
def __call__(
|
||||
self,
|
||||
x: mx.array,
|
||||
cache: "DeepseekV4Cache",
|
||||
state: ArraysCache,
|
||||
offset: Any,
|
||||
key: str = ...,
|
||||
) -> mx.array: ...
|
||||
slot_compressed: int,
|
||||
slot_kv_state: int,
|
||||
slot_score_state: int,
|
||||
) -> Optional[mx.array]: ...
|
||||
|
||||
class Indexer(nn.Module):
|
||||
def __init__(
|
||||
@@ -121,29 +123,14 @@ 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: RotatingKVCache
|
||||
local: Any
|
||||
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
@@ -552,24 +552,15 @@ struct SettingsView: View {
|
||||
let alert = NSAlert()
|
||||
alert.messageText = "Uninstall EXO"
|
||||
alert.informativeText = """
|
||||
This will remove EXO and all its components:
|
||||
This will remove EXO and all its system 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")
|
||||
|
||||
@@ -579,11 +570,11 @@ struct SettingsView: View {
|
||||
|
||||
let response = alert.runModal()
|
||||
if response == .alertFirstButtonReturn {
|
||||
performUninstall(keepModels: checkbox.state == .on)
|
||||
performUninstall()
|
||||
}
|
||||
}
|
||||
|
||||
private func performUninstall(keepModels: Bool) {
|
||||
private func performUninstall() {
|
||||
uninstallInProgress = true
|
||||
|
||||
controller.cancelPendingLaunch()
|
||||
@@ -593,7 +584,6 @@ struct SettingsView: View {
|
||||
DispatchQueue.global(qos: .utility).async {
|
||||
do {
|
||||
try NetworkSetupHelper.uninstall()
|
||||
try Self.removeExoDirectory(keepModels: keepModels)
|
||||
|
||||
DispatchQueue.main.async {
|
||||
LaunchAtLoginHelper.disable()
|
||||
@@ -617,23 +607,6 @@ 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 {
|
||||
|
||||
@@ -3,55 +3,25 @@
|
||||
# EXO Uninstaller Script
|
||||
#
|
||||
# This script removes all EXO system components that persist after deleting the app.
|
||||
# Run with: sudo ./uninstall-exo.sh [--keep-models]
|
||||
#
|
||||
# Options:
|
||||
# --keep-models Preserve ~/.exo/models when removing the EXO data directory.
|
||||
# Run with: sudo ./uninstall-exo.sh
|
||||
#
|
||||
# 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"
|
||||
# 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"
|
||||
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'
|
||||
@@ -99,17 +69,11 @@ else
|
||||
echo_warn "LaunchDaemon plist not found (already removed?)"
|
||||
fi
|
||||
|
||||
# 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
|
||||
# Remove the script and parent directory
|
||||
if [[ -f $SCRIPT_DEST ]]; then
|
||||
rm -f "$SCRIPT_DEST"
|
||||
echo_info "Removed network setup script"
|
||||
else
|
||||
echo_warn "Network setup script not found (already removed?)"
|
||||
fi
|
||||
|
||||
@@ -151,22 +115,6 @@ 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.
|
||||
@@ -196,10 +144,6 @@ 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)."
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
# name, patterns, reasoning
|
||||
#
|
||||
# Optional per-model overrides (CLI flags take priority over these):
|
||||
# temperature, top_p, max_tokens, reasoning_effort, enable_thinking
|
||||
# temperature, top_p, max_tokens, reasoning_effort
|
||||
#
|
||||
# Fallback defaults (when no per-model config):
|
||||
# reasoning: temperature=1.0, max_tokens=131072, reasoning_effort="high"
|
||||
@@ -18,9 +18,10 @@
|
||||
|
||||
# ─── Qwen3.5 (Feb 2026) ─────────────────────────────────────────────
|
||||
# Source: HuggingFace model cards (Qwen/Qwen3.5-*)
|
||||
# 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).
|
||||
# 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]]
|
||||
name = "Qwen3.5 2B"
|
||||
@@ -28,8 +29,7 @@ patterns = ["Qwen3.5-2B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 81920
|
||||
|
||||
[[model]]
|
||||
name = "Qwen3.5 9B"
|
||||
@@ -37,8 +37,7 @@ patterns = ["Qwen3.5-9B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 81920
|
||||
|
||||
[[model]]
|
||||
name = "Qwen3.5 27B"
|
||||
@@ -46,17 +45,15 @@ patterns = ["Qwen3.5-27B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 81920
|
||||
|
||||
[[model]]
|
||||
name = "Qwen3.5 35B A3B"
|
||||
patterns = ["Qwen3.5-35B-A3B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 81920
|
||||
|
||||
[[model]]
|
||||
name = "Qwen3.5 122B A10B"
|
||||
@@ -64,8 +61,7 @@ patterns = ["Qwen3.5-122B-A10B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 81920
|
||||
|
||||
[[model]]
|
||||
name = "Qwen3.5 397B A17B"
|
||||
@@ -73,14 +69,12 @@ patterns = ["Qwen3.5-397B-A17B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 81920
|
||||
|
||||
# ─── Qwen3 (Apr 2025) ───────────────────────────────────────────────
|
||||
# Source: HuggingFace model cards (Qwen/Qwen3-*)
|
||||
# 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
|
||||
# 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, 38912 for complex math/code
|
||||
|
||||
[[model]]
|
||||
@@ -89,7 +83,6 @@ patterns = ["Qwen3-0.6B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 38912
|
||||
|
||||
[[model]]
|
||||
@@ -98,7 +91,6 @@ patterns = ["Qwen3-30B-A3B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 38912
|
||||
|
||||
[[model]]
|
||||
@@ -107,7 +99,6 @@ patterns = ["Qwen3-235B-A22B"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 38912
|
||||
|
||||
[[model]]
|
||||
@@ -116,7 +107,6 @@ patterns = ["Qwen3-Next-80B-A3B-Thinking"]
|
||||
reasoning = true
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 38912
|
||||
|
||||
[[model]]
|
||||
@@ -139,9 +129,9 @@ max_tokens = 16384
|
||||
name = "Qwen3 Coder Next"
|
||||
patterns = ["Qwen3-Coder-Next"]
|
||||
reasoning = false
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
max_tokens = 121072
|
||||
temperature = 0.7
|
||||
top_p = 0.8
|
||||
max_tokens = 16384
|
||||
|
||||
# ─── GPT-OSS (OpenAI) ───────────────────────────────────────────────
|
||||
# Source: OpenAI GitHub README + HuggingFace discussion #21
|
||||
@@ -175,38 +165,10 @@ 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
|
||||
# max_tokens=121072 to match vllm eval (131072 context - 10000 safety margin)
|
||||
# Reasoning tasks: 131072 max_tokens; coding/SWE tasks: temp=0.7
|
||||
|
||||
[[model]]
|
||||
name = "GLM-5"
|
||||
@@ -214,8 +176,7 @@ patterns = ["GLM-5"]
|
||||
reasoning = true
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 131072
|
||||
|
||||
[[model]]
|
||||
name = "GLM 4.5 Air"
|
||||
@@ -230,8 +191,7 @@ patterns = ["GLM-4.7-"]
|
||||
reasoning = true
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 131072
|
||||
# Note: matches both GLM-4.7 and GLM-4.7-Flash
|
||||
|
||||
# ─── Kimi (Moonshot AI) ─────────────────────────────────────────────
|
||||
@@ -253,8 +213,7 @@ patterns = ["Kimi-K2.5"]
|
||||
reasoning = true
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 121072
|
||||
max_tokens = 131072
|
||||
|
||||
[[model]]
|
||||
name = "Kimi K2 Instruct"
|
||||
@@ -264,17 +223,7 @@ temperature = 0.6
|
||||
|
||||
# ─── MiniMax ─────────────────────────────────────────────────────────
|
||||
# Source: HuggingFace model cards + generation_config.json
|
||||
# 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
|
||||
# All models: temp=1.0, top_p=0.95, top_k=40
|
||||
|
||||
[[model]]
|
||||
name = "MiniMax M2.5"
|
||||
@@ -282,8 +231,6 @@ patterns = ["MiniMax-M2.5"]
|
||||
reasoning = true
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
enable_thinking = true
|
||||
max_tokens = 90000
|
||||
|
||||
[[model]]
|
||||
name = "MiniMax M2.1"
|
||||
@@ -304,8 +251,6 @@ 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
|
||||
|
||||
+99
-192
@@ -35,7 +35,6 @@ 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,
|
||||
@@ -80,7 +79,7 @@ def load_tokenizer_for_bench(model_id: str) -> Any:
|
||||
model_path = Path(
|
||||
snapshot_download(
|
||||
model_id,
|
||||
allow_patterns=["*.json", "*.py", "*.tiktoken", "*.model", "*.jinja"],
|
||||
allow_patterns=["*.json", "*.py", "*.tiktoken", "*.model"],
|
||||
)
|
||||
)
|
||||
|
||||
@@ -278,72 +277,28 @@ 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,
|
||||
}
|
||||
|
||||
if not stream:
|
||||
payload["stream"] = False
|
||||
t0 = time.perf_counter()
|
||||
out = client.post_bench_chat_completions(payload)
|
||||
elapsed = time.perf_counter() - t0
|
||||
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 {}
|
||||
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
|
||||
stats = out.get("generation_stats")
|
||||
|
||||
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},
|
||||
}
|
||||
# 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 ""
|
||||
|
||||
return {
|
||||
"elapsed_s": elapsed,
|
||||
@@ -470,11 +425,6 @@ 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",
|
||||
@@ -540,124 +490,81 @@ def main() -> int:
|
||||
logger.error("[exo-bench] tokenizer usable but prompt sizing failed")
|
||||
raise
|
||||
|
||||
# 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"
|
||||
)
|
||||
selected = settle_and_fetch_placements(
|
||||
client, full_model_id, args, settle_timeout=args.settle_timeout
|
||||
)
|
||||
|
||||
if reused_instance_id is not None:
|
||||
# Use the existing instance directly — skip placement iteration
|
||||
selected = []
|
||||
download_duration_s = None
|
||||
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 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:
|
||||
selected = settle_and_fetch_placements(
|
||||
client, full_model_id, args, settle_timeout=args.settle_timeout
|
||||
)
|
||||
|
||||
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 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")
|
||||
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:
|
||||
created_instance = False
|
||||
if preview is not None:
|
||||
instance = preview["instance"]
|
||||
instance_id = instance_id_from_instance(instance)
|
||||
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}"
|
||||
)
|
||||
|
||||
# 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)
|
||||
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
|
||||
|
||||
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}")
|
||||
time.sleep(1)
|
||||
|
||||
sampler: SystemMetricsSampler | None = None
|
||||
if not args.no_system_metrics and preview is not None:
|
||||
if not args.no_system_metrics:
|
||||
nids = node_ids_from_instance(instance)
|
||||
sampler = SystemMetricsSampler(
|
||||
ExoClient(args.host, args.port, timeout_s=30),
|
||||
@@ -666,20 +573,16 @@ 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):
|
||||
_do_one(client, pp_list[0], tg_list[0])
|
||||
run_one_completion(
|
||||
client,
|
||||
full_model_id,
|
||||
pp_list[0],
|
||||
tg_list[0],
|
||||
prompt_sizer,
|
||||
use_prefix_cache=args.use_prefix_cache,
|
||||
)
|
||||
logger.debug(f" warmup {i + 1}/{args.warmup} done")
|
||||
|
||||
# If pp and tg lists have same length, run in tandem (zip)
|
||||
@@ -701,7 +604,14 @@ def main() -> int:
|
||||
# Sequential: single request
|
||||
try:
|
||||
inf_t0 = time.monotonic()
|
||||
row, actual_pp_tokens = _do_one(client, pp, tg)
|
||||
row, actual_pp_tokens = run_one_completion(
|
||||
client,
|
||||
full_model_id,
|
||||
pp,
|
||||
tg,
|
||||
prompt_sizer,
|
||||
use_prefix_cache=args.use_prefix_cache,
|
||||
)
|
||||
inference_windows.append((inf_t0, time.monotonic()))
|
||||
except Exception as e:
|
||||
logger.error(e)
|
||||
@@ -850,12 +760,10 @@ 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} "
|
||||
@@ -880,16 +788,15 @@ def main() -> int:
|
||||
if placement_metrics:
|
||||
all_system_metrics.update(placement_metrics)
|
||||
|
||||
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}")
|
||||
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:
|
||||
|
||||
+56
-427
@@ -47,7 +47,6 @@ 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,
|
||||
@@ -63,15 +62,6 @@ 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
|
||||
@@ -281,7 +271,7 @@ def run_humaneval_test(
|
||||
|
||||
@dataclass
|
||||
class QuestionResult:
|
||||
question_id: int | str
|
||||
question_id: int
|
||||
prompt: str
|
||||
response: str
|
||||
extracted_answer: str | None
|
||||
@@ -291,11 +281,7 @@ 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
|
||||
@@ -531,10 +517,6 @@ 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(
|
||||
@@ -548,9 +530,6 @@ 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:
|
||||
@@ -567,12 +546,6 @@ 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",
|
||||
@@ -581,40 +554,19 @@ async def _call_api(
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
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:
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
if not content or not content.strip():
|
||||
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,
|
||||
@@ -626,14 +578,8 @@ 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,
|
||||
@@ -646,30 +592,8 @@ 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(
|
||||
@@ -694,16 +618,10 @@ 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
|
||||
@@ -734,21 +652,7 @@ 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
|
||||
@@ -756,13 +660,6 @@ 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":
|
||||
@@ -770,64 +667,16 @@ 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)
|
||||
@@ -848,50 +697,24 @@ async def evaluate_benchmark(
|
||||
raise ValueError(f"Unknown benchmark: {benchmark_name}")
|
||||
|
||||
async with semaphore:
|
||||
if instance_failed.is_set():
|
||||
return
|
||||
t0 = time.monotonic()
|
||||
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
|
||||
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,
|
||||
)
|
||||
elapsed = time.monotonic() - t0
|
||||
|
||||
if api_result is None:
|
||||
result = QuestionResult(
|
||||
question_id=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response="",
|
||||
extracted_answer=None,
|
||||
@@ -906,17 +729,13 @@ 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=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response=response,
|
||||
extracted_answer=extracted,
|
||||
@@ -930,7 +749,7 @@ async def evaluate_benchmark(
|
||||
check_aime_answer(extracted, int(gold)) if extracted else False
|
||||
)
|
||||
result = QuestionResult(
|
||||
question_id=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response=response,
|
||||
extracted_answer=extracted,
|
||||
@@ -944,7 +763,7 @@ async def evaluate_benchmark(
|
||||
code = extract_code_block(response, preserve_indent=keep_indent)
|
||||
if code is None:
|
||||
result = QuestionResult(
|
||||
question_id=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response=response,
|
||||
extracted_answer=None,
|
||||
@@ -959,7 +778,7 @@ async def evaluate_benchmark(
|
||||
code,
|
||||
)
|
||||
result = QuestionResult(
|
||||
question_id=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response=response,
|
||||
extracted_answer="pass" if passed else "fail",
|
||||
@@ -974,7 +793,7 @@ async def evaluate_benchmark(
|
||||
exec_meta["sample"],
|
||||
)
|
||||
result = QuestionResult(
|
||||
question_id=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response=response,
|
||||
extracted_answer="pass" if passed else "fail",
|
||||
@@ -985,7 +804,7 @@ async def evaluate_benchmark(
|
||||
)
|
||||
else:
|
||||
result = QuestionResult(
|
||||
question_id=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response=response,
|
||||
extracted_answer=None,
|
||||
@@ -996,7 +815,7 @@ async def evaluate_benchmark(
|
||||
)
|
||||
else:
|
||||
result = QuestionResult(
|
||||
question_id=question_id,
|
||||
question_id=idx,
|
||||
prompt=prompt,
|
||||
response=response,
|
||||
extracted_answer=None,
|
||||
@@ -1008,82 +827,24 @@ 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
|
||||
|
||||
# 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)
|
||||
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%})"
|
||||
)
|
||||
|
||||
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
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1106,8 +867,6 @@ 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%})")
|
||||
@@ -1119,10 +878,6 @@ 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:
|
||||
@@ -1141,8 +896,6 @@ 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,
|
||||
}
|
||||
|
||||
|
||||
@@ -1300,11 +1053,7 @@ 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
|
||||
],
|
||||
@@ -1320,15 +1069,6 @@ 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:
|
||||
@@ -1356,12 +1096,6 @@ 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(
|
||||
@@ -1381,8 +1115,6 @@ 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."
|
||||
)
|
||||
@@ -1416,31 +1148,15 @@ 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(
|
||||
"--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",
|
||||
"--skip-instance-setup",
|
||||
action="store_true",
|
||||
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).",
|
||||
help="Skip exo instance management (assumes model is already running).",
|
||||
)
|
||||
|
||||
args, _ = ap.parse_known_args()
|
||||
@@ -1461,26 +1177,13 @@ def main() -> int:
|
||||
# Instance management
|
||||
client = ExoClient(args.host, args.port, timeout_s=args.timeout)
|
||||
instance_id: str | None = None
|
||||
created_instance = False
|
||||
|
||||
_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:
|
||||
if not args.skip_instance_setup:
|
||||
short_id, full_model_id = resolve_model_short_id(
|
||||
client,
|
||||
args.model,
|
||||
force_download=args.force_download,
|
||||
)
|
||||
selected = settle_and_fetch_placements(
|
||||
client,
|
||||
full_model_id,
|
||||
@@ -1495,7 +1198,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,
|
||||
)
|
||||
@@ -1522,18 +1225,6 @@ 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)
|
||||
@@ -1543,9 +1234,10 @@ def main() -> int:
|
||||
client.request_json("DELETE", f"/instance/{instance_id}")
|
||||
return 1
|
||||
time.sleep(1)
|
||||
created_instance = True
|
||||
|
||||
cluster_snapshot = capture_cluster_snapshot(client)
|
||||
cluster_snapshot = capture_cluster_snapshot(client)
|
||||
else:
|
||||
full_model_id = args.model
|
||||
cluster_snapshot = None
|
||||
|
||||
# Auto-detect reasoning from model config
|
||||
model_config = load_model_config(full_model_id)
|
||||
@@ -1599,57 +1291,16 @@ 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)
|
||||
@@ -1658,11 +1309,6 @@ 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,
|
||||
@@ -1673,8 +1319,9 @@ def main() -> int:
|
||||
concurrency=c,
|
||||
limit=args.limit,
|
||||
timeout=args.request_timeout,
|
||||
checkpoint_path=checkpoint_path,
|
||||
**eval_kwargs,
|
||||
reasoning_effort=reasoning_effort,
|
||||
top_p=top_p,
|
||||
difficulty=args.difficulty,
|
||||
)
|
||||
)
|
||||
if results:
|
||||
@@ -1689,18 +1336,10 @@ 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,
|
||||
@@ -1711,8 +1350,9 @@ def main() -> int:
|
||||
concurrency=args.num_concurrent,
|
||||
limit=args.limit,
|
||||
timeout=args.request_timeout,
|
||||
checkpoint_path=checkpoint_path,
|
||||
**eval_kwargs,
|
||||
reasoning_effort=reasoning_effort,
|
||||
top_p=top_p,
|
||||
difficulty=args.difficulty,
|
||||
)
|
||||
)
|
||||
if results:
|
||||
@@ -1726,25 +1366,14 @@ def main() -> int:
|
||||
scores,
|
||||
cluster=cluster_snapshot,
|
||||
)
|
||||
# Clean up checkpoint on success
|
||||
if checkpoint_path.exists():
|
||||
checkpoint_path.unlink()
|
||||
finally:
|
||||
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"
|
||||
)
|
||||
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)
|
||||
|
||||
return 0
|
||||
|
||||
|
||||
+12
-65
@@ -6,7 +6,6 @@ import http.client
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
from urllib.parse import urlencode
|
||||
|
||||
@@ -70,30 +69,6 @@ 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}")
|
||||
@@ -293,15 +268,11 @@ def sharding_filter(sharding: str, wanted: str) -> bool:
|
||||
|
||||
|
||||
def fetch_and_filter_placements(
|
||||
client: ExoClient,
|
||||
full_model_id: str,
|
||||
args: argparse.Namespace,
|
||||
node_id: str | None = None,
|
||||
client: ExoClient, full_model_id: str, args: argparse.Namespace
|
||||
) -> list[dict[str, Any]]:
|
||||
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_resp = client.request_json(
|
||||
"GET", "/instance/previews", params={"model_id": full_model_id}
|
||||
)
|
||||
previews = previews_resp.get("previews") or []
|
||||
|
||||
selected: list[dict[str, Any]] = []
|
||||
@@ -361,9 +332,8 @@ 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, node_id=node_id)
|
||||
selected = fetch_and_filter_placements(client, full_model_id, args)
|
||||
|
||||
if not selected and settle_timeout > 0:
|
||||
backoff = _SETTLE_INITIAL_BACKOFF_S
|
||||
@@ -376,9 +346,7 @@ 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, node_id=node_id
|
||||
)
|
||||
selected = fetch_and_filter_placements(client, full_model_id, args)
|
||||
|
||||
return selected
|
||||
|
||||
@@ -494,8 +462,9 @@ def run_planning_phase(
|
||||
)
|
||||
logger.info(f"Started download on {node_id}")
|
||||
|
||||
# Wait for downloads (no timeout — poll until complete or failed)
|
||||
while True:
|
||||
# Wait for downloads
|
||||
start = time.time()
|
||||
while time.time() - start < timeout:
|
||||
all_done = True
|
||||
for node_id in node_ids:
|
||||
node_downloads = client.get_node_downloads(node_id) or []
|
||||
@@ -545,24 +514,9 @@ def run_planning_phase(
|
||||
if download_t0 is not None:
|
||||
return time.perf_counter() - download_t0
|
||||
return None
|
||||
time.sleep(10)
|
||||
time.sleep(1)
|
||||
|
||||
|
||||
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
|
||||
raise TimeoutError("Downloads did not complete in time")
|
||||
|
||||
|
||||
def add_common_instance_args(ap: argparse.ArgumentParser) -> None:
|
||||
@@ -589,9 +543,7 @@ 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", "vllm", "both"],
|
||||
default="both",
|
||||
"--instance-meta", choices=["ring", "jaccl", "both"], default="both"
|
||||
)
|
||||
ap.add_argument(
|
||||
"--sharding", choices=["pipeline", "tensor", "both"], default="both"
|
||||
@@ -620,8 +572,3 @@ 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.",
|
||||
)
|
||||
@@ -1,36 +0,0 @@
|
||||
# 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 = "Ryuichi’s MacBook Pro"
|
||||
instance_meta = "ring"
|
||||
sharding = "pipeline"
|
||||
min_nodes = 1
|
||||
max_nodes = 1
|
||||
@@ -1,869 +0,0 @@
|
||||
# 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,7 +202,6 @@
|
||||
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 {
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
*/
|
||||
|
||||
interface Props {
|
||||
/** "macbook pro" | "mac studio" | "mac mini" | "dgx spark" | "linux" etc. */
|
||||
/** "macbook pro" | "mac studio" | "mac mini" etc. */
|
||||
deviceType: string;
|
||||
/** Center X coordinate in SVG space */
|
||||
cx: number;
|
||||
@@ -38,43 +38,10 @@
|
||||
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);
|
||||
@@ -147,264 +114,7 @@
|
||||
const studioClipId = $derived(`di-studio-${uid}`);
|
||||
</script>
|
||||
|
||||
{#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"}
|
||||
{#if modelLower === "mac studio" || modelLower === "mac mini"}
|
||||
<!-- Mac Studio / Mac Mini -->
|
||||
<defs>
|
||||
<clipPath id={studioClipId}>
|
||||
|
||||
@@ -1,8 +1,5 @@
|
||||
<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;
|
||||
@@ -300,28 +297,5 @@
|
||||
</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>
|
||||
@@ -23,7 +23,7 @@
|
||||
} | null;
|
||||
nodes?: Record<string, NodeInfo>;
|
||||
sharding?: "Pipeline" | "Tensor";
|
||||
runtime?: "MlxRing" | "MlxJaccl" | "Vllm";
|
||||
runtime?: "MlxRing" | "MlxJaccl";
|
||||
onLaunch?: () => void;
|
||||
tags?: string[];
|
||||
apiPreview?: PlacementPreview | null;
|
||||
@@ -168,10 +168,8 @@
|
||||
|
||||
function getDeviceType(
|
||||
name: string,
|
||||
): "macbook" | "studio" | "mini" | "dgx" | "linux" | "unknown" {
|
||||
): "macbook" | "studio" | "mini" | "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";
|
||||
@@ -578,17 +576,13 @@
|
||||
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)."
|
||||
: runtime === "MlxJaccl"
|
||||
? "RDMA: direct memory access over Thunderbolt. Significantly faster for multi-device inference."
|
||||
: "vLLM: NVIDIA CUDA inference engine."}
|
||||
: "RDMA: direct memory access over Thunderbolt. Significantly faster for multi-device inference."}
|
||||
>
|
||||
{runtime === "MlxRing"
|
||||
? "MLX Ring"
|
||||
: runtime === "MlxJaccl"
|
||||
? "MLX RDMA"
|
||||
: runtime === "Vllm"
|
||||
? "vLLM"
|
||||
: runtime}
|
||||
: runtime}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
@@ -996,81 +990,6 @@
|
||||
/>
|
||||
{/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,7 +9,6 @@
|
||||
capabilities?: string[];
|
||||
family?: string;
|
||||
is_custom?: boolean;
|
||||
requires_vllm?: boolean;
|
||||
}
|
||||
|
||||
interface ModelGroup {
|
||||
@@ -20,7 +19,6 @@
|
||||
variants: ModelInfo[];
|
||||
smallestVariant: ModelInfo;
|
||||
hasMultipleVariants: boolean;
|
||||
requiresVllm: boolean;
|
||||
}
|
||||
|
||||
type DownloadAvailability = {
|
||||
@@ -215,14 +213,6 @@
|
||||
<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"}
|
||||
@@ -533,15 +523,6 @@
|
||||
{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(
|
||||
@@ -647,7 +628,6 @@
|
||||
variants: [variant],
|
||||
smallestVariant: variant,
|
||||
hasMultipleVariants: false,
|
||||
requiresVllm: variant.requires_vllm === true,
|
||||
});
|
||||
}}
|
||||
title="View variant details"
|
||||
|
||||
@@ -22,7 +22,6 @@
|
||||
is_custom?: boolean;
|
||||
tasks?: string[];
|
||||
hugging_face_id?: string;
|
||||
requires_vllm?: boolean;
|
||||
}
|
||||
|
||||
interface ModelGroup {
|
||||
@@ -33,7 +32,6 @@
|
||||
variants: ModelInfo[];
|
||||
smallestVariant: ModelInfo;
|
||||
hasMultipleVariants: boolean;
|
||||
requiresVllm: boolean;
|
||||
}
|
||||
|
||||
interface FilterState {
|
||||
@@ -398,7 +396,6 @@
|
||||
variants: [],
|
||||
smallestVariant: model,
|
||||
hasMultipleVariants: false,
|
||||
requiresVllm: true,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -433,7 +430,6 @@
|
||||
(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)
|
||||
@@ -591,7 +587,6 @@
|
||||
variants: [model],
|
||||
smallestVariant: model,
|
||||
hasMultipleVariants: false,
|
||||
requiresVllm: model.requires_vllm === true,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -1170,17 +1165,6 @@
|
||||
<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>
|
||||
|
||||
@@ -1,565 +0,0 @@
|
||||
<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,10 +117,6 @@
|
||||
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;
|
||||
@@ -558,13 +554,6 @@
|
||||
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);
|
||||
@@ -634,382 +623,7 @@
|
||||
`${friendlyName}\nID: ${nodeInfo.id.slice(-8)}\nMemory: ${formatBytes(ramUsed)}/${formatBytes(ramTotal)}`,
|
||||
);
|
||||
|
||||
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") {
|
||||
if (modelLower === "mac studio") {
|
||||
// Mac Studio - classic cube with memory fill
|
||||
iconBaseWidth = nodeRadius * 1.25;
|
||||
iconBaseHeight = nodeRadius * 0.85;
|
||||
@@ -1568,12 +1182,8 @@
|
||||
debugLabelY += debugLineHeight;
|
||||
}
|
||||
|
||||
const dbgIdentity = identitiesData[nodeInfo.id];
|
||||
if (dbgIdentity?.osVersion) {
|
||||
const osLabel =
|
||||
dbgIdentity.osVersion === "Linux"
|
||||
? "Linux"
|
||||
: `macOS ${dbgIdentity.osVersion}${dbgIdentity.osBuildVersion ? ` (${dbgIdentity.osBuildVersion})` : ""}`;
|
||||
const identity = identitiesData[nodeInfo.id];
|
||||
if (identity?.osVersion) {
|
||||
nodeG
|
||||
.append("text")
|
||||
.attr("x", nodeInfo.x)
|
||||
@@ -1582,7 +1192,9 @@
|
||||
.attr("fill", "rgba(179,179,179,0.7)")
|
||||
.attr("font-size", debugFontSize)
|
||||
.attr("font-family", "SF Mono, Monaco, monospace")
|
||||
.text(osLabel);
|
||||
.text(
|
||||
`macOS ${identity.osVersion}${identity.osBuildVersion ? ` (${identity.osBuildVersion})` : ""}`,
|
||||
);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
@@ -74,12 +74,6 @@ 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;
|
||||
@@ -229,7 +223,6 @@ interface RawStateResponse {
|
||||
}
|
||||
>;
|
||||
runners?: Record<string, unknown>;
|
||||
instanceLinks?: Record<string, RawInstanceLink>;
|
||||
downloads?: Record<string, unknown[]>;
|
||||
// New granular node state fields
|
||||
nodeIdentities?: Record<string, RawNodeIdentity>;
|
||||
@@ -548,8 +541,6 @@ 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<
|
||||
@@ -1283,7 +1274,6 @@ class AppStore {
|
||||
|
||||
startPolling() {
|
||||
this.fetchState();
|
||||
this.fetchFeatureFlags();
|
||||
this.fetchInterval = setInterval(() => this.fetchState(), 1000);
|
||||
}
|
||||
|
||||
@@ -1295,16 +1285,6 @@ 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");
|
||||
@@ -1330,11 +1310,6 @@ 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;
|
||||
}
|
||||
@@ -1695,15 +1670,7 @@ class AppStore {
|
||||
}
|
||||
}
|
||||
}
|
||||
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;
|
||||
return { role: m.role, content: msgContent };
|
||||
}),
|
||||
];
|
||||
|
||||
@@ -1910,15 +1877,7 @@ class AppStore {
|
||||
const apiMessages = [
|
||||
systemPrompt,
|
||||
...targetConversation.messages.slice(0, -1).map((m) => {
|
||||
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;
|
||||
return { role: m.role, content: m.content };
|
||||
}),
|
||||
];
|
||||
|
||||
@@ -2449,15 +2408,10 @@ class AppStore {
|
||||
contentParts.push({ type: "text", text: textContent });
|
||||
}
|
||||
|
||||
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;
|
||||
return {
|
||||
role: m.role,
|
||||
content: contentParts,
|
||||
};
|
||||
}
|
||||
|
||||
// Text-only message (original path)
|
||||
@@ -2475,15 +2429,10 @@ class AppStore {
|
||||
}
|
||||
}
|
||||
|
||||
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;
|
||||
return {
|
||||
role: m.role,
|
||||
content: msgContent,
|
||||
};
|
||||
}),
|
||||
];
|
||||
|
||||
@@ -3332,60 +3281,6 @@ 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
|
||||
*/
|
||||
@@ -3484,19 +3379,6 @@ 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;
|
||||
|
||||
@@ -1,44 +0,0 @@
|
||||
// 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;
|
||||
}
|
||||
@@ -65,7 +65,6 @@
|
||||
nodeThunderboltBridge,
|
||||
nodeIdentities,
|
||||
isConnected,
|
||||
featureFlags,
|
||||
type DownloadProgress,
|
||||
type PlacementPreview,
|
||||
} from "$lib/stores/app.svelte";
|
||||
@@ -703,10 +702,7 @@
|
||||
? Object.keys(topologyData()!.nodes).length
|
||||
: 1;
|
||||
const sharding = nodeCount <= 1 ? "Pipeline" : selectedSharding;
|
||||
const instanceType =
|
||||
nodeCount <= 1 && selectedInstanceType === "MlxJaccl"
|
||||
? "MlxRing"
|
||||
: selectedInstanceType;
|
||||
const instanceType = nodeCount <= 1 ? "MlxRing" : selectedInstanceType;
|
||||
try {
|
||||
const placementResponse = await fetch(
|
||||
`/instance/placement?model_id=${encodeURIComponent(modelId)}&sharding=${sharding}&instance_meta=${instanceType}&min_nodes=1`,
|
||||
@@ -787,7 +783,6 @@
|
||||
quantization?: string;
|
||||
base_model?: string;
|
||||
capabilities?: string[];
|
||||
requires_vllm?: boolean;
|
||||
}>
|
||||
>([]);
|
||||
type ModelMemoryFitStatus =
|
||||
@@ -891,7 +886,7 @@
|
||||
}
|
||||
|
||||
let selectedSharding = $state<"Pipeline" | "Tensor">("Pipeline");
|
||||
type InstanceMeta = "MlxRing" | "MlxJaccl" | "Vllm";
|
||||
type InstanceMeta = "MlxRing" | "MlxJaccl";
|
||||
|
||||
// Launch defaults persistence
|
||||
const LAUNCH_DEFAULTS_KEY = "exo-launch-defaults-v2";
|
||||
@@ -937,12 +932,7 @@
|
||||
// Apply sharding and instance type unconditionally
|
||||
selectedSharding = defaults.sharding;
|
||||
selectedInstanceType =
|
||||
defaults.instanceType === "MlxRing"
|
||||
? "MlxRing"
|
||||
: defaults.instanceType === "Vllm"
|
||||
? "Vllm"
|
||||
: "MlxJaccl";
|
||||
userPickedInstanceType = true;
|
||||
defaults.instanceType === "MlxRing" ? "MlxRing" : "MlxJaccl";
|
||||
|
||||
// Apply minNodes if valid (between 1 and maxNodes)
|
||||
if (
|
||||
@@ -964,23 +954,6 @@
|
||||
}
|
||||
|
||||
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);
|
||||
@@ -1173,7 +1146,9 @@
|
||||
}
|
||||
|
||||
const matchesSelectedRuntime = (runtime: InstanceMeta): boolean =>
|
||||
runtime === selectedInstanceType;
|
||||
selectedInstanceType === "MlxRing"
|
||||
? runtime === "MlxRing"
|
||||
: runtime === "MlxJaccl";
|
||||
|
||||
// Helper to check if a model can be launched (has valid placement with >= minNodes)
|
||||
function canModelFit(modelId: string): boolean {
|
||||
@@ -2088,7 +2063,6 @@
|
||||
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?: {
|
||||
@@ -5795,18 +5769,14 @@
|
||||
</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 {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'}"
|
||||
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'}"
|
||||
>
|
||||
<span
|
||||
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
|
||||
@@ -5822,18 +5792,14 @@
|
||||
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 {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'}"
|
||||
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'}"
|
||||
>
|
||||
<span
|
||||
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
|
||||
@@ -5848,41 +5814,7 @@
|
||||
</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 -->
|
||||
|
||||
@@ -1,81 +0,0 @@
|
||||
<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>
|
||||
@@ -14,7 +14,6 @@
|
||||
|
||||
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[] = [];
|
||||
@@ -133,7 +132,6 @@
|
||||
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) };
|
||||
@@ -141,27 +139,6 @@
|
||||
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) {
|
||||
@@ -373,25 +350,16 @@
|
||||
try {
|
||||
const resp = await fetch("/v1/models");
|
||||
const data = (await resp.json()) as {
|
||||
data: {
|
||||
id: string;
|
||||
capabilities: string[];
|
||||
context_length: number;
|
||||
reasoning_dialect?: string;
|
||||
}[];
|
||||
data: { id: string; capabilities: string[]; context_length: number }[];
|
||||
};
|
||||
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 */
|
||||
}
|
||||
|
||||
@@ -146,7 +146,7 @@
|
||||
config.treefmt.build.wrapper
|
||||
|
||||
# PYTHON
|
||||
self'.packages.exo.passthru.evenv
|
||||
self'.packages.editableVenv
|
||||
uv
|
||||
|
||||
# RUST
|
||||
|
||||
@@ -40,19 +40,6 @@ 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/
|
||||
|
||||
@@ -1,26 +0,0 @@
|
||||
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
|
||||
+34
-57
@@ -15,22 +15,26 @@ 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",
|
||||
"tiktoken>=0.12.0", # required for kimi k2 tokenizer
|
||||
"mlx==0.31.2; sys_platform == 'darwin'",
|
||||
"mlx-lm; sys_platform=='darwin'",
|
||||
"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]
|
||||
@@ -45,30 +49,23 @@ dev = [
|
||||
|
||||
[project.optional-dependencies]
|
||||
build = ["nanobind"]
|
||||
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'",
|
||||
|
||||
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-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'",
|
||||
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'",
|
||||
]
|
||||
|
||||
###
|
||||
@@ -82,23 +79,12 @@ 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-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')" },
|
||||
{ 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'" },
|
||||
]
|
||||
vllm = { git = "https://github.com/hmellor/vllm.git", branch = "transformers-v5" }
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cu130"
|
||||
@@ -106,8 +92,8 @@ url = "https://download.pytorch.org/whl/cu130"
|
||||
explicit = true
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cu128"
|
||||
url = "https://download.pytorch.org/whl/cu128"
|
||||
name = "pytorch-cu120"
|
||||
url = "https://download.pytorch.org/whl/cu120"
|
||||
explicit = true
|
||||
|
||||
[[tool.uv.index]]
|
||||
@@ -168,19 +154,11 @@ root = "src"
|
||||
required-version = ">=0.8.6"
|
||||
prerelease = "allow"
|
||||
environments = ["sys_platform == 'darwin'", "sys_platform == 'linux'"]
|
||||
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" },
|
||||
],
|
||||
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'",
|
||||
]
|
||||
|
||||
[tool.uv.extra-build-dependencies]
|
||||
@@ -195,7 +173,6 @@ mlx = [
|
||||
"ninja",
|
||||
]
|
||||
mlx-lm = ["setuptools"]
|
||||
mflux = ["uv_build"]
|
||||
xgrammar = [
|
||||
"nanobind",
|
||||
"setuptools",
|
||||
|
||||
+20
-202
@@ -10,10 +10,8 @@ let
|
||||
inherit (pkgs.stdenv.hostPlatform) isLinux isDarwin isx86_64;
|
||||
inherit (pkgs.config) cudaSupport;
|
||||
inherit (pkgs) cudaPackages;
|
||||
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";
|
||||
cuda13Support = cudaSupport && cudaPackages.cudaMajorVersion == "13";
|
||||
libmlx_source = if cuda13Support then "mlx-cuda-13" else if cudaSupport then "mlx-cuda-12" else "mlx-cpu";
|
||||
python = pkgs.python313;
|
||||
cudaLibs = with cudaPackages; [
|
||||
cuda_cudart
|
||||
@@ -115,213 +113,37 @@ 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: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
buildInputs = old.buildInputs ++ [ cudaLibs ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
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";
|
||||
@@ -342,28 +164,24 @@ let
|
||||
buildSystemsOverlay
|
||||
]
|
||||
);
|
||||
# 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 {
|
||||
venv = name: (pythonSet.mkVirtualEnv "${name}-env" members).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" ]; });
|
||||
mkApp = cmd: 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 ];
|
||||
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" ]; });
|
||||
};
|
||||
text = "exec " + lib.optionalString cudaSupport "${lib.getExe pkgs.nix-gl-host} " + cmd;
|
||||
};
|
||||
in
|
||||
{
|
||||
inherit venv;
|
||||
editablePythonSet = pythonSet.overrideScope editableOverlay;
|
||||
mkPythonScript = path: mkApp ''python ${path} "$@"'';
|
||||
mkExo = mkApp ''exo "$@"'';
|
||||
};
|
||||
@@ -373,18 +191,18 @@ in
|
||||
{ self', pkgs, unfreePkgs, lib, ... }:
|
||||
let
|
||||
inherit (pkgs.stdenv.hostPlatform) isLinux;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "mlx-cpu" "vllm-none" ]; }; }) mkExo;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "cpu" ]; }; }) editablePythonSet mkExo;
|
||||
|
||||
# Virtual environment with dev dependencies for testing
|
||||
testVenv = (mkPythonSet {
|
||||
inherit self' pkgs lib; members = {
|
||||
exo = [ "dev" "mlx-cpu" "vllm-none" ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
exo = [ "dev" "cpu" ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
};
|
||||
}).venv "exo-test";
|
||||
|
||||
mkBenchScript = (mkPythonSet {
|
||||
inherit self' pkgs lib; members = {
|
||||
exo = [ "mlx-cpu" "vllm-none" ];
|
||||
exo = [ "cpu" ];
|
||||
exo-bench = [ ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
};
|
||||
}).mkPythonScript;
|
||||
@@ -394,12 +212,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);
|
||||
@@ -408,8 +226,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 = cuda12Set.mkExo "exo-cuda-12";
|
||||
exo-cuda-13 = cuda13Set.mkExo "exo-cuda-13";
|
||||
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";
|
||||
};
|
||||
|
||||
checks = {
|
||||
|
||||
@@ -1,21 +0,0 @@
|
||||
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,7 +8,6 @@ family = "deepseek"
|
||||
quantization = "4bit"
|
||||
base_model = "DeepSeek V3.1"
|
||||
capabilities = ["text", "thinking", "thinking_toggle"]
|
||||
reasoning_dialect = "post_last_user"
|
||||
|
||||
context_length = 131072
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ family = "deepseek"
|
||||
quantization = "8bit"
|
||||
base_model = "DeepSeek V3.1"
|
||||
capabilities = ["text", "thinking", "thinking_toggle"]
|
||||
reasoning_dialect = "post_last_user"
|
||||
|
||||
context_length = 131072
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ family = "deepseek"
|
||||
quantization = "4bit"
|
||||
base_model = "DeepSeek V3.2"
|
||||
capabilities = ["text", "thinking", "thinking_toggle"]
|
||||
reasoning_dialect = "tool_conditional"
|
||||
|
||||
context_length = 131072
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ family = "deepseek"
|
||||
quantization = "8bit"
|
||||
base_model = "DeepSeek V3.2"
|
||||
capabilities = ["text", "thinking", "thinking_toggle"]
|
||||
reasoning_dialect = "tool_conditional"
|
||||
|
||||
context_length = 131072
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ family = "deepseek"
|
||||
quantization = "8bit"
|
||||
base_model = "DeepSeek V4 Flash"
|
||||
capabilities = ["text", "thinking", "thinking_toggle"]
|
||||
reasoning_dialect = "tool_conditional"
|
||||
|
||||
context_length = 1048576
|
||||
|
||||
|
||||
@@ -8,7 +8,6 @@ family = "deepseek"
|
||||
quantization = "8bit"
|
||||
base_model = "DeepSeek V4 Pro"
|
||||
capabilities = ["text", "thinking", "thinking_toggle"]
|
||||
reasoning_dialect = "tool_conditional"
|
||||
|
||||
context_length = 1048576
|
||||
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
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
|
||||
@@ -1,20 +0,0 @@
|
||||
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
|
||||
@@ -1,32 +0,0 @@
|
||||
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
|
||||
@@ -1,222 +0,0 @@
|
||||
#!/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)
|
||||
@@ -1,124 +0,0 @@
|
||||
#!/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
|
||||
@@ -1,170 +0,0 @@
|
||||
#!/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
|
||||
@@ -131,13 +131,9 @@ async def chat_request_to_text_generation(
|
||||
multimodal_content.append({"type": "text", "text": part.text})
|
||||
else:
|
||||
multimodal_content.append({"type": "image"})
|
||||
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)
|
||||
chat_template_messages.append(
|
||||
{"role": msg.role, "content": multimodal_content}
|
||||
)
|
||||
continue
|
||||
msg_copy = msg.model_copy(update={"content": content})
|
||||
|
||||
@@ -172,8 +168,6 @@ 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,
|
||||
)
|
||||
|
||||
|
||||
+4
-91
@@ -79,8 +79,6 @@ from exo.api.types import (
|
||||
ImageListItem,
|
||||
ImageListResponse,
|
||||
ImageSize,
|
||||
InstanceLinkBody,
|
||||
InstanceLinkResponse,
|
||||
ModelList,
|
||||
ModelListModel,
|
||||
PlaceInstanceParams,
|
||||
@@ -124,7 +122,6 @@ 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,
|
||||
@@ -157,7 +154,6 @@ from exo.shared.types.commands import (
|
||||
DeleteCustomModelCard,
|
||||
DeleteDownload,
|
||||
DeleteInstance,
|
||||
DeleteInstanceLink,
|
||||
DownloadCommand,
|
||||
ForwarderCommand,
|
||||
ForwarderDownloadCommand,
|
||||
@@ -165,7 +161,6 @@ from exo.shared.types.commands import (
|
||||
ImageGeneration,
|
||||
PlaceInstance,
|
||||
SendInputChunk,
|
||||
SetInstanceLink,
|
||||
StartDownload,
|
||||
TaskCancelled,
|
||||
TaskFinished,
|
||||
@@ -179,7 +174,6 @@ 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 (
|
||||
@@ -221,17 +215,6 @@ 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,
|
||||
@@ -345,11 +328,6 @@ 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)
|
||||
@@ -358,9 +336,7 @@ class API:
|
||||
self.app.post("/v1/chat/completions", response_model=None)(
|
||||
self.chat_completions
|
||||
)
|
||||
self.app.post("/bench/chat/completions", response_model=None)(
|
||||
self.bench_chat_completions
|
||||
)
|
||||
self.app.post("/bench/chat/completions")(self.bench_chat_completions)
|
||||
self.app.post("/v1/images/generations", response_model=None)(
|
||||
self.image_generations
|
||||
)
|
||||
@@ -526,8 +502,8 @@ class API:
|
||||
)
|
||||
]
|
||||
)
|
||||
if any(self.state.node_vllm.values()):
|
||||
instance_combinations.append((Sharding.Pipeline, InstanceMeta.Vllm, 1))
|
||||
# TODO: PDD
|
||||
# instance_combinations.append((Sharding.PrefillDecodeDisaggregation, InstanceMeta.MlxRing, 1))
|
||||
|
||||
for sharding, instance_meta, min_nodes in instance_combinations:
|
||||
try:
|
||||
@@ -639,52 +615,6 @@ 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(
|
||||
@@ -899,7 +829,7 @@ class API:
|
||||
|
||||
async def bench_chat_completions(
|
||||
self, payload: BenchChatCompletionRequest
|
||||
) -> BenchChatCompletionResponse | StreamingResponse:
|
||||
) -> BenchChatCompletionResponse:
|
||||
task_params = await chat_request_to_text_generation(payload)
|
||||
resolved_model = await self._resolve_and_validate_text_model(
|
||||
ModelId(task_params.model)
|
||||
@@ -916,22 +846,6 @@ 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:
|
||||
@@ -1751,7 +1665,6 @@ class API:
|
||||
capabilities=card.capabilities,
|
||||
reasoning_dialect=card.reasoning_dialect,
|
||||
context_length=card.context_length,
|
||||
requires_vllm=card.requires_vllm,
|
||||
)
|
||||
for card in cards
|
||||
]
|
||||
|
||||
@@ -34,8 +34,6 @@ 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
|
||||
|
||||
@@ -49,7 +49,6 @@ 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):
|
||||
@@ -297,16 +296,6 @@ 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",
|
||||
|
||||
+5
-107
@@ -10,7 +10,6 @@ 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 (
|
||||
@@ -18,7 +17,6 @@ from exo.shared.types.commands import (
|
||||
CreateInstance,
|
||||
DeleteCustomModelCard,
|
||||
DeleteInstance,
|
||||
DeleteInstanceLink,
|
||||
ForwarderCommand,
|
||||
ForwarderDownloadCommand,
|
||||
ImageEdits,
|
||||
@@ -26,7 +24,6 @@ from exo.shared.types.commands import (
|
||||
PlaceInstance,
|
||||
RequestEventLog,
|
||||
SendInputChunk,
|
||||
SetInstanceLink,
|
||||
TaskCancelled,
|
||||
TaskFinished,
|
||||
TestCommand,
|
||||
@@ -41,8 +38,6 @@ from exo.shared.types.events import (
|
||||
IndexedEvent,
|
||||
InputChunkReceived,
|
||||
InstanceDeleted,
|
||||
InstanceLinkCreated,
|
||||
InstanceLinkDeleted,
|
||||
LocalForwarderEvent,
|
||||
NodeGatheredInfo,
|
||||
NodeTimedOut,
|
||||
@@ -53,7 +48,6 @@ 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,
|
||||
@@ -75,46 +69,6 @@ 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,
|
||||
@@ -174,45 +128,15 @@ class Master:
|
||||
case TestCommand():
|
||||
pass
|
||||
case TextGeneration():
|
||||
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
|
||||
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
|
||||
@@ -230,27 +154,20 @@ 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=decode_instance_id,
|
||||
instance_id=available_instance_ids[0],
|
||||
task_status=TaskStatus.Pending,
|
||||
task_params=params,
|
||||
task_params=command.task_params,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
self.command_task_mapping[command.command_id] = task_id
|
||||
case ImageGeneration():
|
||||
for instance in self.state.instances.values():
|
||||
@@ -258,12 +175,10 @@ 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
|
||||
@@ -314,12 +229,10 @@ 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
|
||||
@@ -444,21 +357,6 @@ 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
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import random
|
||||
from collections.abc import Mapping
|
||||
from copy import deepcopy
|
||||
from typing import Sequence
|
||||
@@ -43,10 +44,13 @@ from exo.shared.types.worker.instances import (
|
||||
InstanceMeta,
|
||||
MlxJacclInstance,
|
||||
MlxRingInstance,
|
||||
VllmInstance,
|
||||
)
|
||||
from exo.shared.types.worker.shards import Sharding
|
||||
from exo.utils.ports import random_ephemeral_port
|
||||
|
||||
|
||||
def random_ephemeral_port() -> int:
|
||||
port = random.randint(49153, 65535)
|
||||
return port - 1 if port <= 52415 else port
|
||||
|
||||
|
||||
def add_instance_to_placements(
|
||||
@@ -203,7 +207,7 @@ def place_instance(
|
||||
)
|
||||
|
||||
# Single-node: force Pipeline/Ring (Tensor and Jaccl require multi-node)
|
||||
if len(selected_cycle) == 1 and command.instance_meta != InstanceMeta.Vllm:
|
||||
if len(selected_cycle) == 1:
|
||||
command = command.model_copy(
|
||||
update={
|
||||
"instance_meta": InstanceMeta.MlxRing,
|
||||
@@ -267,11 +271,6 @@ 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
|
||||
|
||||
|
||||
@@ -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,13 +375,7 @@ def find_ip_prioritised(
|
||||
"maybe_ethernet": 3,
|
||||
"thunderbolt": 4,
|
||||
}
|
||||
|
||||
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)
|
||||
return min(ips, key=lambda ip: priority.get(ip_to_type.get(ip, "unknown"), 2))
|
||||
|
||||
|
||||
def get_mlx_ring_hosts_by_node(
|
||||
@@ -419,7 +413,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:
|
||||
@@ -451,7 +445,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:
|
||||
|
||||
File renamed without changes.
@@ -0,0 +1,33 @@
|
||||
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))
|
||||
@@ -0,0 +1,261 @@
|
||||
"""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)
|
||||
@@ -0,0 +1,21 @@
|
||||
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)
|
||||
@@ -0,0 +1,170 @@
|
||||
"""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]
|
||||
@@ -0,0 +1,75 @@
|
||||
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()
|
||||
@@ -0,0 +1,70 @@
|
||||
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()
|
||||
@@ -0,0 +1,75 @@
|
||||
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"
|
||||
@@ -0,0 +1,384 @@
|
||||
"""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"}
|
||||
},
|
||||
)
|
||||
+4
-79
@@ -14,8 +14,6 @@ from exo.shared.types.events import (
|
||||
InputChunkReceived,
|
||||
InstanceCreated,
|
||||
InstanceDeleted,
|
||||
InstanceLinkCreated,
|
||||
InstanceLinkDeleted,
|
||||
NodeDownloadProgress,
|
||||
NodeGatheredInfo,
|
||||
NodeTimedOut,
|
||||
@@ -31,7 +29,6 @@ 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,
|
||||
@@ -44,12 +41,7 @@ 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,
|
||||
RunnerReady,
|
||||
RunnerShutdown,
|
||||
RunnerStatus,
|
||||
)
|
||||
from exo.shared.types.worker.runners import RunnerId, RunnerShutdown, RunnerStatus
|
||||
from exo.utils.info_gatherer.info_gatherer import (
|
||||
MacmonMetrics,
|
||||
MacThunderboltConnections,
|
||||
@@ -59,11 +51,9 @@ from exo.utils.info_gatherer.info_gatherer import (
|
||||
NodeConfig,
|
||||
NodeDiskUsage,
|
||||
NodeNetworkInterfaces,
|
||||
NvmlMetrics,
|
||||
RdmaCtlStatus,
|
||||
StaticNodeInformation,
|
||||
ThunderboltBridgeInfo,
|
||||
VllmCapability,
|
||||
)
|
||||
|
||||
|
||||
@@ -105,10 +95,6 @@ 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:
|
||||
@@ -208,38 +194,7 @@ 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
|
||||
}
|
||||
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})
|
||||
return state.model_copy(update={"instances": new_instances})
|
||||
|
||||
|
||||
def apply_runner_status_updated(event: RunnerStatusUpdated, state: State) -> State:
|
||||
@@ -247,28 +202,12 @@ 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
|
||||
}
|
||||
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}
|
||||
)
|
||||
return state.model_copy(update={"runners": new_runners})
|
||||
new_runners = {
|
||||
**state.runners,
|
||||
event.runner_id: event.runner_status,
|
||||
}
|
||||
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)
|
||||
return state.model_copy(update={"runners": new_runners})
|
||||
|
||||
|
||||
def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
|
||||
@@ -306,9 +245,6 @@ 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 = (
|
||||
@@ -331,7 +267,6 @@ 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,
|
||||
}
|
||||
)
|
||||
@@ -358,11 +293,6 @@ 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():
|
||||
@@ -443,11 +373,6 @@ 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)
|
||||
|
||||
|
||||
@@ -96,8 +96,6 @@ 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
|
||||
@@ -150,7 +150,6 @@ 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)
|
||||
@@ -350,11 +349,7 @@ 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.
|
||||
|
||||
Single-shard repos don't have a `model.safetensors.index.json`; fall back
|
||||
to the HF API for those.
|
||||
"""
|
||||
"""Gets model size from safetensors index or falls back to HF API."""
|
||||
from exo.download.download_utils import (
|
||||
download_file_with_retry,
|
||||
resolve_model_dir,
|
||||
@@ -362,25 +357,21 @@ 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)
|
||||
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
|
||||
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())
|
||||
|
||||
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)
|
||||
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:
|
||||
|
||||
@@ -1,72 +0,0 @@
|
||||
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
|
||||
@@ -85,6 +85,6 @@ class PrefillProgressChunk(BaseChunk):
|
||||
total_tokens: int
|
||||
|
||||
|
||||
StatusChunk = PrefillProgressChunk
|
||||
GenerationChunk = TokenChunk | ImageChunk | ToolCallChunk | ErrorChunk
|
||||
Chunk = StatusChunk | GenerationChunk
|
||||
GenerationChunk = (
|
||||
TokenChunk | ImageChunk | ToolCallChunk | ErrorChunk | PrefillProgressChunk
|
||||
)
|
||||
@@ -7,7 +7,6 @@ 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
|
||||
@@ -90,16 +89,6 @@ 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
|
||||
|
||||
|
||||
@@ -117,8 +106,6 @@ Command = (
|
||||
| SendInputChunk
|
||||
| AddCustomModelCard
|
||||
| DeleteCustomModelCard
|
||||
| SetInstanceLink
|
||||
| DeleteInstanceLink
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -5,9 +5,8 @@ from pydantic import Field
|
||||
|
||||
from exo.shared.models.model_cards import ModelCard
|
||||
from exo.shared.topology import Connection
|
||||
from exo.shared.types.chunks import Chunk, InputImageChunk
|
||||
from exo.shared.types.chunks import GenerationChunk, 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
|
||||
@@ -92,7 +91,7 @@ class NodeDownloadProgress(BaseEvent):
|
||||
|
||||
class ChunkGenerated(BaseEvent):
|
||||
command_id: CommandId
|
||||
chunk: Chunk
|
||||
chunk: GenerationChunk
|
||||
|
||||
|
||||
class InputChunkReceived(BaseEvent):
|
||||
@@ -138,14 +137,6 @@ class TracesMerged(BaseEvent):
|
||||
traces: list[TraceEventData]
|
||||
|
||||
|
||||
class InstanceLinkCreated(BaseEvent):
|
||||
link: InstanceLink
|
||||
|
||||
|
||||
class InstanceLinkDeleted(BaseEvent):
|
||||
link_id: InstanceLinkId
|
||||
|
||||
|
||||
Event = (
|
||||
TestEvent
|
||||
| TaskCreated
|
||||
@@ -167,8 +158,6 @@ Event = (
|
||||
| TracesMerged
|
||||
| CustomModelCardAdded
|
||||
| CustomModelCardDeleted
|
||||
| InstanceLinkCreated
|
||||
| InstanceLinkDeleted
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -1,13 +0,0 @@
|
||||
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]
|
||||
@@ -11,16 +11,10 @@ from mlx_lm.models.cache import (
|
||||
QuantizedKVCache,
|
||||
RotatingKVCache,
|
||||
)
|
||||
from mlx_lm.models.deepseek_v4 import DeepseekV4Cache
|
||||
|
||||
# This list contains one cache entry per transformer layer
|
||||
KVCacheType = Sequence[
|
||||
KVCache
|
||||
| RotatingKVCache
|
||||
| QuantizedKVCache
|
||||
| ArraysCache
|
||||
| CacheList
|
||||
| DeepseekV4Cache
|
||||
KVCache | RotatingKVCache | QuantizedKVCache | ArraysCache | CacheList
|
||||
]
|
||||
|
||||
|
||||
@@ -7,7 +7,6 @@ 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,
|
||||
@@ -58,14 +57,10 @@ 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()
|
||||
|
||||
@@ -101,6 +101,3 @@ Task = (
|
||||
| ImageEdits
|
||||
| Shutdown
|
||||
)
|
||||
TextTask = TextGeneration
|
||||
ImageTask = ImageGeneration | ImageEdits
|
||||
GenerationTask = TextTask | ImageTask
|
||||
@@ -13,20 +13,6 @@ 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"
|
||||
]
|
||||
@@ -132,8 +118,6 @@ 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
|
||||
|
||||
|
||||
@@ -15,7 +15,6 @@ class InstanceId(Id):
|
||||
class InstanceMeta(str, Enum):
|
||||
MlxRing = "MlxRing"
|
||||
MlxJaccl = "MlxJaccl"
|
||||
Vllm = "Vllm"
|
||||
|
||||
|
||||
class BaseInstance(TaggedModel):
|
||||
@@ -36,12 +35,8 @@ class MlxJacclInstance(BaseInstance):
|
||||
jaccl_coordinators: dict[NodeId, str]
|
||||
|
||||
|
||||
class VllmInstance(BaseInstance):
|
||||
pass
|
||||
|
||||
|
||||
# TODO: Single node instance
|
||||
Instance = MlxRingInstance | MlxJacclInstance | VllmInstance
|
||||
Instance = MlxRingInstance | MlxJacclInstance
|
||||
|
||||
|
||||
class BoundInstance(FrozenModel):
|
||||
|
||||
@@ -16,6 +16,10 @@ class BaseRunnerResponse(TaggedModel):
|
||||
pass
|
||||
|
||||
|
||||
class TokenizedResponse(BaseRunnerResponse):
|
||||
prompt_tokens: int
|
||||
|
||||
|
||||
class GenerationResponse(BaseRunnerResponse):
|
||||
text: str
|
||||
token: int
|
||||
@@ -71,10 +75,6 @@ class ModelLoadingResponse(BaseRunnerResponse):
|
||||
total: int
|
||||
|
||||
|
||||
class CancelledResponse(BaseRunnerResponse):
|
||||
pass
|
||||
|
||||
|
||||
class PrefillProgressResponse(BaseRunnerResponse):
|
||||
processed_tokens: int
|
||||
total_tokens: int
|
||||
@@ -47,7 +47,7 @@ class RunnerWarmingUp(BaseRunnerStatus):
|
||||
|
||||
|
||||
class RunnerReady(BaseRunnerStatus):
|
||||
prefill_server_port: int | None = None
|
||||
pass
|
||||
|
||||
|
||||
class RunnerRunning(BaseRunnerStatus):
|
||||
|
||||
@@ -25,12 +25,12 @@ def print_startup_banner(port: int) -> None:
|
||||
banner = f"""
|
||||
╔═══════════════════════════════════════════════════════════════════════╗
|
||||
║ ║
|
||||
║ ███████╗██╗ ██╗ ██████╗ ██████╗ ██████╗ ██╗ ██╗ ║
|
||||
║ ██╔════╝╚██╗██╔╝██╔═══██╗ ██ ██╔══██╗██╔════╝ ╚██╗██╔╝ ║
|
||||
║ █████╗ ╚███╔╝ ██║ ██║ ██████╗ ██║ ██║██║ ███╗ ╚███╔╝ ║
|
||||
║ ██╔══╝ ██╔██╗ ██║ ██║ ╚═██╔═╝ ██║ ██║██║ ██║ ██╔██╗ ║
|
||||
║ ███████╗██╔╝ ██╗╚██████╔╝ ╚═╝ ██████╔╝╚██████╔╝██╔╝ ██╗ ║
|
||||
║ ╚══════╝╚═╝ ╚═╝ ╚═════╝ ╚═════╝ ╚═════╝ ╚═╝ ╚═╝ ║
|
||||
║ ███████╗██╗ ██╗ ██████╗ ║
|
||||
║ ██╔════╝╚██╗██╔╝██╔═══██╗ ║
|
||||
║ █████╗ ╚███╔╝ ██║ ██║ ║
|
||||
║ ██╔══╝ ██╔██╗ ██║ ██║ ║
|
||||
║ ███████╗██╔╝ ██╗╚██████╔╝ ║
|
||||
║ ╚══════╝╚═╝ ╚═╝ ╚═════╝ ║
|
||||
║ ║
|
||||
║ Distributed AI Inference Cluster ║
|
||||
║ ║
|
||||
|
||||
@@ -31,7 +31,6 @@ 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,
|
||||
@@ -354,24 +353,6 @@ 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
|
||||
@@ -380,8 +361,6 @@ GatheredInfo = (
|
||||
| MacThunderboltConnections
|
||||
| RdmaCtlStatus
|
||||
| ThunderboltBridgeInfo
|
||||
| NvmlMetrics
|
||||
| VllmCapability
|
||||
| NodeConfig
|
||||
| MiscData
|
||||
| StaticNodeInformation
|
||||
@@ -440,8 +419,6 @@ 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)
|
||||
@@ -450,10 +427,6 @@ 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()
|
||||
@@ -502,16 +475,6 @@ 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
|
||||
|
||||
@@ -1,70 +0,0 @@
|
||||
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()
|
||||
@@ -1,7 +1,6 @@
|
||||
import platform
|
||||
import socket
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from subprocess import CalledProcessError
|
||||
|
||||
import psutil
|
||||
@@ -118,90 +117,12 @@ 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 system model and chip information.
|
||||
|
||||
On macOS, uses ``system_profiler``. On Linux, reads DMI data from
|
||||
sysfs and CPU info from ``/proc/cpuinfo``.
|
||||
"""
|
||||
"""Get Mac system information using system_profiler."""
|
||||
model = "Unknown Model"
|
||||
chip = "Unknown Chip"
|
||||
|
||||
if sys.platform == "linux":
|
||||
return await _get_linux_model_and_chip()
|
||||
|
||||
# TODO: better non mac support
|
||||
if sys.platform != "darwin":
|
||||
return (model, chip)
|
||||
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
import random
|
||||
|
||||
|
||||
def random_ephemeral_port() -> int:
|
||||
port = random.randint(49153, 65535)
|
||||
return port - 1 if port <= 52415 else port
|
||||
Whitespace-only changes.
@@ -1,204 +0,0 @@
|
||||
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))
|
||||
@@ -1,109 +0,0 @@
|
||||
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
|
||||
@@ -1,60 +0,0 @@
|
||||
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: ...
|
||||
@@ -1,16 +1,12 @@
|
||||
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",
|
||||
]
|
||||
@@ -1,219 +0,0 @@
|
||||
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
|
||||
from typing import Any, Literal, Optional
|
||||
|
||||
import mlx.core as mx
|
||||
from mflux.models.common.config.config import Config
|
||||
@@ -9,11 +9,8 @@ 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.shards import (
|
||||
CfgShardMetadata,
|
||||
PipelineShardMetadata,
|
||||
ShardMetadata,
|
||||
)
|
||||
from exo.shared.types.worker.instances import BoundInstance
|
||||
from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata
|
||||
from exo.worker.engines.image.config import ImageModelConfig
|
||||
from exo.worker.engines.image.models import (
|
||||
create_adapter_for_model,
|
||||
@@ -21,7 +18,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 mx_barrier
|
||||
from exo.worker.engines.mlx.utils_mlx import mlx_distributed_init, mx_barrier
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
|
||||
|
||||
@@ -36,7 +33,7 @@ class DistributedImageModel:
|
||||
model_id: ModelId,
|
||||
local_path: Path,
|
||||
shard_metadata: PipelineShardMetadata | CfgShardMetadata,
|
||||
group: mx.distributed.Group | None,
|
||||
group: Optional[mx.distributed.Group] = None,
|
||||
quantize: int | None = None,
|
||||
):
|
||||
config = get_config_for_model(model_id)
|
||||
@@ -79,21 +76,32 @@ class DistributedImageModel:
|
||||
self._runner = runner
|
||||
|
||||
@classmethod
|
||||
def from_shard_metadata(
|
||||
cls, shard: ShardMetadata, group: mx.distributed.Group | None
|
||||
def from_bound_instance(
|
||||
cls, bound_instance: BoundInstance
|
||||
) -> "DistributedImageModel":
|
||||
model_id = shard.model_card.model_id
|
||||
model_id = bound_instance.bound_shard.model_card.model_id
|
||||
model_path = build_model_path(model_id)
|
||||
|
||||
if not isinstance(shard, (PipelineShardMetadata, CfgShardMetadata)):
|
||||
shard_metadata = bound_instance.bound_shard
|
||||
if not isinstance(shard_metadata, (PipelineShardMetadata, CfgShardMetadata)):
|
||||
raise ValueError(
|
||||
"Expected PipelineShardMetadata or CfgShardMetadata for image generation"
|
||||
)
|
||||
|
||||
is_distributed = (
|
||||
len(bound_instance.instance.shard_assignments.node_to_runner) > 1
|
||||
)
|
||||
|
||||
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,
|
||||
shard_metadata=shard_metadata,
|
||||
group=group,
|
||||
)
|
||||
|
||||
@@ -168,3 +176,7 @@ class DistributedImageModel:
|
||||
else:
|
||||
logger.info("generated image")
|
||||
yield result
|
||||
|
||||
|
||||
def initialize_image_model(bound_instance: BoundInstance) -> DistributedImageModel:
|
||||
return DistributedImageModel.from_bound_instance(bound_instance)
|
||||
@@ -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
|
||||
layer.attn.wo_b = _AllSumLinear(layer.attn.wo_b, self.group) # type: ignore[assignment]
|
||||
|
||||
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
|
||||
layer.ffn = wrapped # type: ignore[assignment]
|
||||
|
||||
mx.eval(layer)
|
||||
mx.clear_cache()
|
||||
|
||||
@@ -1,113 +0,0 @@
|
||||
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,
|
||||
)
|
||||
@@ -13,17 +13,11 @@ 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:
|
||||
@@ -53,9 +47,7 @@ class CacheSnapshot:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
states: list[
|
||||
RotatingKVCache | ArraysCache | CacheList | DeepseekV4Cache | None
|
||||
],
|
||||
states: list[RotatingKVCache | ArraysCache | CacheList | None],
|
||||
token_count: int,
|
||||
):
|
||||
self.states = states
|
||||
@@ -120,71 +112,21 @@ def _copy_cache_list(cl: CacheList) -> CacheList:
|
||||
return CacheList(*copied)
|
||||
|
||||
|
||||
def _detached_copy_or_none(a: mx.array | None) -> mx.array | None:
|
||||
if a is None:
|
||||
def restore_snapshot_entry(
|
||||
entry: ArraysCache | RotatingKVCache | CacheList | None,
|
||||
) -> ArraysCache | RotatingKVCache | CacheList | None:
|
||||
if entry is None:
|
||||
return None
|
||||
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)
|
||||
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)
|
||||
|
||||
|
||||
def snapshot_ssm_states(cache: KVCacheType) -> CacheSnapshot:
|
||||
states: list[
|
||||
RotatingKVCache | ArraysCache | CacheList | DeepseekV4Cache | None
|
||||
] = []
|
||||
states: list[ArraysCache | RotatingKVCache | CacheList | None] = []
|
||||
for c in cache:
|
||||
if isinstance(c, ArraysCache):
|
||||
states.append(_copy_arrays_cache(c))
|
||||
@@ -192,8 +134,6 @@ 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)
|
||||
@@ -213,20 +153,14 @@ 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."""
|
||||
return any(is_non_trimmable_cache_entry(c) for c in cache)
|
||||
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
|
||||
|
||||
|
||||
class KVPrefixCache:
|
||||
@@ -382,10 +316,6 @@ 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
|
||||
|
||||
@@ -481,22 +411,16 @@ def trim_cache(
|
||||
)
|
||||
if non_trimmable:
|
||||
if snapshot is not None and snapshot.states[i] is not None:
|
||||
restored = copy_snapshot_entry(snapshot.states[i])
|
||||
restored = restore_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)
|
||||
|
||||
@@ -514,12 +438,7 @@ def encode_prompt(tokenizer: TokenizerWrapper, prompt: str) -> mx.array:
|
||||
|
||||
|
||||
def _entry_length(
|
||||
c: KVCache
|
||||
| RotatingKVCache
|
||||
| QuantizedKVCache
|
||||
| ArraysCache
|
||||
| CacheList
|
||||
| DeepseekV4Cache,
|
||||
c: KVCache | RotatingKVCache | QuantizedKVCache | ArraysCache | CacheList,
|
||||
) -> int:
|
||||
# Use .offset attribute which KVCache types have (len() not implemented in older QuantizedKVCache).
|
||||
if hasattr(c, "offset"):
|
||||
|
||||
+1
@@ -602,6 +602,7 @@ def encode_messages(
|
||||
|
||||
prompt = bos_token if add_default_bos_token and len(context) == 0 else ""
|
||||
|
||||
# Resolve drop_thinking: if any message has tools defined, don't drop thinking
|
||||
effective_drop_thinking = drop_thinking
|
||||
if any(m.get("tools") for m in full_messages):
|
||||
effective_drop_thinking = False
|
||||
Whitespace-only changes.
@@ -1,272 +0,0 @@
|
||||
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
|
||||
@@ -1,155 +0,0 @@
|
||||
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
|
||||
@@ -1,86 +0,0 @@
|
||||
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.
@@ -1,167 +0,0 @@
|
||||
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()
|
||||
@@ -1,270 +0,0 @@
|
||||
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
|
||||
@@ -1,154 +0,0 @@
|
||||
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)
|
||||
@@ -1,120 +0,0 @@
|
||||
"""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
-7
@@ -34,17 +34,11 @@ def encode_messages(
|
||||
add_default_bos_token: bool = True,
|
||||
tools: Any = None, # pyright: ignore[reportAny]
|
||||
) -> str:
|
||||
# V3.2 (like V4) is `tool_conditional`: when tools are in play, prior-turn
|
||||
# reasoning_content must be retained so multi-step tool chains stay
|
||||
# coherent.
|
||||
effective_drop_thinking = drop_thinking
|
||||
if tools:
|
||||
effective_drop_thinking = False
|
||||
prompt: str = deepseek_v32.encode_messages(
|
||||
messages,
|
||||
thinking_mode=thinking_mode,
|
||||
context=context,
|
||||
drop_thinking=effective_drop_thinking,
|
||||
drop_thinking=drop_thinking,
|
||||
add_default_bos_token=add_default_bos_token,
|
||||
tools=tools,
|
||||
)
|
||||
Loaded 100 of 137 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user