mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-08 11:35:40 -04:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
84db569167 | ||
|
|
b83d5e6a6f | ||
|
|
9e67e89862 | ||
|
|
92ea4ed0a4 | ||
|
|
9f37340b89 | ||
|
|
344381fd74 | ||
|
|
701c9b1cf6 | ||
|
|
dc709e933a | ||
|
|
0a736d7eaf | ||
|
|
94b1813f76 | ||
|
|
8774513367 | ||
|
|
35e3335d6d | ||
|
|
c2b35f4d9e | ||
|
|
d96f8379ce | ||
|
|
c1eca8d026 | ||
|
|
dbc736c845 |
No files matched your search
@@ -1,8 +1 @@
|
||||
use flake
|
||||
|
||||
# creates .venv if doesn't exist and loads its environment
|
||||
export VIRTUAL_ENV=".venv"
|
||||
if ! [ -d "./$VIRTUAL_ENV" ]; then
|
||||
uv venv
|
||||
fi
|
||||
layout python
|
||||
@@ -191,13 +191,10 @@ class RotatingKVCache(_BaseCache):
|
||||
def state(self, v): # -> None:
|
||||
...
|
||||
@property
|
||||
def meta_state(self): # -> tuple[str, ...]:
|
||||
...
|
||||
def meta_state(self) -> tuple[str, ...]: ...
|
||||
@meta_state.setter
|
||||
def meta_state(self, v): # -> None:
|
||||
...
|
||||
def is_trimmable(self): # -> bool:
|
||||
...
|
||||
def meta_state(self, v: tuple[str, ...]) -> None: ...
|
||||
def is_trimmable(self) -> bool: ...
|
||||
def trim(self, n: int) -> int: ...
|
||||
def to_quantized(
|
||||
self, group_size: int = ..., bits: int = ...
|
||||
|
||||
File diff suppressed because it is too large.
Load diff
+3
-1
@@ -589,7 +589,9 @@ def add_common_instance_args(ap: argparse.ArgumentParser) -> None:
|
||||
help="Only consider placements using >= this many nodes.",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--instance-meta", choices=["ring", "jaccl", "both"], default="both"
|
||||
"--instance-meta",
|
||||
choices=["ring", "jaccl", "vllm", "both"],
|
||||
default="both",
|
||||
)
|
||||
ap.add_argument(
|
||||
"--sharding", choices=["pipeline", "tensor", "both"], default="both"
|
||||
|
||||
@@ -12,24 +12,24 @@ timeout = 7200.0
|
||||
settle_timeout = 60.0
|
||||
|
||||
# Workload
|
||||
pp = [4096]
|
||||
tg = [512]
|
||||
pp = [4096, 8192]
|
||||
tg = [128]
|
||||
repeat = 1
|
||||
warmup = 0
|
||||
|
||||
json_out = "bench/prefill_decode_results.json"
|
||||
|
||||
[prefill]
|
||||
model = "mlx-community/gpt-oss-20b-MXFP4-Q8"
|
||||
node = "mike"
|
||||
instance_meta = "ring"
|
||||
model = "sakamakismile/Qwen3.6-27B-NVFP4"
|
||||
node = "gx10-de89"
|
||||
instance_meta = "vllm"
|
||||
sharding = "pipeline"
|
||||
min_nodes = 1
|
||||
max_nodes = 1
|
||||
|
||||
[decode]
|
||||
model = "mlx-community/gpt-oss-20b-MXFP4-Q8"
|
||||
node = "james"
|
||||
model = "mlx-community/Qwen3.6-27B-4bit"
|
||||
node = "Ryuichi’s MacBook Pro"
|
||||
instance_meta = "ring"
|
||||
sharding = "pipeline"
|
||||
min_nodes = 1
|
||||
|
||||
@@ -31,6 +31,7 @@ from typing import Any
|
||||
|
||||
from exo_bench import (
|
||||
PromptSizer,
|
||||
SystemMetricsSampler,
|
||||
format_peak_memory,
|
||||
load_tokenizer_for_bench,
|
||||
parse_int_list,
|
||||
@@ -278,6 +279,7 @@ def _run_phase(
|
||||
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]] = []
|
||||
@@ -288,10 +290,13 @@ def _run_phase(
|
||||
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
|
||||
@@ -315,11 +320,26 @@ def _run_phase(
|
||||
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"
|
||||
f"avg_elapsed={avg_elapsed:.2f}s{energy_str}"
|
||||
)
|
||||
time.sleep(2)
|
||||
return rows
|
||||
@@ -332,14 +352,36 @@ def _summarise(rows: list[dict[str, Any]]) -> dict[tuple[int, int], dict[str, fl
|
||||
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]],
|
||||
@@ -350,14 +392,17 @@ def _print_diff(
|
||||
prefill_alone = _summarise(prefill_alone_rows)
|
||||
keys = set(disagg.keys()) | set(decode_alone.keys()) | set(prefill_alone.keys())
|
||||
|
||||
width = 64
|
||||
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':>10} {'prompt_tps':>11} {'gen_tps':>9}"
|
||||
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)),
|
||||
@@ -365,26 +410,51 @@ def _print_diff(
|
||||
("prefill_alone", prefill_alone.get(key)),
|
||||
):
|
||||
if summary is None:
|
||||
logger.info(f" {label:<16} {'—':>10} {'—':>11} {'—':>9}")
|
||||
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']:>9.2f}s "
|
||||
f"{summary['elapsed_s']:>8.2f}s "
|
||||
f"{norm_str} "
|
||||
f"{summary['prompt_tps']:>11.1f} "
|
||||
f"{summary['gen_tps']:>9.2f}"
|
||||
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)
|
||||
if d and da and d["elapsed_s"] > 0:
|
||||
logger.info(
|
||||
f" speedup vs decode_alone: {da['elapsed_s'] / d['elapsed_s']:.2f}x"
|
||||
)
|
||||
if d and pa and d["elapsed_s"] > 0:
|
||||
logger.info(
|
||||
f" speedup vs prefill_alone: {pa['elapsed_s'] / d['elapsed_s']:.2f}x"
|
||||
)
|
||||
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)
|
||||
|
||||
|
||||
@@ -682,6 +752,16 @@ def main() -> int:
|
||||
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})
|
||||
@@ -700,6 +780,7 @@ def main() -> int:
|
||||
warmup=args.warmup,
|
||||
repeat=args.repeat,
|
||||
common_meta=common_meta,
|
||||
sampler=sampler,
|
||||
)
|
||||
all_rows.extend(prefill_alone_rows)
|
||||
|
||||
@@ -729,6 +810,7 @@ def main() -> int:
|
||||
warmup=args.warmup,
|
||||
repeat=args.repeat,
|
||||
common_meta=common_meta,
|
||||
sampler=sampler,
|
||||
)
|
||||
all_rows.extend(disagg_rows)
|
||||
|
||||
@@ -753,11 +835,13 @@ def main() -> int:
|
||||
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)
|
||||
|
||||
@@ -202,6 +202,7 @@
|
||||
let instanceType: string | null = null;
|
||||
if (instanceTag === "MlxRingInstance") instanceType = "MLX Ring";
|
||||
else if (instanceTag === "MlxJacclInstance") instanceType = "MLX RDMA";
|
||||
else if (instanceTag === "VllmInstance") instanceType = "vLLM";
|
||||
|
||||
let sharding: string | null = null;
|
||||
const inst = instance as {
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
*/
|
||||
|
||||
interface Props {
|
||||
/** "macbook pro" | "mac studio" | "mac mini" etc. */
|
||||
/** "macbook pro" | "mac studio" | "mac mini" | "dgx spark" | "linux" etc. */
|
||||
deviceType: string;
|
||||
/** Center X coordinate in SVG space */
|
||||
cx: number;
|
||||
@@ -38,10 +38,43 @@
|
||||
const LOGO_NATIVE_WIDTH = 814;
|
||||
const LOGO_NATIVE_HEIGHT = 1000;
|
||||
|
||||
// NVIDIA logo SVG path
|
||||
const NVIDIA_LOGO_PATH =
|
||||
"M0.81 0.429V0.299c0.013 -0.001 0.026 -0.002 0.038 -0.002 0.355 -0.011 0.588 0.306 0.588 0.306S1.186 0.952 0.916 0.952c-0.036 0 -0.071 -0.006 -0.105 -0.017V0.542c0.138 0.017 0.166 0.078 0.249 0.216l0.185 -0.155s-0.135 -0.177 -0.362 -0.177c-0.024 -0.001 -0.048 0.001 -0.072 0.003m0 -0.429v0.194l0.038 -0.002c0.494 -0.017 0.816 0.405 0.816 0.405s-0.37 0.45 -0.754 0.45c-0.034 0 -0.066 -0.003 -0.099 -0.009v0.12c0.027 0.003 0.055 0.006 0.082 0.006 0.358 0 0.618 -0.183 0.869 -0.399 0.042 0.034 0.212 0.114 0.247 0.15 -0.238 0.2 -0.794 0.361 -1.11 0.361 -0.03 0 -0.059 -0.002 -0.088 -0.005v0.169h1.362V0zm0 0.935v0.102c-0.331 -0.059 -0.423 -0.404 -0.423 -0.404s0.159 -0.176 0.423 -0.205v0.112h-0.001C0.671 0.524 0.562 0.654 0.562 0.654s0.062 0.218 0.248 0.282m-0.588 -0.316s0.196 -0.29 0.589 -0.32V0.194C0.376 0.229 0 0.597 0 0.597s0.213 0.616 0.81 0.672v-0.112c-0.438 -0.054 -0.588 -0.538 -0.588 -0.538";
|
||||
|
||||
const wireColor = "rgba(179,179,179,0.8)";
|
||||
const strokeWidth = 1.5;
|
||||
|
||||
const modelLower = $derived(deviceType.toLowerCase());
|
||||
const isSpark = $derived(
|
||||
modelLower.includes("dgx") || modelLower.includes("gx10"),
|
||||
);
|
||||
const isLinux = $derived(!isSpark && modelLower.startsWith("linux"));
|
||||
const isLinuxLaptop = $derived(isLinux && modelLower.includes("laptop"));
|
||||
|
||||
// ── DGX Spark dimensions ──
|
||||
const dgxW = $derived(size * 1.55);
|
||||
const dgxH = $derived(size * 0.58);
|
||||
const dgxX = $derived(cx - dgxW / 2);
|
||||
const dgxY = $derived(cy - dgxH / 2);
|
||||
const dgxChassisX = $derived(dgxX - dgxW * 0.03);
|
||||
const dgxChassisW = $derived(dgxW * 1.05);
|
||||
const dgxHandleW = $derived(dgxW * 0.27);
|
||||
const dgxHandleGap = $derived(dgxH * 0.05);
|
||||
const dgxHandleH = $derived(dgxH - dgxHandleGap * 2);
|
||||
const dgxHandleY = $derived(dgxY + dgxHandleGap);
|
||||
const dgxInnerHandleW = $derived(dgxW * 0.12);
|
||||
const dgxInnerHandleH = $derived(dgxHandleH - dgxH * 0.06);
|
||||
const dgxLeftHandleX = $derived(dgxX + 4);
|
||||
const dgxRightHandleX = $derived(dgxX + dgxW - dgxHandleW - 4);
|
||||
const dgxClipId = $derived(`di-dgx-${uid}`);
|
||||
const dgxTextureId = $derived(`di-dgx-tex-${uid}`);
|
||||
|
||||
// ── Linux Desktop dimensions (reuses Mac Studio proportions) ──
|
||||
const linuxDesktopClipId = $derived(`di-linux-desktop-${uid}`);
|
||||
|
||||
// ── Linux Laptop dimensions (reuses MacBook proportions) ──
|
||||
const linuxScreenClipId = $derived(`di-linux-screen-${uid}`);
|
||||
|
||||
// ── Mac Studio dimensions (same ratios as TopologyGraph) ──
|
||||
const studioW = $derived(size * 1.25);
|
||||
@@ -114,7 +147,264 @@
|
||||
const studioClipId = $derived(`di-studio-${uid}`);
|
||||
</script>
|
||||
|
||||
{#if modelLower === "mac studio" || modelLower === "mac mini"}
|
||||
{#if isSpark}
|
||||
<!-- DGX Spark -->
|
||||
<defs>
|
||||
<clipPath id={dgxClipId}>
|
||||
<rect x={dgxX} y={dgxY} width={dgxW} height={dgxH} rx="3" />
|
||||
</clipPath>
|
||||
<pattern
|
||||
id={dgxTextureId}
|
||||
patternUnits="userSpaceOnUse"
|
||||
width="8"
|
||||
height="8"
|
||||
>
|
||||
<rect width="8" height="8" fill="#6f6248" />
|
||||
<circle cx="2" cy="2" r="1" fill="#5a4f3b" opacity="0.5" />
|
||||
<circle cx="6" cy="6" r="1" fill="#4a4232" opacity="0.45" />
|
||||
</pattern>
|
||||
</defs>
|
||||
|
||||
<!-- Main body -->
|
||||
<rect
|
||||
x={dgxChassisX}
|
||||
y={dgxY}
|
||||
width={dgxChassisW}
|
||||
height={dgxH}
|
||||
rx="3"
|
||||
fill="url(#{dgxTextureId})"
|
||||
stroke={wireColor}
|
||||
stroke-width={strokeWidth}
|
||||
/>
|
||||
|
||||
<!-- Side border accents -->
|
||||
<rect
|
||||
x={dgxChassisX}
|
||||
y={dgxY}
|
||||
width={dgxW * 0.02}
|
||||
height={dgxH}
|
||||
fill="#8a7a56"
|
||||
/>
|
||||
<rect
|
||||
x={dgxChassisX + dgxChassisW - dgxW * 0.02}
|
||||
y={dgxY}
|
||||
width={dgxW * 0.02}
|
||||
height={dgxH}
|
||||
fill="#8a7a56"
|
||||
/>
|
||||
|
||||
<!-- Memory fill -->
|
||||
{#if ramPercent > 0}
|
||||
<rect
|
||||
x={dgxX}
|
||||
y={dgxY + dgxH - (ramPercent / 100) * dgxH}
|
||||
width={dgxW}
|
||||
height={(ramPercent / 100) * dgxH}
|
||||
fill="rgba(255,215,0,0.45)"
|
||||
clip-path="url(#{dgxClipId})"
|
||||
/>
|
||||
{/if}
|
||||
|
||||
<!-- Left handle -->
|
||||
<rect
|
||||
x={dgxLeftHandleX}
|
||||
y={dgxHandleY}
|
||||
width={dgxHandleW}
|
||||
height={dgxHandleH}
|
||||
rx="2.4"
|
||||
fill="#b3a170"
|
||||
stroke="#403723"
|
||||
stroke-width="0.7"
|
||||
/>
|
||||
<rect
|
||||
x={dgxLeftHandleX + dgxHandleW * 0.06}
|
||||
y={dgxHandleY + dgxH * 0.03}
|
||||
width={dgxInnerHandleW}
|
||||
height={dgxInnerHandleH}
|
||||
rx="1.6"
|
||||
fill="#8a7a56"
|
||||
/>
|
||||
|
||||
<!-- Right handle -->
|
||||
<rect
|
||||
x={dgxRightHandleX}
|
||||
y={dgxHandleY}
|
||||
width={dgxHandleW}
|
||||
height={dgxHandleH}
|
||||
rx="2.4"
|
||||
fill="#b3a170"
|
||||
stroke="#403723"
|
||||
stroke-width="0.7"
|
||||
/>
|
||||
<rect
|
||||
x={dgxRightHandleX + dgxHandleW - dgxInnerHandleW - dgxHandleW * 0.08}
|
||||
y={dgxHandleY + dgxH * 0.03}
|
||||
width={dgxInnerHandleW}
|
||||
height={dgxInnerHandleH}
|
||||
rx="1.6"
|
||||
fill="#8a7a56"
|
||||
/>
|
||||
|
||||
<!-- NVIDIA logo (rotated 90deg on left handle) -->
|
||||
{@const badgeW = dgxW * 0.09}
|
||||
{@const badgeH = dgxHandleH * 0.5}
|
||||
{@const badgeX = dgxLeftHandleX + dgxHandleW - badgeW - dgxHandleW * 0.06}
|
||||
{@const badgeYPos = dgxHandleY + (dgxHandleH - badgeH) / 2}
|
||||
{@const textSz = badgeW * 0.58}
|
||||
{@const logoW = textSz * 1.2}
|
||||
{@const logoH = logoW * (1.438 / 2.174)}
|
||||
{@const ctrX = badgeX + badgeW / 2 - badgeW * 0.03}
|
||||
{@const ctrY = badgeYPos + badgeH / 2}
|
||||
{@const labelGap = badgeW * 0.15}
|
||||
{@const totalW = logoW + labelGap + textSz * 3.6}
|
||||
<g transform="rotate(90 {ctrX} {ctrY})">
|
||||
<svg
|
||||
x={ctrX - totalW / 2}
|
||||
y={ctrY - logoH / 2}
|
||||
width={logoW}
|
||||
height={logoH}
|
||||
viewBox="0 0 2.174 1.438"
|
||||
>
|
||||
<path d={NVIDIA_LOGO_PATH} fill="#76b900" />
|
||||
</svg>
|
||||
<text
|
||||
x={ctrX - totalW / 2 + logoW + labelGap}
|
||||
y={ctrY}
|
||||
text-anchor="start"
|
||||
dominant-baseline="middle"
|
||||
fill="#8a7a56"
|
||||
font-size={textSz}
|
||||
font-family="monospace"
|
||||
font-weight="700">NVIDIA</text
|
||||
>
|
||||
</g>
|
||||
{:else if isLinuxLaptop}
|
||||
<!-- Linux Laptop — MacBook shape with Tux logo -->
|
||||
<defs>
|
||||
<clipPath id={linuxScreenClipId}>
|
||||
<rect
|
||||
x={mbScreenX + mbBezel}
|
||||
y={mbY + mbBezel}
|
||||
width={mbScreenW - mbBezel * 2}
|
||||
height={mbScreenH - mbBezel * 2}
|
||||
rx="2"
|
||||
/>
|
||||
</clipPath>
|
||||
</defs>
|
||||
|
||||
<rect
|
||||
x={mbScreenX}
|
||||
y={mbY}
|
||||
width={mbScreenW}
|
||||
height={mbScreenH}
|
||||
rx="3"
|
||||
fill="#1a1a1a"
|
||||
stroke={wireColor}
|
||||
stroke-width={strokeWidth}
|
||||
/>
|
||||
<rect
|
||||
x={mbScreenX + mbBezel}
|
||||
y={mbY + mbBezel}
|
||||
width={mbScreenW - mbBezel * 2}
|
||||
height={mbScreenH - mbBezel * 2}
|
||||
rx="2"
|
||||
fill="#0a0a12"
|
||||
/>
|
||||
{#if ramPercent > 0}
|
||||
<rect
|
||||
x={mbScreenX + mbBezel}
|
||||
y={mbY + mbBezel + (mbMemTotalH - mbMemH)}
|
||||
width={mbScreenW - mbBezel * 2}
|
||||
height={mbMemH}
|
||||
fill="rgba(255,215,0,0.85)"
|
||||
clip-path="url(#{linuxScreenClipId})"
|
||||
/>
|
||||
{/if}
|
||||
|
||||
<!-- Terminal prompt on screen -->
|
||||
<text
|
||||
x={cx}
|
||||
y={mbY + mbScreenH / 2}
|
||||
text-anchor="middle"
|
||||
dominant-baseline="middle"
|
||||
fill="#FFFFFF"
|
||||
opacity="0.9"
|
||||
font-size={mbScreenH * 0.25}
|
||||
font-family="SF Mono, Monaco, monospace"
|
||||
font-weight="700">{">_"}</text
|
||||
>
|
||||
|
||||
<path
|
||||
d="M {mbBaseTopX} {mbBaseY} L {mbBaseTopX +
|
||||
mbBaseTopW} {mbBaseY} L {mbBaseBottomX + mbBaseBottomW} {mbBaseY +
|
||||
mbBaseH} L {mbBaseBottomX} {mbBaseY + mbBaseH} Z"
|
||||
fill="#2c2c2c"
|
||||
stroke={wireColor}
|
||||
stroke-width="1"
|
||||
/>
|
||||
<rect
|
||||
x={mbKbX}
|
||||
y={mbKbY}
|
||||
width={mbKbW}
|
||||
height={mbKbH}
|
||||
fill="rgba(0,0,0,0.2)"
|
||||
rx="2"
|
||||
/>
|
||||
<rect
|
||||
x={mbTpX}
|
||||
y={mbTpY}
|
||||
width={mbTpW}
|
||||
height={mbTpH}
|
||||
fill="rgba(255,255,255,0.08)"
|
||||
rx="2"
|
||||
/>
|
||||
{:else if isLinux}
|
||||
<!-- Linux Desktop — Mac Studio shape with Tux logo -->
|
||||
<defs>
|
||||
<clipPath id={linuxDesktopClipId}>
|
||||
<rect
|
||||
x={studioX}
|
||||
y={studioY + studioTopH}
|
||||
width={studioW}
|
||||
height={studioH - studioTopH}
|
||||
rx={studioCorner - 1}
|
||||
/>
|
||||
</clipPath>
|
||||
</defs>
|
||||
|
||||
<rect
|
||||
x={studioX}
|
||||
y={studioY}
|
||||
width={studioW}
|
||||
height={studioH}
|
||||
rx={studioCorner}
|
||||
fill="#1a1a1a"
|
||||
stroke={wireColor}
|
||||
stroke-width={strokeWidth}
|
||||
/>
|
||||
{#if ramPercent > 0}
|
||||
<rect
|
||||
x={studioX}
|
||||
y={studioY + studioTopH + (studioMemTotalH - studioMemH)}
|
||||
width={studioW}
|
||||
height={studioMemH}
|
||||
fill="rgba(255,215,0,0.75)"
|
||||
clip-path="url(#{linuxDesktopClipId})"
|
||||
/>
|
||||
{/if}
|
||||
|
||||
<!-- Terminal prompt on front face -->
|
||||
<text
|
||||
x={cx}
|
||||
y={studioY + studioTopH + (studioH - studioTopH) / 2}
|
||||
text-anchor="middle"
|
||||
dominant-baseline="middle"
|
||||
fill="rgba(255,255,255,0.5)"
|
||||
font-size={(studioH - studioTopH) * 0.4}
|
||||
font-family="SF Mono, Monaco, monospace"
|
||||
font-weight="700">{">_"}</text
|
||||
>
|
||||
{:else if modelLower === "mac studio" || modelLower === "mac mini"}
|
||||
<!-- Mac Studio / Mac Mini -->
|
||||
<defs>
|
||||
<clipPath id={studioClipId}>
|
||||
|
||||
@@ -23,7 +23,7 @@
|
||||
} | null;
|
||||
nodes?: Record<string, NodeInfo>;
|
||||
sharding?: "Pipeline" | "Tensor";
|
||||
runtime?: "MlxRing" | "MlxJaccl";
|
||||
runtime?: "MlxRing" | "MlxJaccl" | "Vllm";
|
||||
onLaunch?: () => void;
|
||||
tags?: string[];
|
||||
apiPreview?: PlacementPreview | null;
|
||||
@@ -168,8 +168,10 @@
|
||||
|
||||
function getDeviceType(
|
||||
name: string,
|
||||
): "macbook" | "studio" | "mini" | "unknown" {
|
||||
): "macbook" | "studio" | "mini" | "dgx" | "linux" | "unknown" {
|
||||
const lower = name.toLowerCase();
|
||||
if (lower.includes("dgx") || lower.includes("gx10")) return "dgx";
|
||||
if (lower.includes("linux")) return "linux";
|
||||
if (lower.includes("macbook")) return "macbook";
|
||||
if (lower.includes("studio")) return "studio";
|
||||
if (lower.includes("mini")) return "mini";
|
||||
@@ -576,13 +578,17 @@
|
||||
class="px-1.5 py-0.5 text-xs font-mono tracking-wider uppercase bg-exo-medium-gray/30 text-exo-light-gray border border-exo-medium-gray/40"
|
||||
title={runtime === "MlxRing"
|
||||
? "Ring: standard networking. Works over any connection (Wi-Fi, Ethernet, Thunderbolt)."
|
||||
: "RDMA: direct memory access over Thunderbolt. Significantly faster for multi-device inference."}
|
||||
: runtime === "MlxJaccl"
|
||||
? "RDMA: direct memory access over Thunderbolt. Significantly faster for multi-device inference."
|
||||
: "vLLM: NVIDIA CUDA inference engine."}
|
||||
>
|
||||
{runtime === "MlxRing"
|
||||
? "MLX Ring"
|
||||
: runtime === "MlxJaccl"
|
||||
? "MLX RDMA"
|
||||
: runtime}
|
||||
: runtime === "Vllm"
|
||||
? "vLLM"
|
||||
: runtime}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
@@ -990,6 +996,81 @@
|
||||
/>
|
||||
{/if}
|
||||
</g>
|
||||
{:else if node.deviceType === "dgx"}
|
||||
<!-- DGX Spark icon -->
|
||||
{@const s = node.iconSize}
|
||||
{@const dgxW = s * 1.4}
|
||||
{@const dgxH = s * 0.52}
|
||||
<g transform="translate({-dgxW / 2}, {-dgxH / 2})">
|
||||
<!-- Chassis -->
|
||||
<rect
|
||||
x="0"
|
||||
y="0"
|
||||
width={dgxW}
|
||||
height={dgxH}
|
||||
rx="2"
|
||||
fill="#6f6248"
|
||||
stroke={node.isUsed ? "#FFD700" : "#4B5563"}
|
||||
stroke-width="1.5"
|
||||
/>
|
||||
<!-- Side accents -->
|
||||
<rect
|
||||
x="0"
|
||||
y="0"
|
||||
width={dgxW * 0.02}
|
||||
height={dgxH}
|
||||
fill="#8a7a56"
|
||||
/>
|
||||
<rect
|
||||
x={dgxW - dgxW * 0.02}
|
||||
y="0"
|
||||
width={dgxW * 0.02}
|
||||
height={dgxH}
|
||||
fill="#8a7a56"
|
||||
/>
|
||||
<!-- Left handle -->
|
||||
<rect
|
||||
x={dgxW * 0.04}
|
||||
y={dgxH * 0.08}
|
||||
width={dgxW * 0.22}
|
||||
height={dgxH * 0.84}
|
||||
rx="2"
|
||||
fill="#b3a170"
|
||||
stroke="#403723"
|
||||
stroke-width="0.5"
|
||||
/>
|
||||
<!-- Right handle -->
|
||||
<rect
|
||||
x={dgxW - dgxW * 0.04 - dgxW * 0.22}
|
||||
y={dgxH * 0.08}
|
||||
width={dgxW * 0.22}
|
||||
height={dgxH * 0.84}
|
||||
rx="2"
|
||||
fill="#b3a170"
|
||||
stroke="#403723"
|
||||
stroke-width="0.5"
|
||||
/>
|
||||
<!-- Memory fill -->
|
||||
<rect
|
||||
x="2"
|
||||
y={dgxH - dgxH * (node.currentPercent / 100)}
|
||||
width={dgxW - 4}
|
||||
height={dgxH * (node.currentPercent / 100)}
|
||||
fill="rgba(255,215,0,0.35)"
|
||||
/>
|
||||
{#if node.modelUsageGB > 0 && node.isUsed}
|
||||
<rect
|
||||
x="2"
|
||||
y={dgxH - dgxH * (node.newPercent / 100)}
|
||||
width={dgxW - 4}
|
||||
height={dgxH *
|
||||
((node.newPercent - node.currentPercent) / 100)}
|
||||
fill="#FFD700"
|
||||
filter="url(#memGlow-{filterId})"
|
||||
class="animate-pulse-slow"
|
||||
/>
|
||||
{/if}
|
||||
</g>
|
||||
{:else}
|
||||
<!-- Unknown device - hexagon -->
|
||||
<g
|
||||
|
||||
@@ -9,6 +9,7 @@
|
||||
capabilities?: string[];
|
||||
family?: string;
|
||||
is_custom?: boolean;
|
||||
requires_vllm?: boolean;
|
||||
}
|
||||
|
||||
interface ModelGroup {
|
||||
@@ -19,6 +20,7 @@
|
||||
variants: ModelInfo[];
|
||||
smallestVariant: ModelInfo;
|
||||
hasMultipleVariants: boolean;
|
||||
requiresVllm: boolean;
|
||||
}
|
||||
|
||||
type DownloadAvailability = {
|
||||
@@ -213,6 +215,14 @@
|
||||
<span class="font-mono text-sm text-white truncate">
|
||||
{group.name}
|
||||
</span>
|
||||
{#if group.requiresVllm}
|
||||
<span
|
||||
class="text-[10px] font-mono px-1.5 py-0.5 rounded bg-orange-500/15 text-orange-300 border border-orange-400/30 flex-shrink-0 tracking-wider uppercase"
|
||||
title="Requires vLLM runtime"
|
||||
>
|
||||
vLLM
|
||||
</span>
|
||||
{/if}
|
||||
<!-- Capability icons -->
|
||||
{#each group.capabilities.filter((c) => c !== "text") as cap}
|
||||
{#if cap === "thinking"}
|
||||
@@ -523,6 +533,15 @@
|
||||
{variant.quantization || "default"}
|
||||
</span>
|
||||
|
||||
{#if variant.requires_vllm}
|
||||
<span
|
||||
class="text-[10px] font-mono px-1.5 py-0.5 rounded bg-orange-500/15 text-orange-300 border border-orange-400/30 flex-shrink-0 tracking-wider uppercase"
|
||||
title="Requires vLLM runtime"
|
||||
>
|
||||
vLLM
|
||||
</span>
|
||||
{/if}
|
||||
|
||||
<!-- Size -->
|
||||
<span
|
||||
class="text-xs font-mono flex-1 {getSizeClassForFitStatus(
|
||||
@@ -628,6 +647,7 @@
|
||||
variants: [variant],
|
||||
smallestVariant: variant,
|
||||
hasMultipleVariants: false,
|
||||
requiresVllm: variant.requires_vllm === true,
|
||||
});
|
||||
}}
|
||||
title="View variant details"
|
||||
|
||||
@@ -22,6 +22,7 @@
|
||||
is_custom?: boolean;
|
||||
tasks?: string[];
|
||||
hugging_face_id?: string;
|
||||
requires_vllm?: boolean;
|
||||
}
|
||||
|
||||
interface ModelGroup {
|
||||
@@ -32,6 +33,7 @@
|
||||
variants: ModelInfo[];
|
||||
smallestVariant: ModelInfo;
|
||||
hasMultipleVariants: boolean;
|
||||
requiresVllm: boolean;
|
||||
}
|
||||
|
||||
interface FilterState {
|
||||
@@ -396,6 +398,7 @@
|
||||
variants: [],
|
||||
smallestVariant: model,
|
||||
hasMultipleVariants: false,
|
||||
requiresVllm: true,
|
||||
});
|
||||
}
|
||||
|
||||
@@ -430,6 +433,7 @@
|
||||
(a.storage_size_megabytes || 0) - (b.storage_size_megabytes || 0),
|
||||
);
|
||||
group.hasMultipleVariants = group.variants.length > 1;
|
||||
group.requiresVllm = group.variants.every((v) => v.requires_vllm);
|
||||
}
|
||||
|
||||
// Convert to array and sort by smallest variant size (biggest first)
|
||||
@@ -587,6 +591,7 @@
|
||||
variants: [model],
|
||||
smallestVariant: model,
|
||||
hasMultipleVariants: false,
|
||||
requiresVllm: model.requires_vllm === true,
|
||||
});
|
||||
}
|
||||
}
|
||||
@@ -1165,6 +1170,17 @@
|
||||
<span class="text-white/40">Variants:</span>
|
||||
<span class="text-white/70">{infoGroup.variants.length}</span>
|
||||
</div>
|
||||
{#if infoGroup.requiresVllm}
|
||||
<div class="flex items-center gap-2">
|
||||
<span class="text-white/40">Runtime:</span>
|
||||
<span
|
||||
class="text-[10px] font-mono px-1.5 py-0.5 rounded bg-orange-500/15 text-orange-300 border border-orange-400/30 tracking-wider uppercase"
|
||||
>
|
||||
vLLM
|
||||
</span>
|
||||
<span class="text-white/40 text-[11px]">required</span>
|
||||
</div>
|
||||
{/if}
|
||||
{#if infoGroup.variants.length > 0}
|
||||
<div class="mt-3 pt-3 border-t border-exo-yellow/10">
|
||||
<span class="text-white/40">Available quantizations:</span>
|
||||
|
||||
@@ -219,7 +219,7 @@
|
||||
Prefill vs Decode
|
||||
</summary>
|
||||
<div class="mt-2 text-white/80 text-sm leading-relaxed">
|
||||
Prefill is the compute-heavy pass that consumes the entire prompt and
|
||||
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
|
||||
|
||||
@@ -117,6 +117,10 @@
|
||||
const LOGO_NATIVE_WIDTH = 814;
|
||||
const LOGO_NATIVE_HEIGHT = 1000;
|
||||
|
||||
// NVIDIA logo SVG path (from exo-nvidia)
|
||||
const NVIDIA_LOGO_PATH =
|
||||
"M0.81 0.429V0.299c0.013 -0.001 0.026 -0.002 0.038 -0.002 0.355 -0.011 0.588 0.306 0.588 0.306S1.186 0.952 0.916 0.952c-0.036 0 -0.071 -0.006 -0.105 -0.017V0.542c0.138 0.017 0.166 0.078 0.249 0.216l0.185 -0.155s-0.135 -0.177 -0.362 -0.177c-0.024 -0.001 -0.048 0.001 -0.072 0.003m0 -0.429v0.194l0.038 -0.002c0.494 -0.017 0.816 0.405 0.816 0.405s-0.37 0.45 -0.754 0.45c-0.034 0 -0.066 -0.003 -0.099 -0.009v0.12c0.027 0.003 0.055 0.006 0.082 0.006 0.358 0 0.618 -0.183 0.869 -0.399 0.042 0.034 0.212 0.114 0.247 0.15 -0.238 0.2 -0.794 0.361 -1.11 0.361 -0.03 0 -0.059 -0.002 -0.088 -0.005v0.169h1.362V0zm0 0.935v0.102c-0.331 -0.059 -0.423 -0.404 -0.423 -0.404s0.159 -0.176 0.423 -0.205v0.112h-0.001C0.671 0.524 0.562 0.654 0.562 0.654s0.062 0.218 0.248 0.282m-0.588 -0.316s0.196 -0.29 0.589 -0.32V0.194C0.376 0.229 0 0.597 0 0.597s0.213 0.616 0.81 0.672v-0.112c-0.438 -0.054 -0.588 -0.538 -0.588 -0.538";
|
||||
|
||||
function formatBytes(bytes: number, decimals = 1): string {
|
||||
if (!bytes || bytes === 0) return "0B";
|
||||
const k = 1024;
|
||||
@@ -554,6 +558,13 @@
|
||||
const clipPathId = `clip-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
|
||||
|
||||
const modelLower = modelId.toLowerCase();
|
||||
const identity = identitiesData[nodeInfo.id];
|
||||
const nameLower = (friendlyName || "").toLowerCase();
|
||||
const isSpark = modelLower.includes("dgx") || modelLower.includes("gx10");
|
||||
const isLinux =
|
||||
!isSpark &&
|
||||
(modelLower.startsWith("linux") || identity?.osVersion === "Linux");
|
||||
const isLinuxLaptop = isLinux && modelLower.includes("laptop");
|
||||
|
||||
// Check node states for styling
|
||||
const isHighlighted = highlightedNodes.has(nodeInfo.id);
|
||||
@@ -623,7 +634,382 @@
|
||||
`${friendlyName}\nID: ${nodeInfo.id.slice(-8)}\nMemory: ${formatBytes(ramUsed)}/${formatBytes(ramTotal)}`,
|
||||
);
|
||||
|
||||
if (modelLower === "mac studio") {
|
||||
if (isSpark) {
|
||||
// NVIDIA DGX Spark — gold chassis with textured front, side handles, and NVIDIA badge
|
||||
iconBaseWidth = nodeRadius * 1.55;
|
||||
iconBaseHeight = nodeRadius * 0.58;
|
||||
const x = nodeInfo.x - iconBaseWidth / 2;
|
||||
const y = nodeInfo.y - iconBaseHeight / 2;
|
||||
const chassisX = x - iconBaseWidth * 0.03;
|
||||
const chassisWidth = iconBaseWidth * 1.05;
|
||||
const cornerRadius = 3;
|
||||
|
||||
const dgxClipId = `dgx-clip-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
|
||||
defs
|
||||
.append("clipPath")
|
||||
.attr("id", dgxClipId)
|
||||
.append("rect")
|
||||
.attr("x", x)
|
||||
.attr("y", y)
|
||||
.attr("width", iconBaseWidth)
|
||||
.attr("height", iconBaseHeight)
|
||||
.attr("rx", cornerRadius);
|
||||
|
||||
// Chassis texture pattern
|
||||
const textureId = `chassis-texture-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
|
||||
defs
|
||||
.append("pattern")
|
||||
.attr("id", textureId)
|
||||
.attr("patternUnits", "userSpaceOnUse")
|
||||
.attr("width", 8)
|
||||
.attr("height", 8);
|
||||
const texturePattern = defs.select(`#${textureId}`);
|
||||
texturePattern
|
||||
.append("rect")
|
||||
.attr("width", 8)
|
||||
.attr("height", 8)
|
||||
.attr("fill", "#6f6248");
|
||||
texturePattern
|
||||
.append("circle")
|
||||
.attr("cx", 2)
|
||||
.attr("cy", 2)
|
||||
.attr("r", 1)
|
||||
.attr("fill", "#5a4f3b")
|
||||
.attr("opacity", 0.5);
|
||||
texturePattern
|
||||
.append("circle")
|
||||
.attr("cx", 6)
|
||||
.attr("cy", 6)
|
||||
.attr("r", 1)
|
||||
.attr("fill", "#4a4232")
|
||||
.attr("opacity", 0.45);
|
||||
|
||||
// Main body
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("class", "node-outline")
|
||||
.attr("x", chassisX)
|
||||
.attr("y", y)
|
||||
.attr("width", chassisWidth)
|
||||
.attr("height", iconBaseHeight)
|
||||
.attr("rx", cornerRadius)
|
||||
.attr("fill", `url(#${textureId})`)
|
||||
.attr("stroke", wireColor)
|
||||
.attr("stroke-width", strokeWidth);
|
||||
|
||||
// Side border accents
|
||||
const sideThickness = iconBaseWidth * 0.02;
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", chassisX)
|
||||
.attr("y", y)
|
||||
.attr("width", sideThickness)
|
||||
.attr("height", iconBaseHeight)
|
||||
.attr("fill", "#8a7a56");
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", chassisX + chassisWidth - sideThickness)
|
||||
.attr("y", y)
|
||||
.attr("width", sideThickness)
|
||||
.attr("height", iconBaseHeight)
|
||||
.attr("fill", "#8a7a56");
|
||||
|
||||
// Memory fill (bottom up)
|
||||
if (ramUsagePercent > 0) {
|
||||
const memFillHeight = (ramUsagePercent / 100) * iconBaseHeight;
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", x)
|
||||
.attr("y", y + iconBaseHeight - memFillHeight)
|
||||
.attr("width", iconBaseWidth)
|
||||
.attr("height", memFillHeight)
|
||||
.attr("fill", "rgba(255,215,0,0.45)")
|
||||
.attr("clip-path", `url(#${dgxClipId})`);
|
||||
}
|
||||
|
||||
// Side handles with inner recess
|
||||
const handleWidth = iconBaseWidth * 0.27;
|
||||
const handleGap = iconBaseHeight * 0.05;
|
||||
const handleHeight = iconBaseHeight - handleGap * 2;
|
||||
const handleY = y + handleGap;
|
||||
const innerHandleWidth = iconBaseWidth * 0.12;
|
||||
const innerHandleHeight = handleHeight - iconBaseHeight * 0.06;
|
||||
const leftHandleX = x + 4;
|
||||
const rightHandleX = x + iconBaseWidth - handleWidth - 4;
|
||||
|
||||
// Left handle
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", leftHandleX)
|
||||
.attr("y", handleY)
|
||||
.attr("width", handleWidth)
|
||||
.attr("height", handleHeight)
|
||||
.attr("rx", 2.4)
|
||||
.attr("fill", "#b3a170")
|
||||
.attr("stroke", "#403723")
|
||||
.attr("stroke-width", 0.7);
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", leftHandleX + handleWidth * 0.06)
|
||||
.attr("y", handleY + iconBaseHeight * 0.03)
|
||||
.attr("width", innerHandleWidth)
|
||||
.attr("height", innerHandleHeight)
|
||||
.attr("rx", 1.6)
|
||||
.attr("fill", "#8a7a56");
|
||||
|
||||
// Right handle
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", rightHandleX)
|
||||
.attr("y", handleY)
|
||||
.attr("width", handleWidth)
|
||||
.attr("height", handleHeight)
|
||||
.attr("rx", 2.4)
|
||||
.attr("fill", "#b3a170")
|
||||
.attr("stroke", "#403723")
|
||||
.attr("stroke-width", 0.7);
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr(
|
||||
"x",
|
||||
rightHandleX + handleWidth - innerHandleWidth - handleWidth * 0.08,
|
||||
)
|
||||
.attr("y", handleY + iconBaseHeight * 0.03)
|
||||
.attr("width", innerHandleWidth)
|
||||
.attr("height", innerHandleHeight)
|
||||
.attr("rx", 1.6)
|
||||
.attr("fill", "#8a7a56");
|
||||
|
||||
// NVIDIA logo + text label (rotated 90 deg on left handle)
|
||||
const badgeWidth = iconBaseWidth * 0.09;
|
||||
const badgeHeight = handleHeight * 0.5;
|
||||
const badgeX =
|
||||
leftHandleX + handleWidth - badgeWidth - handleWidth * 0.06;
|
||||
const badgeY = handleY + (handleHeight - badgeHeight) / 2;
|
||||
const textSize = badgeWidth * 0.58;
|
||||
const logoWidth = textSize * 1.2;
|
||||
const logoHeight = logoWidth * (1.438 / 2.174);
|
||||
const centerX = badgeX + badgeWidth / 2 - badgeWidth * 0.03;
|
||||
const centerY = badgeY + badgeHeight / 2;
|
||||
const gap = badgeWidth * 0.15;
|
||||
const totalWidth = logoWidth + gap + textSize * 3.6;
|
||||
|
||||
const labelGroup = nodeG
|
||||
.append("g")
|
||||
.attr("transform", `rotate(90 ${centerX} ${centerY})`);
|
||||
|
||||
labelGroup
|
||||
.append("svg")
|
||||
.attr("x", centerX - totalWidth / 2)
|
||||
.attr("y", centerY - logoHeight / 2)
|
||||
.attr("width", logoWidth)
|
||||
.attr("height", logoHeight)
|
||||
.attr("viewBox", "0 0 2.174 1.438")
|
||||
.append("path")
|
||||
.attr("d", NVIDIA_LOGO_PATH)
|
||||
.attr("fill", "#76b900");
|
||||
|
||||
labelGroup
|
||||
.append("text")
|
||||
.attr("x", centerX - totalWidth / 2 + logoWidth + gap)
|
||||
.attr("y", centerY)
|
||||
.attr("text-anchor", "start")
|
||||
.attr("dominant-baseline", "middle")
|
||||
.attr("fill", "#8a7a56")
|
||||
.attr("font-size", textSize)
|
||||
.attr("font-family", "monospace")
|
||||
.attr("font-weight", "700")
|
||||
.text("NVIDIA");
|
||||
} else if (isLinuxLaptop) {
|
||||
// Linux Laptop — same shape as MacBook but with Tux logo
|
||||
iconBaseWidth = nodeRadius * 1.6;
|
||||
iconBaseHeight = nodeRadius * 1.15;
|
||||
const x = nodeInfo.x - iconBaseWidth / 2;
|
||||
const y = nodeInfo.y - iconBaseHeight / 2;
|
||||
|
||||
const screenHeight = iconBaseHeight * 0.7;
|
||||
const baseHeight = iconBaseHeight * 0.3;
|
||||
const screenWidth = iconBaseWidth * 0.85;
|
||||
const screenX = nodeInfo.x - screenWidth / 2;
|
||||
const screenBezel = 3;
|
||||
|
||||
const linuxScreenClipId = `linux-screen-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
|
||||
defs
|
||||
.append("clipPath")
|
||||
.attr("id", linuxScreenClipId)
|
||||
.append("rect")
|
||||
.attr("x", screenX + screenBezel)
|
||||
.attr("y", y + screenBezel)
|
||||
.attr("width", screenWidth - screenBezel * 2)
|
||||
.attr("height", screenHeight - screenBezel * 2)
|
||||
.attr("rx", 2);
|
||||
|
||||
// Screen outer frame
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("class", "node-outline")
|
||||
.attr("x", screenX)
|
||||
.attr("y", y)
|
||||
.attr("width", screenWidth)
|
||||
.attr("height", screenHeight)
|
||||
.attr("rx", 3)
|
||||
.attr("fill", "#1a1a1a")
|
||||
.attr("stroke", wireColor)
|
||||
.attr("stroke-width", strokeWidth);
|
||||
|
||||
// Screen inner
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", screenX + screenBezel)
|
||||
.attr("y", y + screenBezel)
|
||||
.attr("width", screenWidth - screenBezel * 2)
|
||||
.attr("height", screenHeight - screenBezel * 2)
|
||||
.attr("rx", 2)
|
||||
.attr("fill", "#0a0a12");
|
||||
|
||||
// Memory fill on screen
|
||||
if (ramUsagePercent > 0) {
|
||||
const memFillTotalHeight = screenHeight - screenBezel * 2;
|
||||
const memFillActualHeight =
|
||||
(ramUsagePercent / 100) * memFillTotalHeight;
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", screenX + screenBezel)
|
||||
.attr(
|
||||
"y",
|
||||
y + screenBezel + (memFillTotalHeight - memFillActualHeight),
|
||||
)
|
||||
.attr("width", screenWidth - screenBezel * 2)
|
||||
.attr("height", memFillActualHeight)
|
||||
.attr("fill", "rgba(255,215,0,0.85)")
|
||||
.attr("clip-path", `url(#${linuxScreenClipId})`);
|
||||
}
|
||||
|
||||
// Terminal prompt on screen
|
||||
nodeG
|
||||
.append("text")
|
||||
.attr("x", nodeInfo.x)
|
||||
.attr("y", y + screenHeight / 2)
|
||||
.attr("text-anchor", "middle")
|
||||
.attr("dominant-baseline", "middle")
|
||||
.attr("fill", "#FFFFFF")
|
||||
.attr("opacity", 0.9)
|
||||
.attr("font-size", screenHeight * 0.25)
|
||||
.attr("font-family", "SF Mono, Monaco, monospace")
|
||||
.attr("font-weight", "700")
|
||||
.text(">_");
|
||||
|
||||
// Keyboard base (trapezoidal)
|
||||
const baseY = y + screenHeight;
|
||||
const baseTopWidth = screenWidth;
|
||||
const baseBottomWidth = iconBaseWidth;
|
||||
const baseTopX = nodeInfo.x - baseTopWidth / 2;
|
||||
const baseBottomX = nodeInfo.x - baseBottomWidth / 2;
|
||||
|
||||
nodeG
|
||||
.append("path")
|
||||
.attr(
|
||||
"d",
|
||||
`M ${baseTopX} ${baseY} L ${baseTopX + baseTopWidth} ${baseY} L ${baseBottomX + baseBottomWidth} ${baseY + baseHeight} L ${baseBottomX} ${baseY + baseHeight} Z`,
|
||||
)
|
||||
.attr("fill", "#2c2c2c")
|
||||
.attr("stroke", wireColor)
|
||||
.attr("stroke-width", 1);
|
||||
|
||||
// Keyboard area
|
||||
const keyboardX = baseTopX + 6;
|
||||
const keyboardY = baseY + 3;
|
||||
const keyboardWidth = baseTopWidth - 12;
|
||||
const keyboardHeight = baseHeight * 0.55;
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", keyboardX)
|
||||
.attr("y", keyboardY)
|
||||
.attr("width", keyboardWidth)
|
||||
.attr("height", keyboardHeight)
|
||||
.attr("fill", "rgba(0,0,0,0.2)")
|
||||
.attr("rx", 2);
|
||||
|
||||
// Trackpad
|
||||
const trackpadWidth = baseTopWidth * 0.4;
|
||||
const trackpadX = nodeInfo.x - trackpadWidth / 2;
|
||||
const trackpadY = baseY + keyboardHeight + 5;
|
||||
const trackpadHeight = baseHeight * 0.3;
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", trackpadX)
|
||||
.attr("y", trackpadY)
|
||||
.attr("width", trackpadWidth)
|
||||
.attr("height", trackpadHeight)
|
||||
.attr("fill", "rgba(255,255,255,0.08)")
|
||||
.attr("rx", 2);
|
||||
} else if (isLinux) {
|
||||
// Linux Desktop — same shape as Mac Studio but with Tux logo
|
||||
iconBaseWidth = nodeRadius * 1.25;
|
||||
iconBaseHeight = nodeRadius * 0.85;
|
||||
const x = nodeInfo.x - iconBaseWidth / 2;
|
||||
const y = nodeInfo.y - iconBaseHeight / 2;
|
||||
const cornerRadius = 4;
|
||||
const topSurfaceHeight = iconBaseHeight * 0.15;
|
||||
|
||||
const linuxDesktopClipId = `linux-desktop-${nodeInfo.id.replace(/[^a-zA-Z0-9]/g, "-")}`;
|
||||
defs
|
||||
.append("clipPath")
|
||||
.attr("id", linuxDesktopClipId)
|
||||
.append("rect")
|
||||
.attr("x", x)
|
||||
.attr("y", y + topSurfaceHeight)
|
||||
.attr("width", iconBaseWidth)
|
||||
.attr("height", iconBaseHeight - topSurfaceHeight)
|
||||
.attr("rx", cornerRadius - 1);
|
||||
|
||||
// Main body
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("class", "node-outline")
|
||||
.attr("x", x)
|
||||
.attr("y", y)
|
||||
.attr("width", iconBaseWidth)
|
||||
.attr("height", iconBaseHeight)
|
||||
.attr("rx", cornerRadius)
|
||||
.attr("fill", "#1a1a1a")
|
||||
.attr("stroke", wireColor)
|
||||
.attr("stroke-width", strokeWidth);
|
||||
|
||||
// Memory fill
|
||||
if (ramUsagePercent > 0) {
|
||||
const memFillTotalHeight = iconBaseHeight - topSurfaceHeight;
|
||||
const memFillActualHeight =
|
||||
(ramUsagePercent / 100) * memFillTotalHeight;
|
||||
nodeG
|
||||
.append("rect")
|
||||
.attr("x", x)
|
||||
.attr(
|
||||
"y",
|
||||
y + topSurfaceHeight + (memFillTotalHeight - memFillActualHeight),
|
||||
)
|
||||
.attr("width", iconBaseWidth)
|
||||
.attr("height", memFillActualHeight)
|
||||
.attr("fill", "rgba(255,215,0,0.75)")
|
||||
.attr("clip-path", `url(#${linuxDesktopClipId})`);
|
||||
}
|
||||
|
||||
// Terminal prompt on front face
|
||||
nodeG
|
||||
.append("text")
|
||||
.attr("x", nodeInfo.x)
|
||||
.attr(
|
||||
"y",
|
||||
y + topSurfaceHeight + (iconBaseHeight - topSurfaceHeight) / 2,
|
||||
)
|
||||
.attr("text-anchor", "middle")
|
||||
.attr("dominant-baseline", "middle")
|
||||
.attr("fill", "rgba(255,255,255,0.5)")
|
||||
.attr("font-size", (iconBaseHeight - topSurfaceHeight) * 0.4)
|
||||
.attr("font-family", "SF Mono, Monaco, monospace")
|
||||
.attr("font-weight", "700")
|
||||
.text(">_");
|
||||
} else if (modelLower === "mac studio") {
|
||||
// Mac Studio - classic cube with memory fill
|
||||
iconBaseWidth = nodeRadius * 1.25;
|
||||
iconBaseHeight = nodeRadius * 0.85;
|
||||
@@ -1182,8 +1568,12 @@
|
||||
debugLabelY += debugLineHeight;
|
||||
}
|
||||
|
||||
const identity = identitiesData[nodeInfo.id];
|
||||
if (identity?.osVersion) {
|
||||
const dbgIdentity = identitiesData[nodeInfo.id];
|
||||
if (dbgIdentity?.osVersion) {
|
||||
const osLabel =
|
||||
dbgIdentity.osVersion === "Linux"
|
||||
? "Linux"
|
||||
: `macOS ${dbgIdentity.osVersion}${dbgIdentity.osBuildVersion ? ` (${dbgIdentity.osBuildVersion})` : ""}`;
|
||||
nodeG
|
||||
.append("text")
|
||||
.attr("x", nodeInfo.x)
|
||||
@@ -1192,9 +1582,7 @@
|
||||
.attr("fill", "rgba(179,179,179,0.7)")
|
||||
.attr("font-size", debugFontSize)
|
||||
.attr("font-family", "SF Mono, Monaco, monospace")
|
||||
.text(
|
||||
`macOS ${identity.osVersion}${identity.osBuildVersion ? ` (${identity.osBuildVersion})` : ""}`,
|
||||
);
|
||||
.text(osLabel);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
@@ -65,6 +65,7 @@
|
||||
nodeThunderboltBridge,
|
||||
nodeIdentities,
|
||||
isConnected,
|
||||
featureFlags,
|
||||
type DownloadProgress,
|
||||
type PlacementPreview,
|
||||
} from "$lib/stores/app.svelte";
|
||||
@@ -702,7 +703,10 @@
|
||||
? Object.keys(topologyData()!.nodes).length
|
||||
: 1;
|
||||
const sharding = nodeCount <= 1 ? "Pipeline" : selectedSharding;
|
||||
const instanceType = nodeCount <= 1 ? "MlxRing" : selectedInstanceType;
|
||||
const instanceType =
|
||||
nodeCount <= 1 && selectedInstanceType === "MlxJaccl"
|
||||
? "MlxRing"
|
||||
: selectedInstanceType;
|
||||
try {
|
||||
const placementResponse = await fetch(
|
||||
`/instance/placement?model_id=${encodeURIComponent(modelId)}&sharding=${sharding}&instance_meta=${instanceType}&min_nodes=1`,
|
||||
@@ -783,6 +787,7 @@
|
||||
quantization?: string;
|
||||
base_model?: string;
|
||||
capabilities?: string[];
|
||||
requires_vllm?: boolean;
|
||||
}>
|
||||
>([]);
|
||||
type ModelMemoryFitStatus =
|
||||
@@ -886,7 +891,7 @@
|
||||
}
|
||||
|
||||
let selectedSharding = $state<"Pipeline" | "Tensor">("Pipeline");
|
||||
type InstanceMeta = "MlxRing" | "MlxJaccl";
|
||||
type InstanceMeta = "MlxRing" | "MlxJaccl" | "Vllm";
|
||||
|
||||
// Launch defaults persistence
|
||||
const LAUNCH_DEFAULTS_KEY = "exo-launch-defaults-v2";
|
||||
@@ -932,7 +937,12 @@
|
||||
// Apply sharding and instance type unconditionally
|
||||
selectedSharding = defaults.sharding;
|
||||
selectedInstanceType =
|
||||
defaults.instanceType === "MlxRing" ? "MlxRing" : "MlxJaccl";
|
||||
defaults.instanceType === "MlxRing"
|
||||
? "MlxRing"
|
||||
: defaults.instanceType === "Vllm"
|
||||
? "Vllm"
|
||||
: "MlxJaccl";
|
||||
userPickedInstanceType = true;
|
||||
|
||||
// Apply minNodes if valid (between 1 and maxNodes)
|
||||
if (
|
||||
@@ -954,6 +964,23 @@
|
||||
}
|
||||
|
||||
let selectedInstanceType = $state<InstanceMeta>("MlxRing");
|
||||
let userPickedInstanceType = $state(false);
|
||||
$effect(() => {
|
||||
if (!userPickedInstanceType && featureFlags()["vllm_available"]) {
|
||||
selectedInstanceType = "Vllm";
|
||||
}
|
||||
});
|
||||
const selectedModelRequiresVllm = $derived.by((): boolean => {
|
||||
const id = selectedPreviewModelId();
|
||||
if (!id) return false;
|
||||
const model = models.find((m) => m.id === id);
|
||||
return model?.requires_vllm === true;
|
||||
});
|
||||
$effect(() => {
|
||||
if (selectedModelRequiresVllm) {
|
||||
selectedInstanceType = "Vllm";
|
||||
}
|
||||
});
|
||||
let selectedMinNodes = $state<number>(1);
|
||||
let minNodesInitialized = $state(false);
|
||||
let launchingModelId = $state<string | null>(null);
|
||||
@@ -1146,9 +1173,7 @@
|
||||
}
|
||||
|
||||
const matchesSelectedRuntime = (runtime: InstanceMeta): boolean =>
|
||||
selectedInstanceType === "MlxRing"
|
||||
? runtime === "MlxRing"
|
||||
: runtime === "MlxJaccl";
|
||||
runtime === selectedInstanceType;
|
||||
|
||||
// Helper to check if a model can be launched (has valid placement with >= minNodes)
|
||||
function canModelFit(modelId: string): boolean {
|
||||
@@ -2063,6 +2088,7 @@
|
||||
let instanceType = "Unknown";
|
||||
if (instanceTag === "MlxRingInstance") instanceType = "MLX Ring";
|
||||
else if (instanceTag === "MlxJacclInstance") instanceType = "MLX RDMA";
|
||||
else if (instanceTag === "VllmInstance") instanceType = "vLLM";
|
||||
|
||||
const inst = instance as {
|
||||
shardAssignments?: {
|
||||
@@ -5769,14 +5795,18 @@
|
||||
</div>
|
||||
<div class="flex gap-2">
|
||||
<button
|
||||
disabled={selectedModelRequiresVllm}
|
||||
onclick={() => {
|
||||
if (selectedModelRequiresVllm) return;
|
||||
selectedInstanceType = "MlxRing";
|
||||
userPickedInstanceType = true;
|
||||
saveLaunchDefaults();
|
||||
}}
|
||||
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 cursor-pointer {selectedInstanceType ===
|
||||
'MlxRing'
|
||||
? 'bg-transparent text-exo-yellow border-exo-yellow'
|
||||
: 'bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
|
||||
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 {selectedModelRequiresVllm
|
||||
? 'opacity-40 cursor-not-allowed bg-transparent text-white/40 border-exo-medium-gray/30'
|
||||
: selectedInstanceType === 'MlxRing'
|
||||
? 'cursor-pointer bg-transparent text-exo-yellow border-exo-yellow'
|
||||
: 'cursor-pointer bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
|
||||
>
|
||||
<span
|
||||
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
|
||||
@@ -5792,14 +5822,18 @@
|
||||
TCP/IP
|
||||
</button>
|
||||
<button
|
||||
disabled={selectedModelRequiresVllm}
|
||||
onclick={() => {
|
||||
if (selectedModelRequiresVllm) return;
|
||||
selectedInstanceType = "MlxJaccl";
|
||||
userPickedInstanceType = true;
|
||||
saveLaunchDefaults();
|
||||
}}
|
||||
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 cursor-pointer {selectedInstanceType ===
|
||||
'MlxJaccl'
|
||||
? 'bg-transparent text-exo-yellow border-exo-yellow'
|
||||
: 'bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
|
||||
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 {selectedModelRequiresVllm
|
||||
? 'opacity-40 cursor-not-allowed bg-transparent text-white/40 border-exo-medium-gray/30'
|
||||
: selectedInstanceType === 'MlxJaccl'
|
||||
? 'cursor-pointer bg-transparent text-exo-yellow border-exo-yellow'
|
||||
: 'cursor-pointer bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
|
||||
>
|
||||
<span
|
||||
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
|
||||
@@ -5814,7 +5848,41 @@
|
||||
</span>
|
||||
RDMA (Fast)
|
||||
</button>
|
||||
{#if featureFlags()["vllm_available"] || selectedModelRequiresVllm}
|
||||
<button
|
||||
onclick={() => {
|
||||
selectedInstanceType = "Vllm";
|
||||
userPickedInstanceType = true;
|
||||
saveLaunchDefaults();
|
||||
}}
|
||||
class="flex items-center gap-2 py-1.5 px-3 text-xs font-mono border rounded transition-all duration-200 cursor-pointer {selectedInstanceType ===
|
||||
'Vllm'
|
||||
? 'bg-transparent text-exo-yellow border-exo-yellow'
|
||||
: 'bg-transparent text-white/70 border-exo-medium-gray/50 hover:border-exo-yellow/50'}"
|
||||
>
|
||||
<span
|
||||
class="w-3 h-3 rounded-full border-2 flex items-center justify-center {selectedInstanceType ===
|
||||
'Vllm'
|
||||
? 'border-exo-yellow'
|
||||
: 'border-exo-medium-gray'}"
|
||||
>
|
||||
{#if selectedInstanceType === "Vllm"}
|
||||
<span
|
||||
class="w-1.5 h-1.5 rounded-full bg-exo-yellow"
|
||||
></span>
|
||||
{/if}
|
||||
</span>
|
||||
vLLM (CUDA)
|
||||
</button>
|
||||
{/if}
|
||||
</div>
|
||||
{#if selectedModelRequiresVllm}
|
||||
<div
|
||||
class="mt-2 text-[11px] font-mono text-orange-300/80"
|
||||
>
|
||||
This model requires vLLM.
|
||||
</div>
|
||||
{/if}
|
||||
</div>
|
||||
|
||||
<!-- Minimum Devices -->
|
||||
|
||||
@@ -146,7 +146,7 @@
|
||||
config.treefmt.build.wrapper
|
||||
|
||||
# PYTHON
|
||||
self'.packages.editableVenv
|
||||
self'.packages.exo.passthru.evenv
|
||||
uv
|
||||
|
||||
# RUST
|
||||
|
||||
@@ -40,6 +40,19 @@ build-app: rust-rebuild sync-clean package
|
||||
xcodebuild build -project app/EXO/EXO.xcodeproj -scheme EXO -configuration Debug -derivedDataPath app/EXO/build
|
||||
@echo "\nBuild complete. Run with:\n open {{justfile_directory()}}/app/EXO/build/Build/Products/Debug/EXO.app"
|
||||
|
||||
sync-cuda:
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
uv sync --extra vllm-cuda13 --extra mlx-cpu --no-install-package vllm
|
||||
dest=".venv/lib/python3.13/site-packages"
|
||||
[[ -d $dest/vllm ]] || {
|
||||
nix build .#exo-cuda-13.passthru.evenv
|
||||
# will also grab vllm-0.19.1-distinfo
|
||||
cp -aL result/lib/python3.13/site-packages/vllm* .venv/lib/python3.13/site-packages
|
||||
chmod -R u+rwX .venv/lib/python3.13/site-packages/vllm*
|
||||
rm result
|
||||
}
|
||||
|
||||
clean:
|
||||
rm -rf **/__pycache__
|
||||
rm -rf target/
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
diff --git a/setup.py b/setup.py
|
||||
index 6dc2ed028..bdcc6354a 100644
|
||||
--- a/setup.py
|
||||
+++ b/setup.py
|
||||
@@ -18,6 +18,13 @@ from setuptools import Extension, setup
|
||||
from setuptools.command.build_ext import build_ext
|
||||
|
||||
|
||||
+if "NIX_ATTRS_JSON_FILE" in os.environ:
|
||||
+ with open(os.environ["NIX_ATTRS_JSON_FILE"], "r") as f:
|
||||
+ NIX_ATTRS = json.load(f)
|
||||
+else:
|
||||
+ NIX_ATTRS = { "cmakeFlags": os.environ.get("cmakeFlags", "").split() }
|
||||
+
|
||||
+
|
||||
def load_module_from_path(module_name, path):
|
||||
spec = importlib.util.spec_from_file_location(module_name, path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
@@ -184,6 +191,7 @@ class cmake_build_ext(build_ext):
|
||||
cmake_args = [
|
||||
"-DCMAKE_BUILD_TYPE={}".format(cfg),
|
||||
"-DVLLM_TARGET_DEVICE={}".format(VLLM_TARGET_DEVICE),
|
||||
+ *NIX_ATTRS["cmakeFlags"],
|
||||
]
|
||||
|
||||
verbose = envs.VERBOSE
|
||||
+57
-36
@@ -15,21 +15,18 @@ dependencies = [
|
||||
"huggingface-hub>=1.8.0",
|
||||
"psutil>=7.0.0",
|
||||
"loguru>=0.7.3",
|
||||
"exo-pyo3-bindings", # rust bindings
|
||||
"exo-pyo3-bindings", # rust bindings
|
||||
"anyio==4.11.0",
|
||||
"mlx==0.31.2; sys_platform == 'darwin'",
|
||||
"mlx-lm; sys_platform=='darwin'",
|
||||
"tiktoken>=0.12.0", # required for kimi k2 tokenizer
|
||||
"tiktoken>=0.12.0", # required for kimi k2 tokenizer
|
||||
"hypercorn>=0.18.0",
|
||||
"openai-harmony>=0.0.8",
|
||||
"httpx>=0.28.1",
|
||||
"tomlkit>=0.14.0",
|
||||
"mflux==0.17.2; sys_platform == 'darwin'",
|
||||
"python-multipart>=0.0.21",
|
||||
"msgspec>=0.19.0",
|
||||
"zstandard>=0.23.0",
|
||||
"mlx-vlm>=0.3.11; sys_platform == 'darwin'",
|
||||
"transformers>=5.6.2",
|
||||
"nvidia-ml-py>=13.595.45",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
@@ -48,26 +45,30 @@ dev = [
|
||||
|
||||
[project.optional-dependencies]
|
||||
build = ["nanobind"]
|
||||
cpu = [
|
||||
"mlx==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-cpu==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-lm; sys_platform == 'linux'",
|
||||
"mlx-vlm>=0.3.11; sys_platform== 'linux'",
|
||||
"torch>=2.10.0; sys_platform == 'linux'",
|
||||
mlx-none = ["anyio"]
|
||||
mlx = [
|
||||
"mlx==0.31.2",
|
||||
"mlx-lm",
|
||||
"mlx-vlm>=0.3.11",
|
||||
"mflux==0.17.5",
|
||||
# pinning vllms versions for consistency.
|
||||
"torch==2.10.0; sys_platform == 'darwin'",
|
||||
"torch==2.10.0; sys_platform == 'linux'",
|
||||
"torchaudio==2.10.0; sys_platform == 'darwin'",
|
||||
"torchaudio==2.10.0; sys_platform == 'linux'",
|
||||
"torchvision==0.25.0; sys_platform == 'darwin'",
|
||||
"torchvision==0.25.0; sys_platform == 'linux'",
|
||||
|
||||
]
|
||||
cuda12 = [
|
||||
"mlx==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-cuda-12==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-lm; sys_platform == 'linux'",
|
||||
"mlx-vlm>=0.3.11; 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'",
|
||||
"mlx-vlm>=0.3.11; 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'",
|
||||
]
|
||||
|
||||
###
|
||||
@@ -81,12 +82,23 @@ members = ["rust/exo_pyo3_bindings", "bench"]
|
||||
exo-pyo3-bindings = { workspace = true }
|
||||
mlx = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git", branch = "address-rdma-gpu-locks", marker = "sys_platform == 'darwin'" }
|
||||
mlx-lm = { git = "https://github.com/rltakashige/mlx-lm", branch = "leo/deepseek-v4" }
|
||||
mflux = { git = "http://github.com/evanev7/mflux", branch = "exo" }
|
||||
vllm = { git = "http://github.com/evanev7/vllm", branch = "exo2" }
|
||||
torch = [
|
||||
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and extra == 'cuda13' and extra != 'cpu' and extra != 'cuda12'" },
|
||||
{ index = "pytorch-cu120", marker = "sys_platform == 'linux' and extra == 'cuda12' and extra != 'cpu' and extra != 'cuda13'" },
|
||||
{ index = "pytorch-cpu", marker = "(extra != 'cuda12' and extra != 'cuda13' and sys_platform == 'linux') or sys_platform == 'darwin'" },
|
||||
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and extra == 'mlx-cpu' and extra != 'vllm-cuda13' and extra != 'mlx-cuda13' and extra != 'mlx-cuda12'" },
|
||||
{ index = "pytorch-cu128", marker = "sys_platform == 'linux' and extra == 'mlx-cuda12' and extra != 'mlx-cuda13' and extra != 'vllm-cuda13'" },
|
||||
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and (extra == 'mlx-cuda13' or extra == 'vllm-cuda13')" },
|
||||
]
|
||||
torchvision = [
|
||||
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and extra == 'mlx-cpu' and extra != 'vllm-cuda13' and extra != 'mlx-cuda13' and extra != 'mlx-cuda12'" },
|
||||
{ index = "pytorch-cu128", marker = "sys_platform == 'linux' and extra == 'mlx-cuda12' and extra != 'mlx-cuda13' and extra != 'vllm-cuda13'" },
|
||||
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and (extra == 'mlx-cuda13' or extra == 'vllm-cuda13')" },
|
||||
]
|
||||
torchaudio = [
|
||||
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and extra == 'mlx-cpu' and extra != 'vllm-cuda13' and extra != 'mlx-cuda13' and extra != 'mlx-cuda12'" },
|
||||
{ index = "pytorch-cu128", marker = "sys_platform == 'linux' and extra == 'mlx-cuda12' and extra != 'mlx-cuda13' and extra != 'vllm-cuda13'" },
|
||||
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and (extra == 'mlx-cuda13' or extra == 'vllm-cuda13')" },
|
||||
]
|
||||
vllm = { git = "https://github.com/hmellor/vllm.git", branch = "transformers-v5" }
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cu130"
|
||||
@@ -94,8 +106,8 @@ url = "https://download.pytorch.org/whl/cu130"
|
||||
explicit = true
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cu120"
|
||||
url = "https://download.pytorch.org/whl/cu120"
|
||||
name = "pytorch-cu128"
|
||||
url = "https://download.pytorch.org/whl/cu128"
|
||||
explicit = true
|
||||
|
||||
[[tool.uv.index]]
|
||||
@@ -156,11 +168,19 @@ root = "src"
|
||||
required-version = ">=0.8.6"
|
||||
prerelease = "allow"
|
||||
environments = ["sys_platform == 'darwin'", "sys_platform == 'linux'"]
|
||||
conflicts = [[{ extra = "cuda12" }, { extra = "cuda13" }, { extra = "cpu" }]]
|
||||
constraint-dependencies = ["transformers>=5.6.2"]
|
||||
override-dependencies = [
|
||||
"mlx==0.31.1; sys_platform=='linux'",
|
||||
"mlx; sys_platform=='darwin'",
|
||||
override-dependencies = ["opencv-python; python_version < '0'"]
|
||||
conflicts = [
|
||||
[
|
||||
{ extra = "mlx-cuda13" },
|
||||
{ extra = "mlx-cuda12" },
|
||||
{ extra = "mlx-cpu" },
|
||||
{ extra = "mlx-none" },
|
||||
],
|
||||
[
|
||||
{ extra = "vllm-cuda13" },
|
||||
{ extra = "mlx-cuda12" },
|
||||
{ extra = "vllm-none" },
|
||||
],
|
||||
]
|
||||
|
||||
[tool.uv.extra-build-dependencies]
|
||||
@@ -175,6 +195,7 @@ mlx = [
|
||||
"ninja",
|
||||
]
|
||||
mlx-lm = ["setuptools"]
|
||||
mflux = ["uv_build"]
|
||||
xgrammar = [
|
||||
"nanobind",
|
||||
"setuptools",
|
||||
|
||||
+202
-20
@@ -10,8 +10,10 @@ let
|
||||
inherit (pkgs.stdenv.hostPlatform) isLinux isDarwin isx86_64;
|
||||
inherit (pkgs.config) cudaSupport;
|
||||
inherit (pkgs) cudaPackages;
|
||||
cuda13Support = cudaSupport && cudaPackages.cudaMajorVersion == "13";
|
||||
libmlx_source = if cuda13Support then "mlx-cuda-13" else if cudaSupport then "mlx-cuda-12" else "mlx-cpu";
|
||||
libmlx_source =
|
||||
if (builtins.elem "mlx-cuda13" members.exo or [ ]) then "mlx-cuda-13"
|
||||
else if (builtins.elem "mlx-cuda12" members.exo or [ ]) then "mlx-cuda-12"
|
||||
else "mlx-cpu";
|
||||
python = pkgs.python313;
|
||||
cudaLibs = with cudaPackages; [
|
||||
cuda_cudart
|
||||
@@ -113,37 +115,213 @@ let
|
||||
});
|
||||
} // lib.optionalAttrs isLinux {
|
||||
mlx = prev.mlx.overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ lib.optionals cudaSupport [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ lib.optionals cudaSupport cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = lib.optionals cudaSupport [ "libcuda.so.1" ];
|
||||
postInstall = ''
|
||||
cp -r "${final.${libmlx_source}}/${final.python.sitePackages}/mlx" "$out/${final.python.sitePackages}/mlx/"
|
||||
'';
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
} // lib.optionalAttrs cudaSupport {
|
||||
"${libmlx_source}" = prev."${libmlx_source}".overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cufile = prev.nvidia-cufile.overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ [ pkgs.rdma-core ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cusolver = prev.nvidia-cusolver.overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-nvshmem-cu13 = prev.nvidia-nvshmem-cu13.overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ [ pkgs.rdma-core pkgs.pmix pkgs.libfabric pkgs.ucx pkgs.openmpi ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cusparse = prev.nvidia-cusparse.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ [ cudaLibs ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
});
|
||||
torch = prev.torch.overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
torchaudio = prev.torchaudio.overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
buildInputs = old.buildInputs ++ [ cudaPackages.cuda_cudart ];
|
||||
preFixup = "addAutoPatchelfSearchPath '${final.torch}'";
|
||||
});
|
||||
torchvision = prev.torchvision.overrideAttrs (old: {
|
||||
nativeBuildInputs = old.nativeBuildInputs ++ [ pkgs.autoAddDriverRunpath ];
|
||||
preFixup = "addAutoPatchelfSearchPath '${final.torch}'";
|
||||
});
|
||||
|
||||
torch-c-dlpack-ext = prev.torch-c-dlpack-ext.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
preFixup = "addAutoPatchelfSearchPath '${final.torch}'";
|
||||
});
|
||||
# Currently treating vllm as a cuda dep. it obviously exists as a non cuda dep
|
||||
vllm = prev.vllm.overrideAttrs (old:
|
||||
let
|
||||
cuda_cccl_compat = pkgs.runCommand "cuda-cccl-compat" { } ''
|
||||
mkdir -p $out/include
|
||||
ln -s ${cudaPackages.cuda_cccl}/include $out/include/cccl
|
||||
'';
|
||||
|
||||
cudaRoot = pkgs.symlinkJoin {
|
||||
name = "cuda-merged-exo";
|
||||
paths = builtins.concatMap (p: [ (lib.getBin p) (lib.getLib p) (lib.getDev p) ]) (cudaLibs ++ [ cudaPackages.cuda_nvcc cuda_cccl_compat ]);
|
||||
};
|
||||
|
||||
cutlass = pkgs.fetchFromGitHub {
|
||||
name = "cutlass-source";
|
||||
owner = "NVIDIA";
|
||||
repo = "cutlass";
|
||||
tag = "v4.2.1";
|
||||
hash = "sha256-iP560D5Vwuj6wX1otJhwbvqe/X4mYVeKTpK533Wr5gY=";
|
||||
};
|
||||
triton-kernels = pkgs.fetchFromGitHub {
|
||||
owner = "triton-lang";
|
||||
repo = "triton";
|
||||
tag = "v3.6.0";
|
||||
hash = "sha256-JFSpQn+WsNnh7CAPlcpOcUp0nyKXNbJEANdXqmkt4Tc=";
|
||||
};
|
||||
|
||||
cutlass-flashmla = pkgs.fetchFromGitHub {
|
||||
owner = "NVIDIA";
|
||||
repo = "cutlass";
|
||||
rev = "147f5673d0c1c3dcf66f78d677fd647e4a020219";
|
||||
hash = "sha256-dHQto08IwTDOIuFUp9jwm1MWkFi8v2YJ/UESrLuG71g=";
|
||||
};
|
||||
|
||||
flashmla = pkgs.stdenv.mkDerivation {
|
||||
pname = "flashmla";
|
||||
version = "1.0.0";
|
||||
|
||||
src = pkgs.fetchFromGitHub {
|
||||
name = "FlashMLA-source";
|
||||
owner = "vllm-project";
|
||||
repo = "FlashMLA";
|
||||
rev = "c2afa9cb93e674d5a9120a170a6da57b89267208";
|
||||
hash = "sha256-pKlwxV6G9iHag/jbu3bAyvYvnu5TbrQwUMFV0AlGC3s=";
|
||||
};
|
||||
|
||||
dontConfigure = true;
|
||||
|
||||
buildPhase = ''
|
||||
rm -rf csrc/cutlass
|
||||
ln -sf ${cutlass-flashmla} csrc/cutlass
|
||||
'';
|
||||
|
||||
installPhase = ''
|
||||
cp -rva . $out
|
||||
'';
|
||||
};
|
||||
qutlass = pkgs.fetchFromGitHub {
|
||||
name = "qutlass-source";
|
||||
owner = "IST-DASLab";
|
||||
repo = "qutlass";
|
||||
rev = "830d2c4537c7396e14a02a46fbddd18b5d107c65";
|
||||
hash = "sha256-aG4qd0vlwP+8gudfvHwhtXCFmBOJKQQTvcwahpEqC84=";
|
||||
};
|
||||
vllm-flash-attn = pkgs.stdenv.mkDerivation {
|
||||
pname = "vllm-flash-attn";
|
||||
version = "2.7.2.post1";
|
||||
|
||||
src = pkgs.fetchFromGitHub {
|
||||
name = "flash-attention-source";
|
||||
owner = "vllm-project";
|
||||
repo = "flash-attention";
|
||||
rev = "188be16520ceefdc625fdf71365585d2ee348fe2";
|
||||
hash = "sha256-Osec+/IF3+UDtbIhDMBXzUeWJ7hDJNb5FpaVaziPSgM=";
|
||||
};
|
||||
|
||||
patches = [
|
||||
(pkgs.fetchpatch {
|
||||
url = "https://github.com/Dao-AILab/flash-attention/commit/dad67c88d4b6122c69d0bed1cebded0cded71cea.patch";
|
||||
hash = "sha256-JSgXWItOp5KRpFbTQj/cZk+Tqez+4mEz5kmH5EUeQN4=";
|
||||
})
|
||||
(pkgs.fetchpatch {
|
||||
url = "https://github.com/Dao-AILab/flash-attention/commit/e26dd28e487117ee3e6bc4908682f41f31e6f83a.patch";
|
||||
hash = "sha256-NkCEowXSi+tiWu74Qt+VPKKavx0H9JeteovSJKToK9A=";
|
||||
})
|
||||
];
|
||||
|
||||
dontConfigure = true;
|
||||
|
||||
buildPhase = ''
|
||||
rm -rf csrc/cutlass
|
||||
ln -sf ${cutlass} csrc/cutlass
|
||||
'';
|
||||
|
||||
installPhase = ''
|
||||
cp -rva . $out
|
||||
'';
|
||||
};
|
||||
in
|
||||
{
|
||||
patches = (old.patches or [ ]) ++ [ ../nix/vllm-setuppy-cmake.patch ];
|
||||
nativeBuildInputs = (old.nativeBuildInputs or [ ]) ++ [
|
||||
pkgs.cmake
|
||||
pkgs.ninja
|
||||
pkgs.autoAddDriverRunpath
|
||||
] ++ lib.optionals cudaSupport [
|
||||
cudaPackages.cuda_nvcc
|
||||
];
|
||||
# TODO: vllm rocm/cpu
|
||||
VLLM_TARGET_DEVICE = "empty";
|
||||
preConfigure = ''
|
||||
export MAX_JOBS="$NIX_BUILD_CORES"
|
||||
'';
|
||||
|
||||
# TODO: vllm non cuda13 support, more arch's, etc.
|
||||
} // lib.optionalAttrs cudaSupport {
|
||||
buildInputs = cudaLibs ++ [ cudaRoot ];
|
||||
|
||||
VLLM_CUDA_VERSION = cudaPackages.cudaMajorMinorVersion;
|
||||
CUDA_HOME = "${cudaRoot}";
|
||||
CUDAToolkit_ROOT = "${cudaRoot}";
|
||||
CUDACXX = "${cudaRoot}/bin/nvcc";
|
||||
VLLM_CUTLASS_SRC_DIR = "${lib.getDev cutlass}";
|
||||
VLLM_TARGET_DEVICE = "cuda";
|
||||
TORCH_CUDA_ARCH_LIST = "12.0;12.1";
|
||||
TRITON_KERNELS_SRC_DIR = "${lib.getDev triton-kernels}/python/triton_kernels/triton_kernels";
|
||||
FLASH_MLA_SRC_DIR = "${lib.getDev flashmla}";
|
||||
QUTLASS_SRC_DIR = "${lib.getDev qutlass}";
|
||||
VLLM_FLASH_ATTN_SRC_DIR = "${lib.getDev vllm-flash-attn}";
|
||||
CAFFE2_USE_CUDNN = "ON";
|
||||
CAFFE2_USE_CUFILE = "ON";
|
||||
CUTLASS_ENABLE_CUBLAS = "ON";
|
||||
CUTLASS_NVCC_ARCHS_ENABLED = "12.0;12.1";
|
||||
|
||||
cmakeFlags = [
|
||||
(lib.cmakeBool "CMAKE_SKIP_INSTALL_RPATH" true)
|
||||
(lib.cmakeBool "CMAKE_BUILD_WITH_INSTALL_RPATH" true)
|
||||
(lib.cmakeFeature "CUDA_HOME" "${cudaRoot}")
|
||||
(lib.cmakeFeature "CUDAToolkit_ROOT" "${cudaRoot}")
|
||||
(lib.cmakeFeature "CMAKE_CUDA_COMPILER" "${cudaRoot}/bin/nvcc")
|
||||
(lib.cmakeFeature "CMAKE_PREFIX_PATH" "${cudaRoot}")
|
||||
(lib.cmakeFeature "FETCHCONTENT_SOURCE_DIR_CUTLASS" "${lib.getDev cutlass}")
|
||||
(lib.cmakeFeature "FLASH_MLA_SRC_DIR" "${lib.getDev flashmla}")
|
||||
(lib.cmakeFeature "VLLM_FLASH_ATTN_SRC_DIR" "${lib.getDev vllm-flash-attn}")
|
||||
(lib.cmakeFeature "QUTLASS_SRC_DIR" "${lib.getDev qutlass}")
|
||||
(lib.cmakeFeature "TORCH_CUDA_ARCH_LIST" "12.0;12.1")
|
||||
(lib.cmakeFeature "CUTLASS_NVCC_ARCHS_ENABLED" "${cudaPackages.flags.cmakeCudaArchitecturesString}")
|
||||
(lib.cmakeFeature "CUDA_TOOLKIT_ROOT_DIR" "${cudaRoot}")
|
||||
(lib.cmakeFeature "CAFFE2_USE_CUDNN" "ON")
|
||||
(lib.cmakeFeature "CAFFE2_USE_CUFILE" "ON")
|
||||
(lib.cmakeFeature "CUTLASS_ENABLE_CUBLAS" "ON")
|
||||
];
|
||||
});
|
||||
|
||||
} // lib.optionalAttrs (cudaSupport && isx86_64) {
|
||||
numba = prev.numba.overrideAttrs (old: {
|
||||
buildInputs = (old.buildInputs or [ ]) ++ [ pkgs.tbb ];
|
||||
});
|
||||
};
|
||||
pyprojectOverlay = workspace.mkPyprojectOverlay {
|
||||
sourcePreference = "wheel";
|
||||
@@ -164,24 +342,28 @@ let
|
||||
buildSystemsOverlay
|
||||
]
|
||||
);
|
||||
venv = name: (pythonSet.mkVirtualEnv "${name}-env" members).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" ]; });
|
||||
mkApp = cmd: name: pkgs.writeShellApplication {
|
||||
# mlx and mlx-cuda ship clashing cmake files - we dont need them at runtime anyway
|
||||
venv = name: (pythonSet.mkVirtualEnv "${name}-venv" members).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" "lib/python${python.pythonVersion}/site-packages/build_backend.py" ]; });
|
||||
mkApp = text: name: pkgs.writeShellApplication {
|
||||
inherit name;
|
||||
text = "exec " + lib.optionalString cudaSupport "nixglhost " + text;
|
||||
runtimeEnv = {
|
||||
EXO_DASHBOARD_DIR = self'.packages.dashboard;
|
||||
EXO_RESOURCES_DIR = inputs.self + /resources;
|
||||
};
|
||||
runtimeInputs = [
|
||||
# mlx and mlx-cuda ship clashing cmake files - we dont need them at runtime anyway
|
||||
(venv name)
|
||||
pkgs.nix-gl-host
|
||||
]
|
||||
++ lib.optionals isDarwin [ pkgs.macmon ];
|
||||
text = "exec " + lib.optionalString cudaSupport "${lib.getExe pkgs.nix-gl-host} " + cmd;
|
||||
passthru = {
|
||||
venv = venv name;
|
||||
evenv = ((pythonSet.overrideScope editableOverlay).mkVirtualEnv "${name}-evenv" (members // { exo = (members.exo or [ ]) ++ [ "dev" ]; })).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" "lib/python${python.pythonVersion}/site-packages/build_backend.py" ]; });
|
||||
};
|
||||
};
|
||||
in
|
||||
{
|
||||
inherit venv;
|
||||
editablePythonSet = pythonSet.overrideScope editableOverlay;
|
||||
mkPythonScript = path: mkApp ''python ${path} "$@"'';
|
||||
mkExo = mkApp ''exo "$@"'';
|
||||
};
|
||||
@@ -191,18 +373,18 @@ in
|
||||
{ self', pkgs, unfreePkgs, lib, ... }:
|
||||
let
|
||||
inherit (pkgs.stdenv.hostPlatform) isLinux;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "cpu" ]; }; }) editablePythonSet mkExo;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "mlx-cpu" "vllm-none" ]; }; }) mkExo;
|
||||
|
||||
# Virtual environment with dev dependencies for testing
|
||||
testVenv = (mkPythonSet {
|
||||
inherit self' pkgs lib; members = {
|
||||
exo = [ "dev" "cpu" ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
exo = [ "dev" "mlx-cpu" "vllm-none" ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
};
|
||||
}).venv "exo-test";
|
||||
|
||||
mkBenchScript = (mkPythonSet {
|
||||
inherit self' pkgs lib; members = {
|
||||
exo = [ "cpu" ];
|
||||
exo = [ "mlx-cpu" "vllm-none" ];
|
||||
exo-bench = [ ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
};
|
||||
}).mkPythonScript;
|
||||
@@ -212,12 +394,12 @@ in
|
||||
runtimeInputs = [ pkgs.python313 ];
|
||||
text = ''exec python ${path} "$@"'';
|
||||
};
|
||||
|
||||
cuda12Set = mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; members = { exo = [ "mlx-cuda12" "vllm-none" ]; }; };
|
||||
cuda13Set = mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; members = { exo = [ "mlx-cpu" "vllm-cuda13" ]; }; };
|
||||
in
|
||||
{
|
||||
packages = {
|
||||
exo = mkExo "exo";
|
||||
editableVenv = editablePythonSet.mkVirtualEnv "exo-dev-env" { exo = [ "dev" ]; };
|
||||
# for running tests in ci
|
||||
exo-test-env = testVenv;
|
||||
exo-bench = mkBenchScript "exo-bench" (inputs.self + /bench/exo_bench.py);
|
||||
@@ -226,8 +408,8 @@ in
|
||||
# used by ./tests/run_exo_on.sh
|
||||
exo-get-all-models-on-cluster = mkSimplePythonScript "exo-get-all-models-on-cluster" (inputs.self + /tests/get_all_models_on_cluster.py);
|
||||
} // lib.optionalAttrs isLinux {
|
||||
exo-cuda-12 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; members = { exo = [ "cuda12" ]; }; }).mkExo "exo-cuda-12";
|
||||
exo-cuda-13 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; members = { exo = [ "cuda13" ]; }; }).mkExo "exo-cuda-13";
|
||||
exo-cuda-12 = cuda12Set.mkExo "exo-cuda-12";
|
||||
exo-cuda-13 = cuda13Set.mkExo "exo-cuda-13";
|
||||
};
|
||||
|
||||
checks = {
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
model_id = "2imi9/gpt-oss-20B-NVFP4A16-BF16"
|
||||
n_layers = 24
|
||||
hidden_size = 2880
|
||||
num_key_value_heads = 8
|
||||
supports_tensor = false
|
||||
tasks = ["TextGeneration"]
|
||||
family = "gpt-oss"
|
||||
quantization = "nvfp4"
|
||||
base_model = "GPT-OSS 20B"
|
||||
capabilities = ["text", "thinking"]
|
||||
reasoning_dialect = "channel"
|
||||
context_length = 131072
|
||||
requires_vllm = true
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 41829514752
|
||||
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 1.0
|
||||
top_k = 0
|
||||
@@ -0,0 +1,27 @@
|
||||
model_id = "nvidia/Qwen3-30B-A3B-NVFP4"
|
||||
n_layers = 48
|
||||
hidden_size = 2048
|
||||
num_key_value_heads = 4
|
||||
supports_tensor = false
|
||||
tasks = ["TextGeneration"]
|
||||
family = "qwen"
|
||||
quantization = "nvfp4"
|
||||
base_model = "Qwen3 30B"
|
||||
capabilities = ["text", "thinking", "thinking_toggle"]
|
||||
context_length = 32768
|
||||
requires_vllm = true
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 18087458688
|
||||
|
||||
[sampling_defaults]
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
|
||||
[sampling_defaults.non_thinking]
|
||||
temperature = 0.7
|
||||
top_p = 0.8
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
@@ -0,0 +1,20 @@
|
||||
model_id = "openai/gpt-oss-120b"
|
||||
n_layers = 36
|
||||
hidden_size = 2880
|
||||
num_key_value_heads = 8
|
||||
supports_tensor = false
|
||||
tasks = ["TextGeneration"]
|
||||
family = "gpt-oss"
|
||||
quantization = "mxfp4"
|
||||
base_model = "GPT-OSS 120B"
|
||||
capabilities = ["text", "thinking"]
|
||||
reasoning_dialect = "channel"
|
||||
context_length = 131072
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 65248815744
|
||||
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 1.0
|
||||
top_k = 0
|
||||
@@ -0,0 +1,32 @@
|
||||
model_id = "sakamakismile/Qwen3.6-27B-NVFP4"
|
||||
n_layers = 64
|
||||
hidden_size = 5120
|
||||
num_key_value_heads = 4
|
||||
supports_tensor = false
|
||||
tasks = ["TextGeneration"]
|
||||
family = "qwen"
|
||||
quantization = "nvfp4"
|
||||
base_model = "Qwen3.6 27B"
|
||||
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
|
||||
reasoning_dialect = "post_last_user"
|
||||
context_length = 262144
|
||||
requires_vllm = true
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 16703361232
|
||||
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
|
||||
[sampling_defaults.non_thinking]
|
||||
temperature = 0.7
|
||||
top_p = 0.8
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
@@ -0,0 +1,222 @@
|
||||
#!/usr/bin/env python
|
||||
"""Standalone smoke test for VllmEngine.serve_prefill.
|
||||
|
||||
Loads a real vLLM engine, runs serve_prefill against an in-memory buffer
|
||||
twice in a row with the same prompt, and verifies both runs produce a
|
||||
well-formed wire stream (header -> KV chunks -> Done).
|
||||
|
||||
The second run is the regression guard: with vLLM APC enabled this would
|
||||
trip the chunked-prefill + APC + custom kv-connector CUDA assert
|
||||
(`vectorized_gather_kernel: ind >= ind_dim_size`) and the server would
|
||||
close the socket before the Done frame.
|
||||
|
||||
Usage on the Spark (gx10-de89):
|
||||
|
||||
cd /home/larry/exo
|
||||
/nix/store/2b82iz9ac0pxqafrgxmgdkq8sr2hwlx6-exo-cuda-13-venv/bin/python \\
|
||||
scripts/check_serve_prefill.py Qwen/Qwen3-0.6B
|
||||
|
||||
Exits 0 on success, non-zero with a diagnostic on failure.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import io
|
||||
import os
|
||||
import sys
|
||||
import traceback
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
|
||||
def _ensure_repo_on_path() -> None:
|
||||
repo = Path(__file__).resolve().parent.parent
|
||||
src = repo / "src"
|
||||
if str(src) not in sys.path:
|
||||
sys.path.insert(0, str(src))
|
||||
|
||||
|
||||
_ensure_repo_on_path()
|
||||
|
||||
from exo.shared.types.common import ModelId # noqa: E402
|
||||
from exo.worker.disaggregated.protocol import ( # noqa: E402
|
||||
ArraysState,
|
||||
Done,
|
||||
ErrorMessage,
|
||||
KVChunk,
|
||||
read_header,
|
||||
read_message,
|
||||
)
|
||||
from exo.worker.disaggregated.server import PrefillRequest # noqa: E402
|
||||
|
||||
|
||||
def _make_token_ids(n: int) -> list[int]:
|
||||
return [(i * 1009 + 17) % 30000 + 100 for i in range(n)]
|
||||
|
||||
|
||||
def _decode(
|
||||
payload: bytes,
|
||||
) -> tuple[list[KVChunk], list[ArraysState], Done | None, ErrorMessage | None]:
|
||||
buf = io.BytesIO(payload)
|
||||
_ = read_header(buf)
|
||||
chunks: list[KVChunk] = []
|
||||
arrays: list[ArraysState] = []
|
||||
done: Done | None = None
|
||||
error: ErrorMessage | None = None
|
||||
while True:
|
||||
msg = read_message(buf)
|
||||
if msg is None:
|
||||
break
|
||||
if isinstance(msg, KVChunk):
|
||||
chunks.append(msg)
|
||||
elif isinstance(msg, ArraysState):
|
||||
arrays.append(msg)
|
||||
elif isinstance(msg, Done):
|
||||
done = msg
|
||||
break
|
||||
elif isinstance(msg, ErrorMessage):
|
||||
error = msg
|
||||
break
|
||||
return chunks, arrays, done, error
|
||||
|
||||
|
||||
def _build_engine(model_id: ModelId) -> object:
|
||||
from exo.worker.engines.vllm.engine import VllmEngine
|
||||
from exo.worker.engines.vllm.generator import VllmBatchEngine, load_vllm_engine
|
||||
from exo.worker.engines.vllm.kv_connector import (
|
||||
ExoKVProducerConnector,
|
||||
_patch_gdn_capture,
|
||||
_patch_vllm_for_connector,
|
||||
)
|
||||
|
||||
_patch_vllm_for_connector(ExoKVProducerConnector)
|
||||
_patch_gdn_capture()
|
||||
|
||||
llm_engine, tool_parser = load_vllm_engine(
|
||||
model_id=model_id,
|
||||
trust_remote_code=False,
|
||||
n_layers=1,
|
||||
kv_connector_cls=ExoKVProducerConnector,
|
||||
)
|
||||
gen = VllmBatchEngine(engine=llm_engine, model_id=model_id)
|
||||
|
||||
class _S:
|
||||
def send(self, _: object) -> None: ...
|
||||
|
||||
class _R:
|
||||
def collect(self) -> list[object]:
|
||||
return []
|
||||
|
||||
return VllmEngine(
|
||||
tool_parser=tool_parser,
|
||||
model_id=model_id,
|
||||
cancel_receiver=cast("object", _R()), # pyright: ignore[reportArgumentType]
|
||||
event_sender=cast("object", _S()), # pyright: ignore[reportArgumentType]
|
||||
_gen=gen,
|
||||
max_concurrent_requests=1,
|
||||
)
|
||||
|
||||
|
||||
def _run_one(engine: object, n_tokens: int, label: str) -> int:
|
||||
request = PrefillRequest(
|
||||
request_id=f"check-{label}-{os.getpid()}",
|
||||
model_id="ignored",
|
||||
token_ids=_make_token_ids(n_tokens),
|
||||
start_pos=0,
|
||||
use_prefix_cache=True,
|
||||
)
|
||||
buf = io.BytesIO()
|
||||
engine.serve_prefill(request, buf) # pyright: ignore[reportAttributeAccessIssue]
|
||||
payload = buf.getvalue()
|
||||
if not payload:
|
||||
raise AssertionError(f"{label}: server wrote nothing")
|
||||
|
||||
chunks, arrays, done, error = _decode(payload)
|
||||
if error is not None:
|
||||
raise AssertionError(
|
||||
f"{label}: server returned ErrorMessage [{error.code}]: {error.message}"
|
||||
)
|
||||
if done is None:
|
||||
raise AssertionError(
|
||||
f"{label}: stream did not end with Done "
|
||||
f"({len(chunks)} kv chunks, {len(arrays)} arrays)"
|
||||
)
|
||||
if done.total_tokens <= 0:
|
||||
raise AssertionError(f"{label}: Done reported {done.total_tokens} tokens")
|
||||
if not chunks:
|
||||
raise AssertionError(f"{label}: no KV chunks shipped")
|
||||
|
||||
expected = max(0, n_tokens - 2)
|
||||
if done.total_tokens < expected - 64:
|
||||
raise AssertionError(
|
||||
f"{label}: got {done.total_tokens} tokens, expected ~{expected}"
|
||||
)
|
||||
print(
|
||||
f" [{label}] OK: tokens={done.total_tokens} "
|
||||
f"kv_chunks={len(chunks)} arrays={len(arrays)}"
|
||||
)
|
||||
return done.total_tokens
|
||||
|
||||
|
||||
def main(argv: list[str]) -> int:
|
||||
if len(argv) < 2:
|
||||
print(__doc__)
|
||||
return 2
|
||||
model_id = ModelId(argv[1])
|
||||
|
||||
from exo.download.download_utils import build_model_path
|
||||
|
||||
model_path = build_model_path(model_id)
|
||||
if not model_path.exists():
|
||||
print(f"FAIL: model {model_id} not found at {model_path}")
|
||||
return 1
|
||||
print(f"Loading vLLM engine for {model_id} ({model_path}) ...")
|
||||
|
||||
engine = _build_engine(model_id)
|
||||
failures: list[str] = []
|
||||
try:
|
||||
try:
|
||||
t1 = _run_one(engine, n_tokens=512, label="run1-fresh")
|
||||
except AssertionError as e:
|
||||
failures.append(f"run1: {e}")
|
||||
t1 = 0
|
||||
try:
|
||||
t2 = _run_one(engine, n_tokens=512, label="run2-same-prompt")
|
||||
except AssertionError as e:
|
||||
failures.append(f"run2: {e}")
|
||||
t2 = 0
|
||||
if t1 and t2 and t1 != t2:
|
||||
failures.append(
|
||||
f"run1 returned {t1} tokens but run2 returned {t2} (should match)"
|
||||
)
|
||||
try:
|
||||
ta = _run_one(engine, n_tokens=256, label="run3-shorter")
|
||||
tb = _run_one(engine, n_tokens=768, label="run4-longer")
|
||||
if ta and tb and tb <= ta:
|
||||
failures.append(
|
||||
f"longer prompt should produce more tokens: 256->{ta} 768->{tb}"
|
||||
)
|
||||
except AssertionError as e:
|
||||
failures.append(f"length-variation: {e}")
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
engine.close() # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
if failures:
|
||||
print()
|
||||
print("FAIL")
|
||||
for f in failures:
|
||||
print(f" - {f}")
|
||||
return 1
|
||||
print()
|
||||
print("PASS")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
sys.exit(main(sys.argv))
|
||||
except Exception:
|
||||
traceback.print_exc()
|
||||
sys.exit(1)
|
||||
Executable
+124
@@ -0,0 +1,124 @@
|
||||
#!/usr/bin/env bash
|
||||
set -Eeuo pipefail
|
||||
|
||||
SELF_IP="169.254.100.1"
|
||||
PEER_IP="169.254.100.2"
|
||||
PREFIX="16"
|
||||
IFACE="enP7s7"
|
||||
USE_NM="auto"
|
||||
DRY_RUN=0
|
||||
|
||||
usage() {
|
||||
cat <<EOF
|
||||
Usage: sudo $(basename "$0") [options]
|
||||
|
||||
Configure a Linux Ethernet interface with a static IPv4 for a host-to-host
|
||||
link to a Mac peer.
|
||||
|
||||
Defaults: this host = ${SELF_IP}/${PREFIX}, peer = ${PEER_IP}, iface = ${IFACE}.
|
||||
|
||||
Options:
|
||||
--iface IFACE Default: ${IFACE}
|
||||
--self-ip IP Default: ${SELF_IP}
|
||||
--peer-ip IP For verification ping. Default: ${PEER_IP}
|
||||
--prefix N Default: ${PREFIX}
|
||||
--no-nm Use 'ip addr' directly (transient, no NetworkManager).
|
||||
--dry-run Print actions without applying.
|
||||
-h, --help Show this help.
|
||||
EOF
|
||||
}
|
||||
|
||||
while (($#)); do
|
||||
case "$1" in
|
||||
--iface)
|
||||
shift
|
||||
IFACE="${1:?}"
|
||||
;;
|
||||
--self-ip)
|
||||
shift
|
||||
SELF_IP="${1:?}"
|
||||
;;
|
||||
--peer-ip)
|
||||
shift
|
||||
PEER_IP="${1:?}"
|
||||
;;
|
||||
--prefix)
|
||||
shift
|
||||
PREFIX="${1:?}"
|
||||
;;
|
||||
--no-nm) USE_NM=no ;;
|
||||
--dry-run) DRY_RUN=1 ;;
|
||||
-h | --help)
|
||||
usage
|
||||
exit 0
|
||||
;;
|
||||
*)
|
||||
echo "Unknown arg: $1" >&2
|
||||
usage >&2
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
shift
|
||||
done
|
||||
|
||||
[[ $EUID -eq 0 ]] || {
|
||||
echo "Run as root." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
run() {
|
||||
printf '+'
|
||||
printf ' %q' "$@"
|
||||
printf '\n'
|
||||
((DRY_RUN)) || "$@"
|
||||
}
|
||||
|
||||
ip link show "$IFACE" >/dev/null 2>&1 || {
|
||||
echo "Interface $IFACE does not exist." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
if [[ $USE_NM == "auto" ]]; then
|
||||
if command -v nmcli >/dev/null 2>&1 && systemctl is-active --quiet NetworkManager 2>/dev/null; then
|
||||
USE_NM=yes
|
||||
else
|
||||
USE_NM=no
|
||||
fi
|
||||
fi
|
||||
|
||||
if [[ $USE_NM == "yes" ]]; then
|
||||
CONN="$(nmcli -g GENERAL.CONNECTION device show "$IFACE" 2>/dev/null | head -n1 || true)"
|
||||
if [[ -z $CONN || $CONN == "--" ]]; then
|
||||
CONN="static-${IFACE}"
|
||||
run nmcli connection add type ethernet ifname "$IFACE" con-name "$CONN"
|
||||
fi
|
||||
run nmcli connection modify "$CONN" \
|
||||
connection.interface-name "$IFACE" \
|
||||
connection.autoconnect yes \
|
||||
connection.autoconnect-priority 100 \
|
||||
ipv4.method manual \
|
||||
ipv4.addresses "${SELF_IP}/${PREFIX}" \
|
||||
ipv4.gateway "" \
|
||||
ipv4.dns "" \
|
||||
ipv4.never-default yes \
|
||||
ipv6.method link-local \
|
||||
ipv6.addr-gen-mode stable-privacy
|
||||
run nmcli connection up "$CONN"
|
||||
else
|
||||
run ip link set "$IFACE" up
|
||||
run ip addr flush dev "$IFACE"
|
||||
run ip addr add "${SELF_IP}/${PREFIX}" dev "$IFACE"
|
||||
fi
|
||||
|
||||
if ((!DRY_RUN)); then
|
||||
printf '\n'
|
||||
ip -br addr show "$IFACE"
|
||||
printf '\n'
|
||||
if ping -c2 -W2 "$PEER_IP" >/dev/null 2>&1; then
|
||||
echo "OK: $PEER_IP reachable on $IFACE."
|
||||
else
|
||||
echo "WARN: $PEER_IP not reachable yet."
|
||||
echo " Verify the peer is configured (run setup_linklocal_mac.sh on the Mac)."
|
||||
echo " ip neigh show dev $IFACE # check for the peer MAC"
|
||||
fi
|
||||
fi
|
||||
Executable
+170
@@ -0,0 +1,170 @@
|
||||
#!/usr/bin/env bash
|
||||
set -Eeuo pipefail
|
||||
|
||||
SELF_IP="169.254.100.2"
|
||||
PEER_IP="169.254.100.1"
|
||||
NETMASK="255.255.0.0"
|
||||
IFACE=""
|
||||
DRY_RUN=0
|
||||
|
||||
usage() {
|
||||
cat <<EOF
|
||||
Usage: sudo $(basename "$0") [options]
|
||||
|
||||
Configure a Mac Ethernet interface with a static IPv4 for a host-to-host link
|
||||
to the DGX/GX10 peer.
|
||||
|
||||
Defaults: this Mac = ${SELF_IP}, peer = ${PEER_IP}, mask = ${NETMASK}.
|
||||
|
||||
Options:
|
||||
--iface IFACE Interface (e.g. en12). Default: auto-detect.
|
||||
--self-ip IP This Mac's address. Default: ${SELF_IP}.
|
||||
--peer-ip IP Peer for verification ping. Default: ${PEER_IP}.
|
||||
--netmask MASK Default: ${NETMASK}.
|
||||
--dry-run Print actions without applying.
|
||||
-h, --help Show this help.
|
||||
EOF
|
||||
}
|
||||
|
||||
while (($#)); do
|
||||
case "$1" in
|
||||
--iface)
|
||||
shift
|
||||
IFACE="${1:?}"
|
||||
;;
|
||||
--self-ip)
|
||||
shift
|
||||
SELF_IP="${1:?}"
|
||||
;;
|
||||
--peer-ip)
|
||||
shift
|
||||
PEER_IP="${1:?}"
|
||||
;;
|
||||
--netmask)
|
||||
shift
|
||||
NETMASK="${1:?}"
|
||||
;;
|
||||
--dry-run) DRY_RUN=1 ;;
|
||||
-h | --help)
|
||||
usage
|
||||
exit 0
|
||||
;;
|
||||
*)
|
||||
echo "Unknown arg: $1" >&2
|
||||
usage >&2
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
shift
|
||||
done
|
||||
|
||||
[[ $EUID -eq 0 ]] || {
|
||||
echo "Run with sudo." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
run() {
|
||||
printf '+'
|
||||
printf ' %q' "$@"
|
||||
printf '\n'
|
||||
((DRY_RUN)) || "$@"
|
||||
}
|
||||
|
||||
target_subnet_prefix() {
|
||||
local ip="$1"
|
||||
printf '%s.' "${ip%.*}"
|
||||
}
|
||||
|
||||
iface_score() {
|
||||
local iface="$1" info subnet
|
||||
info="$(ifconfig "$iface" 2>/dev/null || true)"
|
||||
[[ -n $info ]] || {
|
||||
echo 0
|
||||
return
|
||||
}
|
||||
grep -q 'status: active' <<<"$info" || {
|
||||
echo 0
|
||||
return
|
||||
}
|
||||
subnet="$(target_subnet_prefix "$SELF_IP")"
|
||||
if grep -qE "inet ${subnet//./\\.}" <<<"$info"; then
|
||||
echo 100
|
||||
return
|
||||
fi
|
||||
if grep -qE 'inet 169\.254\.' <<<"$info"; then
|
||||
echo 80
|
||||
return
|
||||
fi
|
||||
if ! grep -qE '^[[:space:]]*inet ' <<<"$info"; then
|
||||
echo 60
|
||||
return
|
||||
fi
|
||||
echo 10
|
||||
}
|
||||
|
||||
detect_iface() {
|
||||
local best="" best_score=0 iface score
|
||||
for iface in $(ifconfig -l); do
|
||||
[[ $iface =~ ^en[0-9]+$ ]] || continue
|
||||
score="$(iface_score "$iface")"
|
||||
if ((score > best_score)); then
|
||||
best="$iface"
|
||||
best_score="$score"
|
||||
fi
|
||||
done
|
||||
((best_score >= 60)) || return 1
|
||||
printf '%s\n' "$best"
|
||||
}
|
||||
|
||||
iface_to_service() {
|
||||
local iface="$1" line port=""
|
||||
while IFS= read -r line; do
|
||||
if [[ $line == "Hardware Port: "* ]]; then
|
||||
port="${line#Hardware Port: }"
|
||||
elif [[ $line == "Device: $iface" ]]; then
|
||||
printf '%s\n' "$port"
|
||||
return 0
|
||||
fi
|
||||
done < <(networksetup -listallhardwareports)
|
||||
return 1
|
||||
}
|
||||
|
||||
if [[ -z $IFACE ]]; then
|
||||
IFACE="$(detect_iface || true)"
|
||||
[[ -n $IFACE ]] || {
|
||||
echo "Could not auto-detect a wired interface. Pass --iface enX." >&2
|
||||
echo "Active interfaces:" >&2
|
||||
ifconfig -l | tr ' ' '\n' | grep -E '^en[0-9]+$' | while read -r i; do
|
||||
printf ' %-6s %s\n' "$i" "$(ifconfig "$i" | grep -E 'status:|inet ' | tr '\n' ' ')" >&2
|
||||
done
|
||||
exit 1
|
||||
}
|
||||
echo "Auto-detected interface: $IFACE"
|
||||
fi
|
||||
|
||||
ifconfig "$IFACE" >/dev/null 2>&1 || {
|
||||
echo "Interface $IFACE does not exist." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
SERVICE="$(iface_to_service "$IFACE" || true)"
|
||||
[[ -n $SERVICE ]] || {
|
||||
echo "No network service maps to $IFACE. Check System Settings -> Network." >&2
|
||||
exit 1
|
||||
}
|
||||
echo "Network service: $SERVICE"
|
||||
|
||||
run networksetup -setmanual "$SERVICE" "$SELF_IP" "$NETMASK" ""
|
||||
|
||||
if ((!DRY_RUN)); then
|
||||
printf '\n'
|
||||
ifconfig "$IFACE" | grep -E 'inet |status:'
|
||||
printf '\n'
|
||||
if ping -c2 -t3 "$PEER_IP" >/dev/null 2>&1; then
|
||||
echo "OK: $PEER_IP reachable on $IFACE."
|
||||
else
|
||||
echo "WARN: $PEER_IP not reachable yet."
|
||||
echo " Verify the peer is configured (run setup_linklocal_dgx.sh on the GX10)."
|
||||
echo " arp -an -i $IFACE # check for the peer MAC"
|
||||
fi
|
||||
fi
|
||||
+7
-3
@@ -526,8 +526,8 @@ class API:
|
||||
)
|
||||
]
|
||||
)
|
||||
# TODO: PDD
|
||||
# instance_combinations.append((Sharding.PrefillDecodeDisaggregation, InstanceMeta.MlxRing, 1))
|
||||
if any(self.state.node_vllm.values()):
|
||||
instance_combinations.append((Sharding.Pipeline, InstanceMeta.Vllm, 1))
|
||||
|
||||
for sharding, instance_meta, min_nodes in instance_combinations:
|
||||
try:
|
||||
@@ -640,7 +640,10 @@ class API:
|
||||
)
|
||||
|
||||
async def get_feature_flags(self) -> dict[str, bool]:
|
||||
return {"disaggregation": ENABLE_DISAGGREGATION}
|
||||
return {
|
||||
"disaggregation": ENABLE_DISAGGREGATION,
|
||||
"vllm_available": any(self.state.node_vllm.values()),
|
||||
}
|
||||
|
||||
async def list_instance_links(self) -> list[InstanceLink]:
|
||||
if not ENABLE_DISAGGREGATION:
|
||||
@@ -1748,6 +1751,7 @@ class API:
|
||||
capabilities=card.capabilities,
|
||||
reasoning_dialect=card.reasoning_dialect,
|
||||
context_length=card.context_length,
|
||||
requires_vllm=card.requires_vllm,
|
||||
)
|
||||
for card in cards
|
||||
]
|
||||
|
||||
@@ -49,6 +49,7 @@ class ModelListModel(BaseModel):
|
||||
base_model: str = Field(default="")
|
||||
capabilities: list[str] = Field(default_factory=list)
|
||||
reasoning_dialect: ReasoningDialect = "none"
|
||||
requires_vllm: bool = Field(default=False)
|
||||
|
||||
|
||||
class ModelList(BaseModel):
|
||||
|
||||
+23
-2
@@ -180,10 +180,31 @@ class Master:
|
||||
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():
|
||||
if (
|
||||
model_match = (
|
||||
instance.shard_assignments.model_id
|
||||
== command.task_params.model
|
||||
== requested_model
|
||||
) or (instance.instance_id in linked_decode_ids)
|
||||
if (
|
||||
model_match
|
||||
and instance.instance_id not in prefill_only
|
||||
):
|
||||
in_flight = {TaskStatus.Pending, TaskStatus.Running}
|
||||
|
||||
@@ -43,6 +43,7 @@ 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
|
||||
@@ -202,7 +203,7 @@ def place_instance(
|
||||
)
|
||||
|
||||
# Single-node: force Pipeline/Ring (Tensor and Jaccl require multi-node)
|
||||
if len(selected_cycle) == 1:
|
||||
if len(selected_cycle) == 1 and command.instance_meta != InstanceMeta.Vllm:
|
||||
command = command.model_copy(
|
||||
update={
|
||||
"instance_meta": InstanceMeta.MlxRing,
|
||||
@@ -266,6 +267,11 @@ def place_instance(
|
||||
hosts_by_node=hosts_by_node,
|
||||
ephemeral_port=ephemeral_port,
|
||||
)
|
||||
case InstanceMeta.Vllm:
|
||||
target_instances[instance_id] = VllmInstance(
|
||||
instance_id=instance_id,
|
||||
shard_assignments=shard_assignments,
|
||||
)
|
||||
|
||||
return target_instances
|
||||
|
||||
|
||||
@@ -375,7 +375,13 @@ def find_ip_prioritised(
|
||||
"maybe_ethernet": 3,
|
||||
"thunderbolt": 4,
|
||||
}
|
||||
return min(ips, key=lambda ip: priority.get(ip_to_type.get(ip, "unknown"), 2))
|
||||
|
||||
def _key(ip: str) -> tuple[int, int]:
|
||||
link_local = 0 if ip.startswith("169.254.") else 1
|
||||
type_pri = priority.get(ip_to_type.get(ip, "unknown"), 2)
|
||||
return (link_local, type_pri)
|
||||
|
||||
return min(ips, key=_key)
|
||||
|
||||
|
||||
def get_mlx_ring_hosts_by_node(
|
||||
|
||||
@@ -59,9 +59,11 @@ from exo.utils.info_gatherer.info_gatherer import (
|
||||
NodeConfig,
|
||||
NodeDiskUsage,
|
||||
NodeNetworkInterfaces,
|
||||
NvmlMetrics,
|
||||
RdmaCtlStatus,
|
||||
StaticNodeInformation,
|
||||
ThunderboltBridgeInfo,
|
||||
VllmCapability,
|
||||
)
|
||||
|
||||
|
||||
@@ -304,6 +306,9 @@ def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
|
||||
node_rdma_ctl = {
|
||||
key: value for key, value in state.node_rdma_ctl.items() if key != event.node_id
|
||||
}
|
||||
node_vllm = {
|
||||
key: value for key, value in state.node_vllm.items() if key != event.node_id
|
||||
}
|
||||
# Only recompute cycles if the leaving node had TB bridge enabled
|
||||
leaving_node_status = state.node_thunderbolt_bridge.get(event.node_id)
|
||||
leaving_node_had_tb_enabled = (
|
||||
@@ -326,6 +331,7 @@ def apply_node_timed_out(event: NodeTimedOut, state: State) -> State:
|
||||
"node_thunderbolt": node_thunderbolt,
|
||||
"node_thunderbolt_bridge": node_thunderbolt_bridge,
|
||||
"node_rdma_ctl": node_rdma_ctl,
|
||||
"node_vllm": node_vllm,
|
||||
"thunderbolt_bridge_cycles": thunderbolt_bridge_cycles,
|
||||
}
|
||||
)
|
||||
@@ -352,6 +358,11 @@ def apply_node_gathered_info(event: NodeGatheredInfo, state: State) -> State:
|
||||
event.node_id: info.system_profile,
|
||||
}
|
||||
update["node_memory"] = {**state.node_memory, event.node_id: info.memory}
|
||||
case NvmlMetrics():
|
||||
update["node_system"] = {
|
||||
**state.node_system,
|
||||
event.node_id: info.system_profile,
|
||||
}
|
||||
case MemoryUsage():
|
||||
update["node_memory"] = {**state.node_memory, event.node_id: info}
|
||||
case NodeDiskUsage():
|
||||
@@ -432,6 +443,11 @@ def apply_node_gathered_info(event: NodeGatheredInfo, state: State) -> State:
|
||||
**state.node_rdma_ctl,
|
||||
event.node_id: NodeRdmaCtlStatus(enabled=info.enabled),
|
||||
}
|
||||
case VllmCapability():
|
||||
update["node_vllm"] = {
|
||||
**state.node_vllm,
|
||||
event.node_id: info.available,
|
||||
}
|
||||
|
||||
return state.model_copy(update=update)
|
||||
|
||||
|
||||
@@ -150,6 +150,7 @@ class ModelCard(FrozenModel):
|
||||
context_length: int = 0
|
||||
uses_cfg: bool = False
|
||||
trust_remote_code: bool = True
|
||||
requires_vllm: bool = False
|
||||
is_custom: bool = False
|
||||
vision: VisionCardConfig | None = None
|
||||
sampling_defaults: SamplingDefaults = Field(default_factory=SamplingDefaults)
|
||||
@@ -349,7 +350,11 @@ async def fetch_config_data(model_id: ModelId) -> ConfigData:
|
||||
|
||||
|
||||
async def fetch_safetensors_size(model_id: ModelId) -> Memory:
|
||||
"""Gets model size from safetensors index or falls back to HF API."""
|
||||
"""Gets model size from safetensors index or falls back to HF API.
|
||||
|
||||
Single-shard repos don't have a `model.safetensors.index.json`; fall back
|
||||
to the HF API for those.
|
||||
"""
|
||||
from exo.download.download_utils import (
|
||||
download_file_with_retry,
|
||||
resolve_model_dir,
|
||||
@@ -357,21 +362,25 @@ async def fetch_safetensors_size(model_id: ModelId) -> Memory:
|
||||
from exo.shared.types.worker.downloads import ModelSafetensorsIndex
|
||||
|
||||
target_dir = await resolve_model_dir(model_id)
|
||||
index_path = await download_file_with_retry(
|
||||
model_id,
|
||||
"main",
|
||||
"model.safetensors.index.json",
|
||||
target_dir,
|
||||
lambda curr_bytes, total_bytes, is_renamed: logger.debug(
|
||||
f"Downloading model.safetensors.index.json for {model_id}: {curr_bytes}/{total_bytes} ({is_renamed=})"
|
||||
),
|
||||
)
|
||||
async with aiofiles.open(index_path, "r") as f:
|
||||
index_data = ModelSafetensorsIndex.model_validate_json(await f.read())
|
||||
try:
|
||||
index_path = await download_file_with_retry(
|
||||
model_id,
|
||||
"main",
|
||||
"model.safetensors.index.json",
|
||||
target_dir,
|
||||
lambda curr_bytes, total_bytes, is_renamed: logger.debug(
|
||||
f"Downloading model.safetensors.index.json for {model_id}: {curr_bytes}/{total_bytes} ({is_renamed=})"
|
||||
),
|
||||
)
|
||||
except FileNotFoundError:
|
||||
index_path = None
|
||||
|
||||
metadata = index_data.metadata
|
||||
if metadata is not None and metadata.total_size is not None:
|
||||
return Memory.from_bytes(metadata.total_size)
|
||||
if index_path is not None:
|
||||
async with aiofiles.open(index_path, "r") as f:
|
||||
index_data = ModelSafetensorsIndex.model_validate_json(await f.read())
|
||||
metadata = index_data.metadata
|
||||
if metadata is not None and metadata.total_size is not None:
|
||||
return Memory.from_bytes(metadata.total_size)
|
||||
|
||||
info = model_info(model_id)
|
||||
if info.safetensors is None:
|
||||
|
||||
@@ -58,6 +58,7 @@ 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]] = []
|
||||
|
||||
@@ -15,6 +15,7 @@ class InstanceId(Id):
|
||||
class InstanceMeta(str, Enum):
|
||||
MlxRing = "MlxRing"
|
||||
MlxJaccl = "MlxJaccl"
|
||||
Vllm = "Vllm"
|
||||
|
||||
|
||||
class BaseInstance(TaggedModel):
|
||||
@@ -35,8 +36,12 @@ class MlxJacclInstance(BaseInstance):
|
||||
jaccl_coordinators: dict[NodeId, str]
|
||||
|
||||
|
||||
class VllmInstance(BaseInstance):
|
||||
pass
|
||||
|
||||
|
||||
# TODO: Single node instance
|
||||
Instance = MlxRingInstance | MlxJacclInstance
|
||||
Instance = MlxRingInstance | MlxJacclInstance | VllmInstance
|
||||
|
||||
|
||||
class BoundInstance(FrozenModel):
|
||||
|
||||
@@ -25,12 +25,12 @@ def print_startup_banner(port: int) -> None:
|
||||
banner = f"""
|
||||
╔═══════════════════════════════════════════════════════════════════════╗
|
||||
║ ║
|
||||
║ ███████╗██╗ ██╗ ██████╗ ║
|
||||
║ ██╔════╝╚██╗██╔╝██╔═══██╗ ║
|
||||
║ █████╗ ╚███╔╝ ██║ ██║ ║
|
||||
║ ██╔══╝ ██╔██╗ ██║ ██║ ║
|
||||
║ ███████╗██╔╝ ██╗╚██████╔╝ ║
|
||||
║ ╚══════╝╚═╝ ╚═╝ ╚═════╝ ║
|
||||
║ ███████╗██╗ ██╗ ██████╗ ██████╗ ██████╗ ██╗ ██╗ ║
|
||||
║ ██╔════╝╚██╗██╔╝██╔═══██╗ ██ ██╔══██╗██╔════╝ ╚██╗██╔╝ ║
|
||||
║ █████╗ ╚███╔╝ ██║ ██║ ██████╗ ██║ ██║██║ ███╗ ╚███╔╝ ║
|
||||
║ ██╔══╝ ██╔██╗ ██║ ██║ ╚═██╔═╝ ██║ ██║██║ ██║ ██╔██╗ ║
|
||||
║ ███████╗██╔╝ ██╗╚██████╔╝ ╚═╝ ██████╔╝╚██████╔╝██╔╝ ██╗ ║
|
||||
║ ╚══════╝╚═╝ ╚═╝ ╚═════╝ ╚═════╝ ╚═════╝ ╚═╝ ╚═╝ ║
|
||||
║ ║
|
||||
║ Distributed AI Inference Cluster ║
|
||||
║ ║
|
||||
|
||||
@@ -31,6 +31,7 @@ from exo.utils.pydantic_ext import TaggedModel
|
||||
from exo.utils.task_group import TaskGroup
|
||||
|
||||
from .macmon import MacmonMetrics
|
||||
from .nvml import NvmlMetrics, gather_nvidia_metrics, has_nvml
|
||||
from .system_info import (
|
||||
get_friendly_name,
|
||||
get_model_and_chip,
|
||||
@@ -353,6 +354,24 @@ async def _gather_iface_map() -> dict[str, str] | None:
|
||||
return ports
|
||||
|
||||
|
||||
class VllmCapability(TaggedModel):
|
||||
available: bool
|
||||
version: str | None = None
|
||||
|
||||
@classmethod
|
||||
async def gather(cls) -> Self:
|
||||
try:
|
||||
import importlib
|
||||
|
||||
vllm = importlib.import_module("vllm")
|
||||
return cls(
|
||||
available=True,
|
||||
version=cast(str | None, getattr(vllm, "__version__", None)),
|
||||
)
|
||||
except ImportError:
|
||||
return cls(available=False)
|
||||
|
||||
|
||||
GatheredInfo = (
|
||||
MacmonMetrics
|
||||
| MemoryUsage
|
||||
@@ -361,6 +380,8 @@ GatheredInfo = (
|
||||
| MacThunderboltConnections
|
||||
| RdmaCtlStatus
|
||||
| ThunderboltBridgeInfo
|
||||
| NvmlMetrics
|
||||
| VllmCapability
|
||||
| NodeConfig
|
||||
| MiscData
|
||||
| StaticNodeInformation
|
||||
@@ -419,6 +440,8 @@ class InfoGatherer:
|
||||
tg.start_soon(self._monitor_rdma_ctl_status, 10)
|
||||
if not IS_DARWIN:
|
||||
tg.start_soon(self._monitor_memory_usage, 1)
|
||||
if has_nvml():
|
||||
tg.start_soon(self._monitor_nvml_metrics, 1)
|
||||
tg.start_soon(self._watch_system_info, 10)
|
||||
tg.start_soon(self._monitor_misc, 60)
|
||||
tg.start_soon(self._monitor_static_info, 60)
|
||||
@@ -427,6 +450,10 @@ class InfoGatherer:
|
||||
nc = await NodeConfig.gather()
|
||||
if nc is not None:
|
||||
await self.info_sender.send(nc)
|
||||
try:
|
||||
await self.info_sender.send(await VllmCapability.gather())
|
||||
except Exception as e:
|
||||
logger.warning(f"Error gathering vLLM capability: {e}")
|
||||
|
||||
def shutdown(self):
|
||||
self._tg.cancel_tasks()
|
||||
@@ -475,6 +502,16 @@ class InfoGatherer:
|
||||
logger.opt(exception=e).warning("Error gathering Thunderbolt data")
|
||||
await anyio.sleep(system_profiler_interval)
|
||||
|
||||
async def _monitor_nvml_metrics(self, nvml_poll_rate: float):
|
||||
while True:
|
||||
try:
|
||||
metrics = gather_nvidia_metrics()
|
||||
if metrics is not None:
|
||||
await self.info_sender.send(metrics)
|
||||
except Exception as e:
|
||||
logger.opt(exception=e).warning("Error gathering NVML metrics")
|
||||
await anyio.sleep(nvml_poll_rate)
|
||||
|
||||
async def _monitor_memory_usage(self, memory_poll_rate: float):
|
||||
if self._psutil_enabled:
|
||||
return
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
from exo.shared.types.profiling import SystemPerformanceProfile
|
||||
from exo.utils.pydantic_ext import TaggedModel
|
||||
|
||||
try:
|
||||
import pynvml as nvml
|
||||
except ImportError:
|
||||
nvml = None
|
||||
|
||||
_CPU_POWER_IDLE = 20.0
|
||||
_CPU_POWER_MAX = 100.0
|
||||
_GPU_POWER_MAX = 120.0
|
||||
|
||||
|
||||
class NvmlMetrics(TaggedModel):
|
||||
system_profile: SystemPerformanceProfile
|
||||
|
||||
|
||||
def has_nvml() -> bool:
|
||||
if nvml is None:
|
||||
return False
|
||||
try:
|
||||
nvml.nvmlInit()
|
||||
count = nvml.nvmlDeviceGetCount()
|
||||
nvml.nvmlShutdown()
|
||||
return count > 0
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def gather_nvidia_metrics() -> NvmlMetrics | None:
|
||||
if nvml is None:
|
||||
return None
|
||||
|
||||
is_init = False
|
||||
try:
|
||||
nvml.nvmlInit()
|
||||
is_init = True
|
||||
count = nvml.nvmlDeviceGetCount()
|
||||
if count == 0:
|
||||
return None
|
||||
|
||||
total_gpu_util = 0.0
|
||||
total_temp = 0.0
|
||||
total_gpu_power = 0.0
|
||||
for i in range(count):
|
||||
handle = nvml.nvmlDeviceGetHandleByIndex(i)
|
||||
util = nvml.nvmlDeviceGetUtilizationRates(handle)
|
||||
total_gpu_util += float(util.gpu)
|
||||
total_temp += float(
|
||||
nvml.nvmlDeviceGetTemperatureV(handle, nvml.NVML_TEMPERATURE_GPU)
|
||||
)
|
||||
total_gpu_power += float(nvml.nvmlDeviceGetPowerUsage(handle)) / 1000.0
|
||||
|
||||
gpu_load_fraction = min(total_gpu_power / _GPU_POWER_MAX, 1.0)
|
||||
estimated_cpu_power = (
|
||||
_CPU_POWER_IDLE + (_CPU_POWER_MAX - _CPU_POWER_IDLE) * gpu_load_fraction
|
||||
)
|
||||
|
||||
return NvmlMetrics(
|
||||
system_profile=SystemPerformanceProfile(
|
||||
gpu_usage=total_gpu_util / count / 100.0,
|
||||
temp=total_temp / count,
|
||||
sys_power=total_gpu_power + estimated_cpu_power,
|
||||
),
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
finally:
|
||||
if is_init:
|
||||
nvml.nvmlShutdown()
|
||||
@@ -1,6 +1,7 @@
|
||||
import platform
|
||||
import socket
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from subprocess import CalledProcessError
|
||||
|
||||
import psutil
|
||||
@@ -117,12 +118,90 @@ async def get_network_interfaces() -> list[NetworkInterfaceInfo]:
|
||||
return interfaces_info
|
||||
|
||||
|
||||
def _read_dmi_field(name: str) -> str | None:
|
||||
try:
|
||||
path = Path(f"/sys/class/dmi/id/{name}")
|
||||
if path.exists():
|
||||
return path.read_text().strip()
|
||||
except (OSError, PermissionError):
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
async def _get_linux_model_and_chip() -> tuple[str, str]:
|
||||
model = "Linux"
|
||||
chip = "Unknown Chip"
|
||||
|
||||
product_name = _read_dmi_field("product_name")
|
||||
sys_vendor = _read_dmi_field("sys_vendor")
|
||||
|
||||
# DGX Spark: DMI product_name may be "DGX_Spark" or "gx10" variant
|
||||
product_lower = (product_name or "").lower()
|
||||
if product_name and ("dgx" in product_lower or "gx10" in product_lower):
|
||||
model = "DGX Spark"
|
||||
try:
|
||||
process = await run_process(
|
||||
["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"]
|
||||
)
|
||||
gpu_name = process.stdout.decode().strip().split("\n")[0]
|
||||
chip = gpu_name if gpu_name and gpu_name != "[N/A]" else "NVIDIA GB10"
|
||||
except (CalledProcessError, FileNotFoundError):
|
||||
chip = "NVIDIA GB10"
|
||||
return (model, chip)
|
||||
|
||||
# Other NVIDIA systems (sys_vendor contains "NVIDIA")
|
||||
if sys_vendor and "NVIDIA" in sys_vendor:
|
||||
model = product_name.replace("_", " ") if product_name else "NVIDIA System"
|
||||
try:
|
||||
process = await run_process(
|
||||
["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"]
|
||||
)
|
||||
gpu_name = process.stdout.decode().strip().split("\n")[0]
|
||||
if gpu_name and gpu_name != "[N/A]":
|
||||
chip = gpu_name
|
||||
except (CalledProcessError, FileNotFoundError):
|
||||
pass
|
||||
return (model, chip)
|
||||
|
||||
# Generic Linux — detect laptop vs desktop via chassis_type
|
||||
# SMBIOS chassis types: 8,9,10,14,31,32 = portable/laptop
|
||||
chassis_type = _read_dmi_field("chassis_type")
|
||||
laptop_chassis_types = {"8", "9", "10", "14", "31", "32"}
|
||||
if chassis_type in laptop_chassis_types:
|
||||
model = "Linux Laptop"
|
||||
elif chassis_type is not None:
|
||||
model = "Linux Desktop"
|
||||
|
||||
# Also check for battery as a fallback laptop indicator
|
||||
if model == "Linux" and Path("/sys/class/power_supply/BAT0").exists():
|
||||
model = "Linux Laptop"
|
||||
|
||||
# Use /proc/cpuinfo for chip
|
||||
cpuinfo_path = Path("/proc/cpuinfo")
|
||||
if cpuinfo_path.exists():
|
||||
try:
|
||||
for line in cpuinfo_path.read_text().splitlines():
|
||||
if line.startswith("model name"):
|
||||
chip = line.split(":", 1)[1].strip()
|
||||
break
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
return (model, chip)
|
||||
|
||||
|
||||
async def get_model_and_chip() -> tuple[str, str]:
|
||||
"""Get Mac system information using system_profiler."""
|
||||
"""Get system model and chip information.
|
||||
|
||||
On macOS, uses ``system_profiler``. On Linux, reads DMI data from
|
||||
sysfs and CPU info from ``/proc/cpuinfo``.
|
||||
"""
|
||||
model = "Unknown Model"
|
||||
chip = "Unknown Chip"
|
||||
|
||||
# TODO: better non mac support
|
||||
if sys.platform == "linux":
|
||||
return await _get_linux_model_and_chip()
|
||||
|
||||
if sys.platform != "darwin":
|
||||
return (model, chip)
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import BinaryIO, Literal
|
||||
|
||||
import msgspec
|
||||
@@ -23,7 +24,28 @@ class TensorBlob(msgspec.Struct):
|
||||
data: bytes
|
||||
|
||||
|
||||
class KVChunk(msgspec.Struct, tag="kv_chunk"):
|
||||
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
|
||||
@@ -51,10 +73,13 @@ class ErrorMessage(msgspec.Struct, tag="error"):
|
||||
message: str
|
||||
|
||||
|
||||
_WireMessage = _KVChunkHeader | ArraysState | Done | ErrorMessage
|
||||
Message = KVChunk | ArraysState | Done | ErrorMessage
|
||||
|
||||
_msg_encoder = msgspec.msgpack.Encoder()
|
||||
_msg_decoder: msgspec.msgpack.Decoder[Message] = msgspec.msgpack.Decoder(Message)
|
||||
_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)
|
||||
|
||||
@@ -99,7 +124,7 @@ def read_header(stream: BinaryIO) -> Header:
|
||||
raise ProtocolError(f"Bad header: {exc}") from exc
|
||||
|
||||
|
||||
def write_message(stream: BinaryIO, msg: Message) -> None:
|
||||
def write_message(stream: BinaryIO, msg: _WireMessage) -> None:
|
||||
write_frame(stream, _msg_encoder.encode(msg))
|
||||
|
||||
|
||||
@@ -108,9 +133,22 @@ def read_message(stream: BinaryIO) -> Message | None:
|
||||
if not payload:
|
||||
return None
|
||||
try:
|
||||
return _msg_decoder.decode(payload)
|
||||
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(
|
||||
@@ -121,21 +159,35 @@ def write_kv_chunk(
|
||||
n_heads: int,
|
||||
head_dim: int,
|
||||
dtype: DType,
|
||||
keys: bytes,
|
||||
values: bytes,
|
||||
keys: "bytes | memoryview",
|
||||
values: "bytes | memoryview",
|
||||
) -> None:
|
||||
write_message(
|
||||
stream,
|
||||
KVChunk(
|
||||
"""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=keys,
|
||||
values=values,
|
||||
),
|
||||
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(
|
||||
|
||||
@@ -21,6 +21,7 @@ class PrefillRequest(msgspec.Struct):
|
||||
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()
|
||||
@@ -56,7 +57,10 @@ class _PrefillHandler(socketserver.StreamRequestHandler):
|
||||
super().setup()
|
||||
sock = cast(socket.socket, self.request)
|
||||
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1)
|
||||
sock.setsockopt(socket.SOL_SOCKET, socket.SO_SNDBUF, 4 * 1024 * 1024)
|
||||
# 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)
|
||||
|
||||
@@ -34,7 +34,7 @@ class MlxBuilder(Builder):
|
||||
model_id: ModelId
|
||||
event_sender: MpSender[Event]
|
||||
cancel_receiver: MpReceiver[TaskId]
|
||||
inference_model: Model | None = None
|
||||
model: Model | None = None
|
||||
tokenizer: TokenizerWrapper | None = None
|
||||
group: mx.distributed.Group | None = None
|
||||
vision_processor: VisionProcessor | None = None
|
||||
@@ -44,14 +44,14 @@ class MlxBuilder(Builder):
|
||||
|
||||
def load(self, bound_instance: BoundInstance) -> Generator[ModelLoadingResponse]:
|
||||
(
|
||||
self.inference_model,
|
||||
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.inference_model
|
||||
del self.model
|
||||
with contextlib.suppress(NameError, AttributeError):
|
||||
del self.tokenizer
|
||||
with contextlib.suppress(NameError, AttributeError):
|
||||
@@ -60,7 +60,7 @@ class MlxBuilder(Builder):
|
||||
def build(
|
||||
self,
|
||||
) -> Engine:
|
||||
assert self.inference_model
|
||||
assert self.model
|
||||
assert self.tokenizer
|
||||
|
||||
vision_processor = self.vision_processor
|
||||
@@ -86,7 +86,7 @@ class MlxBuilder(Builder):
|
||||
if os.environ.get("EXO_NO_BATCH"):
|
||||
logger.info("using SequentialGenerator (batching disabled)")
|
||||
return SequentialGenerator(
|
||||
model=self.inference_model,
|
||||
model=self.model,
|
||||
tokenizer=self.tokenizer,
|
||||
group=self.group,
|
||||
tool_parser=tool_parser,
|
||||
@@ -100,7 +100,7 @@ class MlxBuilder(Builder):
|
||||
else:
|
||||
logger.info("using BatchGenerator")
|
||||
return BatchGenerator(
|
||||
model=self.inference_model,
|
||||
model=self.model,
|
||||
tokenizer=self.tokenizer,
|
||||
group=self.group,
|
||||
tool_parser=tool_parser,
|
||||
|
||||
@@ -22,7 +22,6 @@ from exo.worker.disaggregated.protocol import (
|
||||
write_kv_chunk,
|
||||
)
|
||||
from exo.worker.engines.mlx.types import KVCacheType
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
|
||||
_STR_TO_MX: dict[DType, mx.Dtype] = {
|
||||
"bfloat16": mx.bfloat16,
|
||||
@@ -90,6 +89,18 @@ def nhd_to_bhsd(t: mx.array) -> mx.array:
|
||||
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,
|
||||
@@ -103,7 +114,7 @@ def send_mlx_kv_cache(
|
||||
match c:
|
||||
case QuantizedKVCache() | CacheList() | DeepseekV4Cache():
|
||||
raise NotImplementedError
|
||||
case KVCache() | RotatingKVCache():
|
||||
case KVCache():
|
||||
keys = c.keys
|
||||
values = c.values
|
||||
if keys is None or values is None:
|
||||
@@ -132,11 +143,39 @@ def send_mlx_kv_cache(
|
||||
keys=array_to_bytes(k_nhd),
|
||||
values=array_to_bytes(v_nhd),
|
||||
)
|
||||
if tokens_sent != 0 and num_tokens != tokens_sent:
|
||||
logger.critical(
|
||||
f"Unexpected number of tokens sent {num_tokens} != {tokens_sent}"
|
||||
)
|
||||
tokens_sent = num_tokens
|
||||
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:
|
||||
|
||||
@@ -79,6 +79,7 @@ def remote_prefill_fetch(
|
||||
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)
|
||||
@@ -93,10 +94,17 @@ def remote_prefill_fetch(
|
||||
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:
|
||||
|
||||
@@ -215,14 +215,18 @@ class ExoBatchGenerator:
|
||||
with vision_ctx:
|
||||
if use_remote and task_params.prefill_endpoint is not None:
|
||||
try:
|
||||
# Send full prompt; producer's vLLM APC handles the prefix
|
||||
# match. `start_pos` aligns the writer's skip_tokens with
|
||||
# the consumer's locally-cached prefix.
|
||||
_prefill_tps, _prefill_tokens, cache_snapshots = remote_prefill(
|
||||
prompt_tokens[:-1],
|
||||
all_prompt_tokens[:-1],
|
||||
cache,
|
||||
on_prefill_progress,
|
||||
endpoint=task_params.prefill_endpoint,
|
||||
request_id=str(uuid.uuid4()),
|
||||
model_id=str(task_params.model),
|
||||
start_pos=prefix_hit_length,
|
||||
use_prefix_cache=not is_bench or task_params.use_prefix_cache,
|
||||
)
|
||||
remote_prefilled = True
|
||||
except Exception:
|
||||
|
||||
@@ -648,14 +648,20 @@ def mlx_generate(
|
||||
with maybe_vision_ctx:
|
||||
if use_remote and task.prefill_endpoint is not None:
|
||||
try:
|
||||
# Send the FULL prompt to the producer (not the cache-stripped
|
||||
# suffix). vLLM's APC handles the prefix match internally;
|
||||
# `start_pos` tells our extractor / wire writer how much of the
|
||||
# producer-side capture corresponds to tokens the consumer
|
||||
# already has, so the writer's skip_tokens math aligns.
|
||||
prefill_tps, prefill_tokens, ssm_snapshots_list = remote_prefill(
|
||||
prompt_tokens[:-1],
|
||||
all_prompt_tokens[:-1],
|
||||
caches,
|
||||
on_prefill_progress,
|
||||
endpoint=task.prefill_endpoint,
|
||||
request_id=str(uuid.uuid4()),
|
||||
model_id=str(task.model),
|
||||
start_pos=prefix_hit_length,
|
||||
use_prefix_cache=not is_bench or task.use_prefix_cache,
|
||||
)
|
||||
remote_prefilled = True
|
||||
except Exception:
|
||||
|
||||
@@ -25,23 +25,25 @@ def remote_prefill(
|
||||
request_id: str,
|
||||
model_id: str,
|
||||
start_pos: int = 0,
|
||||
use_prefix_cache: bool = True,
|
||||
) -> tuple[float, int, list[CacheSnapshot]]:
|
||||
t0 = time.perf_counter()
|
||||
total_prompt_tokens = int(prompt_tokens.shape[0])
|
||||
num_layers: int = 0
|
||||
tokens_received_total: int = 0
|
||||
|
||||
def _on_header(header: Header) -> None:
|
||||
nonlocal num_layers
|
||||
num_layers = header.num_layers
|
||||
|
||||
def _on_chunk(_chunk: KVChunk, chunks_received: int) -> None:
|
||||
nonlocal num_layers
|
||||
def _on_chunk(chunk: KVChunk, chunks_received: int) -> None:
|
||||
nonlocal num_layers, tokens_received_total
|
||||
tokens_received_total += chunk.num_tokens
|
||||
if on_prefill_progress is None:
|
||||
return
|
||||
if num_layers > 0 and chunks_received % num_layers == 0:
|
||||
tokens_so_far = chunks_received // num_layers
|
||||
on_prefill_progress(
|
||||
min(tokens_so_far, total_prompt_tokens),
|
||||
min(tokens_received_total // num_layers, total_prompt_tokens),
|
||||
total_prompt_tokens,
|
||||
)
|
||||
|
||||
@@ -50,6 +52,7 @@ def remote_prefill(
|
||||
token_ids=cast(list[int], prompt_tokens.tolist()),
|
||||
start_pos=start_pos,
|
||||
request_id=request_id,
|
||||
use_prefix_cache=use_prefix_cache,
|
||||
)
|
||||
result = remote_prefill_fetch(
|
||||
endpoint, request, on_header=_on_header, on_kv_chunk=_on_chunk
|
||||
@@ -61,6 +64,28 @@ def remote_prefill(
|
||||
t_done = time.perf_counter()
|
||||
|
||||
num_tokens = final_offset - start_pos
|
||||
# The producer strips the last 2 tokens of the prompt (consumer warm-starts
|
||||
# decode from those locally). Anything within `producer_strip` of the full
|
||||
# suffix is the expected outcome, not a bug.
|
||||
producer_strip = 2 if total_prompt_tokens > 2 else 0
|
||||
expected_min = max(0, total_prompt_tokens - start_pos - producer_strip)
|
||||
expected_max = max(0, total_prompt_tokens - start_pos)
|
||||
if num_tokens <= 0:
|
||||
raise RuntimeError(
|
||||
f"Remote prefill returned no KV (start_pos={start_pos}, "
|
||||
f"final_offset={final_offset}, expected={expected_min}, "
|
||||
f"transfer={(t_received - t0) * 1000:.0f}ms)"
|
||||
)
|
||||
if num_tokens < expected_min:
|
||||
logger.warning(
|
||||
f"Remote prefill returned {num_tokens} tokens, expected at least "
|
||||
f"{expected_min} (start_pos={start_pos}, final_offset={final_offset})"
|
||||
)
|
||||
elif num_tokens > expected_max:
|
||||
logger.warning(
|
||||
f"Remote prefill returned {num_tokens} tokens, expected at most "
|
||||
f"{expected_max} (start_pos={start_pos}, final_offset={final_offset})"
|
||||
)
|
||||
tps = num_tokens / max(t_done - t0, 0.001)
|
||||
|
||||
logger.info(
|
||||
|
||||
@@ -50,6 +50,7 @@ from exo.shared.types.worker.instances import (
|
||||
BoundInstance,
|
||||
MlxJacclInstance,
|
||||
MlxRingInstance,
|
||||
VllmInstance,
|
||||
)
|
||||
from exo.shared.types.worker.runner_response import ModelLoadingResponse
|
||||
from exo.shared.types.worker.shards import (
|
||||
@@ -140,6 +141,8 @@ def mlx_distributed_init(
|
||||
os.environ["MLX_RANK"] = str(rank)
|
||||
os.environ["MLX_JACCL_COORDINATOR"] = jaccl_coordinator
|
||||
group = mx.distributed.init(backend="jaccl", strict=True)
|
||||
case VllmInstance():
|
||||
raise ValueError("loaded VllmInstance in MLX engine")
|
||||
|
||||
logger.info(f"Rank {rank} mlx distributed initialization complete")
|
||||
|
||||
|
||||
Whitespace-only changes.
@@ -0,0 +1,88 @@
|
||||
import contextlib
|
||||
import os
|
||||
from collections.abc import Generator
|
||||
from dataclasses import dataclass
|
||||
|
||||
from exo.shared.constants import EXO_MAX_CONCURRENT_REQUESTS
|
||||
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.engines.vllm.engine import VllmEngine
|
||||
from exo.worker.engines.vllm.generator import VllmBatchEngine, load_vllm_engine
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
|
||||
|
||||
@dataclass
|
||||
class VllmBuilder(Builder):
|
||||
model_id: ModelId
|
||||
event_sender: MpSender[Event]
|
||||
cancel_receiver: MpReceiver[TaskId]
|
||||
|
||||
def connect(self, bound_instance: BoundInstance) -> None:
|
||||
raise NotImplementedError(
|
||||
"Multiple node VLLM instances are not supported at the moment!"
|
||||
)
|
||||
|
||||
def load(
|
||||
self,
|
||||
bound_instance: BoundInstance,
|
||||
) -> Generator[ModelLoadingResponse]:
|
||||
from exo.worker.engines.vllm.kv_connector import (
|
||||
ExoKVProducerConnector,
|
||||
_patch_gdn_capture,
|
||||
_patch_vllm_for_connector,
|
||||
)
|
||||
|
||||
# Apply bypass patches before vLLM init reads its connector registry
|
||||
# and the unifier touches hybrid kv-cache specs.
|
||||
_patch_vllm_for_connector(ExoKVProducerConnector)
|
||||
_patch_gdn_capture()
|
||||
|
||||
kv_connector_cls: type[object] | None = ExoKVProducerConnector
|
||||
# overlapping = not os.environ.get("EXO_NO_OVERLAPPING_PREFILL_SENDS")
|
||||
|
||||
def on_layer_loaded(loaded: int, total: int) -> None:
|
||||
pass
|
||||
|
||||
self._bound_runner_id = bound_instance.bound_runner_id
|
||||
self._engine, self._tool_parser = load_vllm_engine(
|
||||
model_id=self.model_id,
|
||||
trust_remote_code=bound_instance.bound_shard.model_card.trust_remote_code,
|
||||
n_layers=bound_instance.bound_shard.model_card.n_layers,
|
||||
on_layer_loaded=on_layer_loaded,
|
||||
kv_connector_cls=kv_connector_cls,
|
||||
)
|
||||
return
|
||||
yield
|
||||
|
||||
def build(self) -> Engine:
|
||||
gen = VllmBatchEngine(
|
||||
engine=self._engine,
|
||||
model_id=self.model_id,
|
||||
)
|
||||
try:
|
||||
max_concurrent = (
|
||||
1
|
||||
if bool(os.getenv("EXO_NO_BATCH", False))
|
||||
else EXO_MAX_CONCURRENT_REQUESTS
|
||||
)
|
||||
except Exception:
|
||||
max_concurrent = EXO_MAX_CONCURRENT_REQUESTS
|
||||
|
||||
logger.info(f"using VllmEngine (max_concurrent={max_concurrent})")
|
||||
return VllmEngine(
|
||||
tool_parser=self._tool_parser,
|
||||
model_id=self.model_id,
|
||||
cancel_receiver=self.cancel_receiver,
|
||||
event_sender=self.event_sender,
|
||||
_gen=gen,
|
||||
max_concurrent_requests=max_concurrent,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
with contextlib.suppress(NameError, AttributeError):
|
||||
del self._engine, self._tool_parser
|
||||
Whitespace-only changes.
@@ -0,0 +1,238 @@
|
||||
"""vLLM-side disaggregation adapter.
|
||||
|
||||
Mirrors `engines/mlx/disaggregated/adapter.py` for the vLLM engine: owns
|
||||
torch dtype ↔ wire dtype, byte (de)serialization, layout conversion (vLLM's
|
||||
paged block storage → NHD per-token), and the wire-write helpers used by
|
||||
the producer connector + serve_prefill flow.
|
||||
|
||||
Wire format is `engines/.../disaggregated/protocol.py` (msgpack), shared with
|
||||
the MLX side.
|
||||
"""
|
||||
|
||||
import os
|
||||
from typing import BinaryIO
|
||||
|
||||
import torch
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig
|
||||
|
||||
from exo.worker.disaggregated.protocol import (
|
||||
DType,
|
||||
Header,
|
||||
TensorBlob,
|
||||
write_arrays_state,
|
||||
write_done,
|
||||
write_header,
|
||||
write_kv_chunk,
|
||||
)
|
||||
|
||||
_TORCH_TO_WIRE: dict[torch.dtype, DType] = {
|
||||
torch.bfloat16: "bfloat16",
|
||||
torch.float16: "float16",
|
||||
torch.float32: "float32",
|
||||
}
|
||||
_WIRE_TO_TORCH: dict[DType, torch.dtype] = {v: k for k, v in _TORCH_TO_WIRE.items()}
|
||||
|
||||
|
||||
def torch_dtype_to_wire(dtype: torch.dtype) -> DType:
|
||||
if dtype not in _TORCH_TO_WIRE:
|
||||
raise ValueError(f"Unsupported torch dtype on wire: {dtype}")
|
||||
return _TORCH_TO_WIRE[dtype]
|
||||
|
||||
|
||||
def wire_to_torch_dtype(dtype: DType) -> torch.dtype:
|
||||
return _WIRE_TO_TORCH[dtype]
|
||||
|
||||
|
||||
def tensor_to_wire_bytes(t: torch.Tensor) -> bytes:
|
||||
"""Serialize an NHD-laid-out tensor to wire bytes.
|
||||
|
||||
bfloat16 has no native numpy dtype — bitcast through uint16.
|
||||
"""
|
||||
t = t.detach().contiguous().cpu()
|
||||
if t.dtype == torch.bfloat16:
|
||||
return bytes(t.view(torch.uint16).numpy().tobytes())
|
||||
return bytes(t.numpy().tobytes())
|
||||
|
||||
|
||||
def to_nhd(t: torch.Tensor) -> torch.Tensor:
|
||||
"""Permute HND → NHD when vLLM's KV cache layout is HND."""
|
||||
if os.environ.get("VLLM_KV_CACHE_LAYOUT", "HND") == "HND" and t.dim() == 3:
|
||||
return t.permute(1, 0, 2)
|
||||
return t
|
||||
|
||||
|
||||
def to_bf16(t: torch.Tensor) -> torch.Tensor:
|
||||
"""Coerce to bfloat16, dequantizing fp8 / uint8-encoded fp8 if needed."""
|
||||
if t.dtype == torch.uint8:
|
||||
t = t.view(torch.float8_e4m3fn)
|
||||
if t.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
return t.to(torch.float32).to(torch.bfloat16)
|
||||
if t.dtype in (torch.bfloat16, torch.float16, torch.float32):
|
||||
return t
|
||||
return t.to(torch.bfloat16)
|
||||
|
||||
|
||||
def extract_kv_via_slot_mapping(
|
||||
kv_layer: torch.Tensor, slot_mapping: torch.Tensor
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Pull (keys, values) for the fresh tokens of one layer using slot_mapping.
|
||||
|
||||
`kv_layer` is vLLM's per-layer paged storage. Layout depends on attention
|
||||
backend: either `[2, num_blocks, block_size, H, D]` or `[num_blocks, 2,
|
||||
block_size, H, D]`. NHD is enforced via `VLLM_KV_CACHE_LAYOUT=NHD`.
|
||||
`slot_mapping` is the per-token slot index; entries `< 0` are padding.
|
||||
|
||||
Returned tensors stay on the GPU — the D2H copy is deferred to the
|
||||
writer thread (`tensor_to_wire_bytes` calls `.cpu()`) so it doesn't
|
||||
block forward of subsequent layers.
|
||||
"""
|
||||
if kv_layer.shape[0] == 2:
|
||||
k_all = to_nhd(kv_layer[0])
|
||||
v_all = to_nhd(kv_layer[1])
|
||||
else:
|
||||
k_all = to_nhd(kv_layer[:, 0])
|
||||
v_all = to_nhd(kv_layer[:, 1])
|
||||
k_flat = k_all.reshape(-1, *k_all.shape[-2:])
|
||||
v_flat = v_all.reshape(-1, *v_all.shape[-2:])
|
||||
valid = slot_mapping >= 0
|
||||
safe_sm = slot_mapping.clamp(min=0)
|
||||
keys = to_bf16(k_flat[safe_sm][valid])
|
||||
values = to_bf16(v_flat[safe_sm][valid])
|
||||
return keys, values
|
||||
|
||||
|
||||
def write_kv_layer_chunk(
|
||||
wfile: BinaryIO,
|
||||
layer_idx: int,
|
||||
keys: torch.Tensor,
|
||||
values: torch.Tensor,
|
||||
) -> None:
|
||||
"""Serialize one layer's NHD-shaped K/V to a `KVChunk` on the wire."""
|
||||
if keys.dim() == 4:
|
||||
keys = keys.reshape(-1, keys.shape[-2], keys.shape[-1])
|
||||
values = values.reshape(-1, values.shape[-2], values.shape[-1])
|
||||
num_tokens = int(keys.shape[0])
|
||||
n_heads = int(keys.shape[1])
|
||||
head_dim = int(keys.shape[2])
|
||||
write_kv_chunk(
|
||||
wfile,
|
||||
layer_idx=layer_idx,
|
||||
num_tokens=num_tokens,
|
||||
n_heads=n_heads,
|
||||
head_dim=head_dim,
|
||||
dtype=torch_dtype_to_wire(keys.dtype),
|
||||
keys=tensor_to_wire_bytes(keys),
|
||||
values=tensor_to_wire_bytes(values),
|
||||
)
|
||||
|
||||
|
||||
def arrays_to_blobs(arrays: list[torch.Tensor]) -> list[TensorBlob]:
|
||||
"""Convert torch tensors (CPU or GPU) to wire-ready `TensorBlob`s."""
|
||||
return [
|
||||
TensorBlob(
|
||||
dtype=torch_dtype_to_wire(arr.dtype),
|
||||
shape=tuple(int(d) for d in arr.shape),
|
||||
data=tensor_to_wire_bytes(arr),
|
||||
)
|
||||
for arr in arrays
|
||||
]
|
||||
|
||||
|
||||
def write_layer_arrays_blobs(
|
||||
wfile: BinaryIO,
|
||||
layer_idx: int,
|
||||
blobs: list[TensorBlob],
|
||||
) -> None:
|
||||
write_arrays_state(wfile, layer_idx, blobs)
|
||||
|
||||
|
||||
def write_layer_arrays(
|
||||
wfile: BinaryIO,
|
||||
layer_idx: int,
|
||||
arrays: list[torch.Tensor],
|
||||
) -> None:
|
||||
"""Serialize a layer's auxiliary state (SSM/conv) as `ArraysState`."""
|
||||
write_layer_arrays_blobs(wfile, layer_idx, arrays_to_blobs(arrays))
|
||||
|
||||
|
||||
def write_prefill_header(
|
||||
wfile: BinaryIO,
|
||||
*,
|
||||
request_id: str,
|
||||
model_id: str,
|
||||
num_layers: int,
|
||||
dtype: DType = "bfloat16",
|
||||
start_pos: int = 0,
|
||||
) -> None:
|
||||
write_header(
|
||||
wfile,
|
||||
Header(
|
||||
request_id=request_id,
|
||||
model_id=model_id,
|
||||
num_layers=num_layers,
|
||||
dtype=dtype,
|
||||
start_pos=start_pos,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def write_prefill_done(wfile: BinaryIO, total_tokens: int) -> None:
|
||||
write_done(wfile, total_tokens)
|
||||
|
||||
|
||||
def build_layer_to_group(kv_cache_config: KVCacheConfig) -> list[int]:
|
||||
"""Map each layer index (model_runner.kv_caches order) to its kv_cache group.
|
||||
|
||||
vLLM's hybrid models split layers across multiple KV cache groups (e.g.
|
||||
full attention vs sliding-window attention). `request_finished_all_groups`
|
||||
returns block_ids per group; we need this map to look up the right group
|
||||
when reading a layer's blocks.
|
||||
"""
|
||||
group_lookup: dict[str, int] = {}
|
||||
for group_idx, group_spec in enumerate(kv_cache_config.kv_cache_groups):
|
||||
for layer_name in group_spec.layer_names:
|
||||
group_lookup[layer_name] = group_idx
|
||||
|
||||
layer_to_group: list[int] = []
|
||||
for tensor_spec in kv_cache_config.kv_cache_tensors:
|
||||
for name in tensor_spec.shared_by:
|
||||
layer_to_group.append(group_lookup[name])
|
||||
return layer_to_group
|
||||
|
||||
|
||||
def gather_layer_kv_from_blocks(
|
||||
layer_kv: torch.Tensor,
|
||||
block_ids: list[int],
|
||||
num_tokens: int,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Read K and V for `num_tokens` from a layer's paged block storage.
|
||||
|
||||
Captures APC-cached blocks identically to freshly-computed ones — the
|
||||
block pool doesn't distinguish.
|
||||
|
||||
`layer_kv` shapes (NHD, set via `VLLM_KV_CACHE_LAYOUT=NHD`):
|
||||
- `[2, num_pool_blocks, block_size, n_kv_heads, head_dim]`, or
|
||||
- `[num_pool_blocks, 2, block_size, n_kv_heads, head_dim]`.
|
||||
Returns NHD-shaped K and V of shape `[num_tokens, n_kv_heads, head_dim]`
|
||||
on the same CUDA device as `layer_kv`. The caller is responsible for
|
||||
issuing the D2H copy on a side stream so the scheduler thread isn't
|
||||
blocked.
|
||||
"""
|
||||
if not block_ids:
|
||||
return torch.empty(0, device=layer_kv.device), torch.empty(
|
||||
0, device=layer_kv.device
|
||||
)
|
||||
block_idx_tensor = torch.tensor(block_ids, dtype=torch.long, device=layer_kv.device)
|
||||
if layer_kv.shape[0] == 2:
|
||||
# [2, blocks, block, H, D]
|
||||
gathered_k = layer_kv[0][block_idx_tensor]
|
||||
gathered_v = layer_kv[1][block_idx_tensor]
|
||||
else:
|
||||
# [blocks, 2, block, H, D]
|
||||
gathered = layer_kv[block_idx_tensor]
|
||||
gathered_k = gathered[:, 0]
|
||||
gathered_v = gathered[:, 1]
|
||||
# gathered_k/v: [num_blocks, block_size, H, D]. Concat blocks along seq.
|
||||
keys = gathered_k.reshape(-1, *gathered_k.shape[-2:])[:num_tokens]
|
||||
values = gathered_v.reshape(-1, *gathered_v.shape[-2:])[:num_tokens]
|
||||
return to_bf16(keys), to_bf16(values)
|
||||
@@ -0,0 +1,10 @@
|
||||
import pytest
|
||||
|
||||
|
||||
def pytest_addoption(parser: pytest.Parser) -> None:
|
||||
parser.addoption(
|
||||
"--model-id",
|
||||
action="store",
|
||||
default=None,
|
||||
help="HuggingFace-style model id (e.g. Qwen/Qwen3-0.6B) — must be downloaded",
|
||||
)
|
||||
@@ -0,0 +1,201 @@
|
||||
"""End-to-end test for VllmEngine.serve_prefill.
|
||||
|
||||
Boots a real vLLM engine with a small model, calls serve_prefill twice in a
|
||||
row against an in-memory wire buffer, and verifies both runs produce a
|
||||
well-formed stream (header -> KV chunks -> Done).
|
||||
|
||||
The second run is the regression case: with vLLM APC enabled this would hit
|
||||
the chunked-prefill + APC + custom kv-connector CUDA assert
|
||||
(`vectorized_gather_kernel: ind >= ind_dim_size`) and the server would close
|
||||
the socket without a Done frame. With APC disabled at engine creation time
|
||||
each request runs a full forward pass and the stream is well-formed.
|
||||
|
||||
Run on Spark (gx10-de89):
|
||||
cd /home/larry/exo
|
||||
uv run pytest -q -s -m "" \\
|
||||
src/exo/worker/engines/vllm/disaggregated/tests/test_serve_prefill_integration.py \\
|
||||
--model-id Qwen/Qwen3-0.6B
|
||||
|
||||
The test is gated on `--model-id` being passed; the model must already be
|
||||
present at `EXO_DEFAULT_MODELS_DIR/<id-with-/-as--->` (the standard exo
|
||||
download layout). On machines without CUDA / vLLM the test is skipped.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import io
|
||||
from collections.abc import Iterator
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.worker.disaggregated.protocol import (
|
||||
ArraysState,
|
||||
Done,
|
||||
ErrorMessage,
|
||||
KVChunk,
|
||||
read_header,
|
||||
read_message,
|
||||
)
|
||||
from exo.worker.disaggregated.server import PrefillRequest
|
||||
from exo.worker.engines.base import Engine
|
||||
from exo.worker.engines.vllm.engine import VllmEngine
|
||||
|
||||
|
||||
def _has_cuda() -> bool:
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
return False
|
||||
return bool(torch.cuda.is_available())
|
||||
|
||||
|
||||
def _make_token_ids(n: int) -> list[int]:
|
||||
# Deterministic synthetic tokens. Vocab >= ~30k for the Qwen tokenizers we
|
||||
# care about, so 100..30099 is safe.
|
||||
return [(i * 1009 + 17) % 30000 + 100 for i in range(n)]
|
||||
|
||||
|
||||
def _decode_stream(
|
||||
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:
|
||||
match msg := read_message(buf):
|
||||
case None:
|
||||
break
|
||||
case KVChunk():
|
||||
chunks.append(msg)
|
||||
case ArraysState():
|
||||
arrays.append(msg)
|
||||
case Done():
|
||||
done = msg
|
||||
break
|
||||
case ErrorMessage():
|
||||
error = msg
|
||||
break
|
||||
return chunks, arrays, done, error
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def vllm_engine(request: pytest.FixtureRequest) -> Iterator[object]:
|
||||
"""Build a real VllmEngine pointed at a downloaded HF model."""
|
||||
if not _has_cuda():
|
||||
pytest.skip("CUDA not available")
|
||||
model_id_str = cast(str, request.config.getoption("--model-id"))
|
||||
if not model_id_str:
|
||||
pytest.skip("pass --model-id <hf-id> to run this test")
|
||||
|
||||
model_id = ModelId(model_id_str)
|
||||
|
||||
from exo.download.download_utils import build_model_path
|
||||
|
||||
if not build_model_path(model_id).exists():
|
||||
pytest.skip(f"model {model_id} not downloaded locally")
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
# Mirror VllmBuilder.load() — patches must run before LLMEngine init.
|
||||
_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)
|
||||
|
||||
# serve_prefill only touches self._gen.engine; the channel fields exist
|
||||
# for the (unused-here) generation path.
|
||||
class _DummySender:
|
||||
def send(self, _: object) -> None: ...
|
||||
|
||||
class _DummyReceiver:
|
||||
def collect(self) -> list[object]:
|
||||
return []
|
||||
|
||||
engine = VllmEngine(
|
||||
tool_parser=tool_parser,
|
||||
model_id=model_id,
|
||||
cancel_receiver=_DummyReceiver(), # pyright: ignore[reportArgumentType]
|
||||
event_sender=_DummySender(), # pyright: ignore[reportArgumentType]
|
||||
_gen=gen,
|
||||
max_concurrent_requests=1,
|
||||
)
|
||||
try:
|
||||
yield engine
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
engine.close()
|
||||
|
||||
|
||||
def _run_one(engine: Engine, n_tokens: int, label: str) -> Done:
|
||||
request = PrefillRequest(
|
||||
request_id=f"itest-{label}",
|
||||
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)
|
||||
|
||||
payload = buf.getvalue()
|
||||
assert payload, f"{label}: server wrote nothing"
|
||||
|
||||
chunks, arrays, done, error = _decode_stream(payload)
|
||||
if error is not None:
|
||||
pytest.fail(
|
||||
f"{label}: server returned ErrorMessage [{error.code}]: {error.message}"
|
||||
)
|
||||
assert done is not None, (
|
||||
f"{label}: stream did not end with Done "
|
||||
f"(received {len(chunks)} kv chunks, {len(arrays)} arrays)"
|
||||
)
|
||||
expected = max(0, n_tokens - 2) # serve_prefill drops the last 2 tokens
|
||||
assert done.total_tokens > 0, f"{label}: Done reported 0 tokens"
|
||||
assert done.total_tokens >= expected - 64, (
|
||||
f"{label}: got {done.total_tokens} tokens, expected ~{expected}"
|
||||
)
|
||||
assert chunks, f"{label}: no KV chunks shipped"
|
||||
return done
|
||||
|
||||
|
||||
def test_serve_prefill_two_runs_no_apc_assert(vllm_engine: VllmEngine) -> None:
|
||||
"""Two consecutive prefills against the same engine must both succeed.
|
||||
|
||||
Before the fix, the second call hit a CUDA assert (vLLM APC + chunked
|
||||
prefill + custom kv-connector). With APC off, each request runs a full
|
||||
forward and the stream is well-formed both times.
|
||||
"""
|
||||
first = _run_one(vllm_engine, n_tokens=512, label="run1")
|
||||
second = _run_one(vllm_engine, n_tokens=512, label="run2-same-prompt")
|
||||
assert second.total_tokens == first.total_tokens, (
|
||||
f"run2 returned {second.total_tokens}, run1 returned {first.total_tokens}"
|
||||
)
|
||||
|
||||
|
||||
def test_serve_prefill_different_lengths(vllm_engine: VllmEngine) -> None:
|
||||
"""A second prefill with a different prompt length still succeeds."""
|
||||
a = _run_one(vllm_engine, n_tokens=256, label="run-256")
|
||||
b = _run_one(vllm_engine, n_tokens=768, label="run-768")
|
||||
assert b.total_tokens > a.total_tokens, (
|
||||
f"longer prompt should ship more tokens: 256->{a.total_tokens} 768->{b.total_tokens}"
|
||||
)
|
||||
@@ -0,0 +1,623 @@
|
||||
import contextlib
|
||||
import itertools
|
||||
import threading
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Generator, Iterator
|
||||
from dataclasses import dataclass, field
|
||||
from typing import BinaryIO
|
||||
|
||||
import torch
|
||||
from vllm import SamplingParams
|
||||
from vllm.outputs import RequestOutput
|
||||
|
||||
from exo.shared.constants import EXO_MAX_CONCURRENT_REQUESTS
|
||||
from exo.shared.types.chunks import ErrorChunk, GenerationChunk, PrefillProgressChunk
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.events import ChunkGenerated, Event
|
||||
from exo.shared.types.tasks import (
|
||||
CANCEL_ALL_TASKS,
|
||||
GenerationTask,
|
||||
TaskId,
|
||||
TextGeneration,
|
||||
)
|
||||
from exo.shared.types.text_generation import TextGenerationTaskParams
|
||||
from exo.shared.types.worker.runner_response import (
|
||||
CancelledResponse,
|
||||
FinishedResponse,
|
||||
GenerationResponse,
|
||||
)
|
||||
from exo.utils.channels import MpReceiver, MpSender
|
||||
from exo.worker.disaggregated.protocol import write_error, write_kv_chunk
|
||||
from exo.worker.disaggregated.server import PrefillRequest
|
||||
from exo.worker.engines.base import Engine
|
||||
from exo.worker.engines.vllm.disaggregated.adapter import (
|
||||
arrays_to_blobs,
|
||||
tensor_to_wire_bytes,
|
||||
torch_dtype_to_wire,
|
||||
write_layer_arrays_blobs,
|
||||
write_prefill_done,
|
||||
write_prefill_header,
|
||||
)
|
||||
from exo.worker.engines.vllm.generator import VllmBatchEngine
|
||||
from exo.worker.engines.vllm.growable_cache import get_model_runner
|
||||
from exo.worker.engines.vllm.kv_connector import (
|
||||
get_arrays_queue,
|
||||
get_gdn_shipped,
|
||||
get_gdn_states,
|
||||
get_kv_queue,
|
||||
get_save_kv_layer_diag,
|
||||
init_gdn_layer_order,
|
||||
reset_capture_state,
|
||||
)
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
from exo.worker.runner.llm_inference.model_output_parsers import (
|
||||
apply_all_parsers,
|
||||
map_responses_to_chunks,
|
||||
)
|
||||
from exo.worker.runner.llm_inference.tool_parsers import ToolParser
|
||||
|
||||
|
||||
class GeneratorQueue[T]:
|
||||
def __init__(self) -> None:
|
||||
self._q = deque[T]()
|
||||
|
||||
def push(self, t: T) -> None:
|
||||
self._q.append(t)
|
||||
|
||||
def gen(self) -> Generator[T | None]:
|
||||
while True:
|
||||
if len(self._q) == 0:
|
||||
yield None
|
||||
else:
|
||||
yield self._q.popleft()
|
||||
|
||||
|
||||
EXO_RUNNER_MUST_FAIL = "EXO RUNNER MUST FAIL"
|
||||
EXO_RUNNER_MUST_TIMEOUT = "EXO RUNNER MUST TIMEOUT"
|
||||
|
||||
|
||||
def _check_for_debug_prompts(task_params: TextGenerationTaskParams) -> None:
|
||||
"""Keep the cheap debug prompt hooks without importing the MLX engine."""
|
||||
if len(task_params.input) == 0:
|
||||
return
|
||||
prompt = task_params.input[0].content
|
||||
if not prompt:
|
||||
return
|
||||
if EXO_RUNNER_MUST_FAIL in prompt:
|
||||
raise Exception("Artificial runner exception - for testing purposes only.")
|
||||
if EXO_RUNNER_MUST_TIMEOUT in prompt:
|
||||
time.sleep(100)
|
||||
|
||||
|
||||
@dataclass(eq=False)
|
||||
class VllmEngine(Engine):
|
||||
"""Single-node vLLM implementation of the exo Engine interface.
|
||||
|
||||
This intentionally duplicates the local orchestration from the MLX
|
||||
BatchGenerator instead of trying to share a batch abstraction too early.
|
||||
The vLLM-specific tokenization/sampling/stepping remains inside
|
||||
VllmBatchEngine.
|
||||
"""
|
||||
|
||||
tool_parser: ToolParser | None
|
||||
model_id: ModelId
|
||||
cancel_receiver: MpReceiver[TaskId]
|
||||
event_sender: MpSender[Event]
|
||||
_gen: VllmBatchEngine
|
||||
max_concurrent_requests: int = EXO_MAX_CONCURRENT_REQUESTS
|
||||
check_for_cancel_every: int = 50
|
||||
|
||||
_cancelled_tasks: set[TaskId] = field(default_factory=set, init=False)
|
||||
_all_tasks: dict[TaskId, TextGeneration] = field(default_factory=dict, init=False)
|
||||
_queue: deque[TextGeneration] = field(default_factory=deque, init=False)
|
||||
_active_tasks: dict[
|
||||
TaskId,
|
||||
tuple[
|
||||
TextGeneration,
|
||||
GeneratorQueue[GenerationResponse],
|
||||
Iterator[GenerationChunk | None],
|
||||
],
|
||||
] = field(default_factory=dict, init=False)
|
||||
|
||||
def warmup(self) -> None:
|
||||
self.check_for_cancel_every = self._gen.warmup()
|
||||
|
||||
def submit(self, task: GenerationTask) -> None:
|
||||
assert isinstance(task, TextGeneration)
|
||||
self._cancelled_tasks.discard(CANCEL_ALL_TASKS)
|
||||
self._all_tasks[task.task_id] = task
|
||||
self._queue.append(task)
|
||||
|
||||
def step(
|
||||
self,
|
||||
) -> Iterator[
|
||||
tuple[TaskId, GenerationChunk | CancelledResponse | FinishedResponse]
|
||||
]:
|
||||
self._collect_cancellations()
|
||||
output: list[
|
||||
tuple[TaskId, GenerationChunk | CancelledResponse | FinishedResponse]
|
||||
] = list(self._apply_cancellations())
|
||||
|
||||
while self._queue and len(self._active_tasks) < self.max_concurrent_requests:
|
||||
task = self._queue.popleft()
|
||||
if self.should_cancel(task.task_id):
|
||||
output.append((task.task_id, CancelledResponse()))
|
||||
self._all_tasks.pop(task.task_id, None)
|
||||
continue
|
||||
|
||||
try:
|
||||
task_id, queue, output_generator = self._start_task(task)
|
||||
except Exception as e:
|
||||
self._send_error(task, e)
|
||||
self._all_tasks.pop(task.task_id, None)
|
||||
raise
|
||||
|
||||
self._active_tasks[task_id] = (task, queue, output_generator)
|
||||
|
||||
if not self._gen.has_work:
|
||||
return iter(output)
|
||||
|
||||
results = self._gen.step()
|
||||
for task_id, response in results:
|
||||
if task_id not in self._active_tasks:
|
||||
logger.warning(f"{task_id=} not found in active vLLM tasks")
|
||||
continue
|
||||
|
||||
task, queue, output_generator = self._active_tasks[task_id]
|
||||
queue.push(response)
|
||||
while (parsed := next(output_generator, None)) is not None:
|
||||
output.append((task.task_id, parsed))
|
||||
|
||||
if response.finish_reason is not None:
|
||||
output.append((task.task_id, FinishedResponse()))
|
||||
del self._active_tasks[task_id]
|
||||
self._all_tasks.pop(task.task_id, None)
|
||||
|
||||
return itertools.chain(output, self._apply_cancellations())
|
||||
|
||||
def close(self) -> None:
|
||||
self._gen.close()
|
||||
|
||||
def serve_prefill(self, request: PrefillRequest, wfile: BinaryIO) -> None:
|
||||
engine = self._gen.engine
|
||||
if engine.has_unfinished_requests():
|
||||
logger.warning("serve_prefill: engine busy, refusing prefill request")
|
||||
write_error(wfile, code=503, message="engine busy")
|
||||
return
|
||||
|
||||
model_runner = get_model_runner()
|
||||
if model_runner is None:
|
||||
logger.warning("serve_prefill: model runner not initialized")
|
||||
write_error(wfile, code=503, message="model runner not initialized")
|
||||
return
|
||||
|
||||
init_gdn_layer_order(model_runner.kv_caches)
|
||||
|
||||
prefill_token_ids = (
|
||||
request.token_ids[:-2]
|
||||
if len(request.token_ids) > 2
|
||||
else list(request.token_ids)
|
||||
)
|
||||
n_layers = len(model_runner.kv_caches)
|
||||
|
||||
reset_capture_state()
|
||||
arrays_queue = get_arrays_queue()
|
||||
kv_queue = get_kv_queue()
|
||||
|
||||
# We strip the trailing 2 tokens because the consumer warm-starts
|
||||
# decode from them locally.
|
||||
sp = SamplingParams(max_tokens=2, temperature=0.0, detokenize=False)
|
||||
engine.add_request(
|
||||
request.request_id,
|
||||
{"prompt_token_ids": prefill_token_ids},
|
||||
sp,
|
||||
)
|
||||
|
||||
write_prefill_header(
|
||||
wfile,
|
||||
request_id=request.request_id,
|
||||
model_id=request.model_id,
|
||||
num_layers=n_layers,
|
||||
start_pos=request.start_pos,
|
||||
)
|
||||
|
||||
skip_tokens = request.start_pos
|
||||
chunks_sent = 0
|
||||
arrays_streamed = 0
|
||||
layer_token_counts: dict[int, int] = {}
|
||||
# Both writer threads serialize through this lock — BufferedWriter
|
||||
# is not thread-safe and we don't want partial frames interleaved.
|
||||
wfile_lock = threading.Lock()
|
||||
# Diag for end-of-request bandwidth report.
|
||||
writer_stats = {
|
||||
"bytes_shipped": 0,
|
||||
"wait_event_secs": 0.0,
|
||||
"socket_secs": 0.0,
|
||||
"first_byte_t": 0.0,
|
||||
"last_byte_t": 0.0,
|
||||
"started_t": 0.0,
|
||||
}
|
||||
|
||||
def writer_loop() -> None:
|
||||
nonlocal chunks_sent
|
||||
writer_stats["started_t"] = time.perf_counter()
|
||||
last_hb = time.perf_counter()
|
||||
try:
|
||||
while True:
|
||||
try:
|
||||
item = kv_queue.get(timeout=3.0)
|
||||
except Exception:
|
||||
item = ... # sentinel for "no item yet"
|
||||
if item is ...:
|
||||
now = time.perf_counter()
|
||||
if now - last_hb > 3.0:
|
||||
logger.info(
|
||||
f"serve_prefill writer idle: "
|
||||
f"chunks_sent={chunks_sent} "
|
||||
f"kv_queue_size={kv_queue.qsize()} "
|
||||
f"arrays_queue_size={arrays_queue.qsize()}"
|
||||
)
|
||||
last_hb = now
|
||||
continue
|
||||
if item is None:
|
||||
break
|
||||
layer_idx, count, keys, values, copy_event = item
|
||||
# Wait for the side-stream D2H to finish populating the
|
||||
# pinned host buffers. CPU-side wait, doesn't block GPU.
|
||||
t_wait = time.perf_counter()
|
||||
copy_event.synchronize()
|
||||
writer_stats["wait_event_secs"] += time.perf_counter() - t_wait
|
||||
previous = layer_token_counts.get(layer_idx, 0)
|
||||
new_total = previous + count
|
||||
layer_token_counts[layer_idx] = new_total
|
||||
|
||||
if new_total <= skip_tokens:
|
||||
continue
|
||||
# Reshape paged 4-D layouts to per-token 3-D up front so
|
||||
# the trim slice operates on the token axis.
|
||||
if keys.dim() == 4:
|
||||
keys = keys.reshape(-1, keys.shape[-2], keys.shape[-1])
|
||||
values = values.reshape(-1, values.shape[-2], values.shape[-1])
|
||||
# Slice keys/values to exactly `count` tokens — the source
|
||||
# tensor may be larger if shape disagrees with logical
|
||||
# token count (e.g. paged storage gathered over more
|
||||
# blocks than tokens consumed).
|
||||
if int(keys.shape[0]) > count:
|
||||
keys = keys[:count]
|
||||
values = values[:count]
|
||||
if previous < skip_tokens:
|
||||
trim = skip_tokens - previous
|
||||
keys = keys[trim:]
|
||||
values = values[trim:]
|
||||
num_tokens = int(keys.shape[0])
|
||||
n_heads = int(keys.shape[1])
|
||||
head_dim = int(keys.shape[2])
|
||||
dtype_w = torch_dtype_to_wire(keys.dtype)
|
||||
keys_bytes = tensor_to_wire_bytes(keys)
|
||||
values_bytes = tensor_to_wire_bytes(values)
|
||||
payload_bytes = len(keys_bytes) + len(values_bytes)
|
||||
if chunks_sent == 0:
|
||||
writer_stats["first_byte_t"] = time.perf_counter()
|
||||
logger.info(
|
||||
f"First KV chunk: layer={layer_idx} keys={keys.shape} "
|
||||
f"keys.dtype={keys.dtype} values.dtype={values.dtype}"
|
||||
)
|
||||
t_sock = time.perf_counter()
|
||||
with wfile_lock:
|
||||
write_kv_chunk(
|
||||
wfile,
|
||||
layer_idx=layer_idx,
|
||||
num_tokens=num_tokens,
|
||||
n_heads=n_heads,
|
||||
head_dim=head_dim,
|
||||
dtype=dtype_w,
|
||||
keys=keys_bytes,
|
||||
values=values_bytes,
|
||||
)
|
||||
writer_stats["socket_secs"] += time.perf_counter() - t_sock
|
||||
writer_stats["bytes_shipped"] += payload_bytes
|
||||
writer_stats["last_byte_t"] = time.perf_counter()
|
||||
chunks_sent += 1
|
||||
except Exception:
|
||||
logger.opt(exception=True).warning(
|
||||
"serve_prefill writer thread crashed"
|
||||
)
|
||||
|
||||
def arrays_writer_loop() -> None:
|
||||
nonlocal arrays_streamed
|
||||
try:
|
||||
while True:
|
||||
item = arrays_queue.get()
|
||||
if item is None:
|
||||
break
|
||||
layer_idx, arrays, copy_event = item
|
||||
if copy_event is not None:
|
||||
copy_event.synchronize()
|
||||
with wfile_lock:
|
||||
write_layer_arrays_blobs(
|
||||
wfile, layer_idx, arrays_to_blobs(arrays)
|
||||
)
|
||||
arrays_streamed += 1
|
||||
except Exception:
|
||||
logger.opt(exception=True).warning(
|
||||
"serve_prefill arrays writer thread crashed"
|
||||
)
|
||||
|
||||
writer_thread = threading.Thread(target=writer_loop, daemon=True)
|
||||
writer_thread.start()
|
||||
arrays_writer_thread = threading.Thread(target=arrays_writer_loop, daemon=True)
|
||||
arrays_writer_thread.start()
|
||||
|
||||
t0 = time.perf_counter()
|
||||
forward_error: Exception | None = None
|
||||
step_count = 0
|
||||
last_step_log = time.perf_counter()
|
||||
first_output_logged = False
|
||||
try:
|
||||
while engine.has_unfinished_requests():
|
||||
outputs = engine.step()
|
||||
step_count += 1
|
||||
now = time.perf_counter()
|
||||
if now - last_step_log > 3.0:
|
||||
logger.info(
|
||||
f"serve_prefill {request.request_id}: "
|
||||
f"step #{step_count} (kv_queue={kv_queue.qsize()})"
|
||||
)
|
||||
last_step_log = now
|
||||
aborted = False
|
||||
for output in outputs:
|
||||
if not first_output_logged:
|
||||
first_output_logged = True
|
||||
logger.info(
|
||||
f"serve_prefill {request.request_id}: first output "
|
||||
f"id={output.request_id!r} "
|
||||
f"tokens={len(output.outputs[0].token_ids) if isinstance(output, RequestOutput) and output.outputs else 0}"
|
||||
)
|
||||
# Match either the external id we passed or the
|
||||
# internal-suffixed id vLLM may surface ('-XXXXXXXX').
|
||||
if (
|
||||
isinstance(output, RequestOutput)
|
||||
and (
|
||||
output.request_id == request.request_id
|
||||
or output.request_id.startswith(request.request_id)
|
||||
)
|
||||
and output.outputs
|
||||
and output.outputs[0].token_ids
|
||||
):
|
||||
engine.abort_request([request.request_id])
|
||||
aborted = True
|
||||
break
|
||||
if aborted:
|
||||
break
|
||||
# Post-abort drain. Bail out hard after 5s — if the request
|
||||
# didn't finish by then something is wrong upstream and we'd
|
||||
# otherwise spin forever in a no-op step loop.
|
||||
drain_deadline = time.perf_counter() + 5.0
|
||||
while engine.has_unfinished_requests():
|
||||
if time.perf_counter() > drain_deadline:
|
||||
logger.warning(
|
||||
f"serve_prefill {request.request_id}: post-abort drain "
|
||||
f"timeout, force-aborting and breaking out"
|
||||
)
|
||||
with contextlib.suppress(Exception):
|
||||
engine.abort_request([request.request_id])
|
||||
break
|
||||
_ = engine.step()
|
||||
step_count += 1
|
||||
except Exception as exc:
|
||||
forward_error = exc
|
||||
with contextlib.suppress(Exception):
|
||||
engine.abort_request([request.request_id])
|
||||
finally:
|
||||
logger.info(
|
||||
f"serve_prefill {request.request_id}: "
|
||||
f"kv_queue={kv_queue.qsize()} arrays_queue={arrays_queue.qsize()}"
|
||||
)
|
||||
kv_queue.put(None)
|
||||
arrays_queue.put(None)
|
||||
writer_thread.join(timeout=30)
|
||||
arrays_writer_thread.join(timeout=30)
|
||||
if writer_thread.is_alive():
|
||||
logger.warning("serve_prefill: kv writer thread did not exit")
|
||||
if arrays_writer_thread.is_alive():
|
||||
logger.warning("serve_prefill: arrays writer thread did not exit")
|
||||
|
||||
if forward_error is not None:
|
||||
logger.opt(exception=forward_error).error(
|
||||
f"serve_prefill {request.request_id}: engine.step() raised"
|
||||
)
|
||||
with contextlib.suppress(Exception):
|
||||
write_error(wfile, code=500, message=f"engine.step: {forward_error!r}")
|
||||
return
|
||||
|
||||
# The K/V writer and arrays writer both drained their queues during
|
||||
# forward (see writer_loop / arrays_writer_loop above). What remains
|
||||
# here is the fallback for any GDN layer whose conv+ssm pair never
|
||||
# reached `_try_ship_gdn` — e.g., ssm captured but not conv. We skip
|
||||
# layers already shipped by the streaming path.
|
||||
gdn = get_gdn_states()
|
||||
gdn_shipped = get_gdn_shipped()
|
||||
unshipped = [li for li in sorted(gdn.keys()) if li not in gdn_shipped]
|
||||
arrays_layers = arrays_streamed
|
||||
if unshipped:
|
||||
torch.cuda.synchronize()
|
||||
for layer_idx in unshipped:
|
||||
state = gdn[layer_idx]
|
||||
arrs: list[torch.Tensor] = []
|
||||
if "conv" in state:
|
||||
arrs.append(state["conv"])
|
||||
if "ssm" in state:
|
||||
arrs.append(state["ssm"])
|
||||
if arrs:
|
||||
write_layer_arrays_blobs(wfile, layer_idx, arrays_to_blobs(arrs))
|
||||
arrays_layers += 1
|
||||
|
||||
forwarded_per_layer = max(layer_token_counts.values(), default=0)
|
||||
tokens_sent = max(0, forwarded_per_layer - skip_tokens)
|
||||
write_prefill_done(wfile, tokens_sent)
|
||||
elapsed = time.perf_counter() - t0
|
||||
diag = get_save_kv_layer_diag()
|
||||
diag_summary = ", ".join(
|
||||
f"L{li}:{','.join(str(s) for s in sizes)}"
|
||||
for li, sizes in sorted(diag.items())
|
||||
)
|
||||
logger.info(
|
||||
f"serve_prefill {request.request_id}: save_kv_layer calls per layer "
|
||||
f"(positive=non-list/tuple kv, negative=list/tuple kv) → {diag_summary}"
|
||||
)
|
||||
logger.info(
|
||||
f"serve_prefill {request.request_id}: layer_token_counts="
|
||||
f"{dict(sorted(layer_token_counts.items()))}"
|
||||
)
|
||||
# Bandwidth + per-stage breakdown for the writer thread.
|
||||
bytes_shipped = writer_stats["bytes_shipped"]
|
||||
wait_secs = writer_stats["wait_event_secs"]
|
||||
sock_secs = writer_stats["socket_secs"]
|
||||
first_byte_dt = (
|
||||
writer_stats["first_byte_t"] - t0 if writer_stats["first_byte_t"] else 0.0
|
||||
)
|
||||
ship_secs = (
|
||||
writer_stats["last_byte_t"] - writer_stats["first_byte_t"]
|
||||
if writer_stats["last_byte_t"]
|
||||
else 0.0
|
||||
)
|
||||
eff_bw_mbps = (bytes_shipped / 1e6 / ship_secs) if ship_secs > 0 else 0.0
|
||||
peak_bw_mbps = (bytes_shipped / 1e6 / sock_secs) if sock_secs > 0 else 0.0
|
||||
logger.info(
|
||||
f"serve_prefill {request.request_id}: "
|
||||
f"streamed_chunks={chunks_sent} arrays_layers={arrays_layers} "
|
||||
f"tokens={tokens_sent} elapsed_ms={elapsed * 1000:.0f} "
|
||||
f"bytes={bytes_shipped / 1e6:.0f}MB ttfb_ms={first_byte_dt * 1000:.0f} "
|
||||
f"ship_ms={ship_secs * 1000:.0f} "
|
||||
f"wait_event_ms={wait_secs * 1000:.0f} sock_ms={sock_secs * 1000:.0f} "
|
||||
f"eff_bw={eff_bw_mbps:.0f}MB/s peak_bw={peak_bw_mbps:.0f}MB/s"
|
||||
)
|
||||
|
||||
def _start_task(
|
||||
self, task: TextGeneration
|
||||
) -> tuple[
|
||||
TaskId,
|
||||
GeneratorQueue[GenerationResponse],
|
||||
Iterator[GenerationChunk | None],
|
||||
]:
|
||||
from exo.worker.engines.vllm.prompt_format import format_vllm_prompt
|
||||
|
||||
_check_for_debug_prompts(task.task_params)
|
||||
|
||||
token_ids, prompt_text, _ = format_vllm_prompt(
|
||||
self._gen.engine, task.task_params
|
||||
)
|
||||
|
||||
queue = GeneratorQueue[GenerationResponse]()
|
||||
if task.task_params.bench:
|
||||
output_generator: Iterator[GenerationChunk | None] = map(
|
||||
lambda r: map_responses_to_chunks(r, self.model_id), queue.gen()
|
||||
)
|
||||
else:
|
||||
from mlx_lm.tokenizer_utils import TokenizerWrapper
|
||||
|
||||
output_generator = apply_all_parsers(
|
||||
queue.gen(),
|
||||
prompt_text,
|
||||
self.tool_parser,
|
||||
TokenizerWrapper(self._gen.engine.get_tokenizer()),
|
||||
self.model_id,
|
||||
task.task_params.tools,
|
||||
)
|
||||
|
||||
check_for_cancel_every = max(self.check_for_cancel_every, 1)
|
||||
tokens_since_cancel_check = check_for_cancel_every
|
||||
|
||||
def on_prefill_progress(processed: int, total: int) -> None:
|
||||
self._collect_cancellations()
|
||||
if self.should_cancel(task.task_id):
|
||||
self._cancelled_tasks.add(task.task_id)
|
||||
self.event_sender.send(
|
||||
ChunkGenerated(
|
||||
command_id=task.command_id,
|
||||
chunk=PrefillProgressChunk(
|
||||
model=self.model_id,
|
||||
processed_tokens=processed,
|
||||
total_tokens=total,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
def on_generation_token() -> None:
|
||||
nonlocal tokens_since_cancel_check
|
||||
tokens_since_cancel_check += 1
|
||||
if tokens_since_cancel_check >= check_for_cancel_every:
|
||||
tokens_since_cancel_check = 0
|
||||
self._collect_cancellations()
|
||||
if self.should_cancel(task.task_id):
|
||||
self._cancelled_tasks.add(task.task_id)
|
||||
|
||||
task_id = self._gen.submit(
|
||||
task_id=task.task_id,
|
||||
task_params=task.task_params,
|
||||
on_prefill_progress=on_prefill_progress,
|
||||
on_generation_token=on_generation_token,
|
||||
token_ids=token_ids,
|
||||
)
|
||||
return task_id, queue, output_generator
|
||||
|
||||
def _collect_cancellations(self) -> None:
|
||||
for task_id in self.cancel_receiver.collect():
|
||||
if task_id == CANCEL_ALL_TASKS:
|
||||
self._cancelled_tasks.add(CANCEL_ALL_TASKS)
|
||||
elif task_id in self._all_tasks:
|
||||
self._cancelled_tasks.add(task_id)
|
||||
|
||||
def _apply_cancellations(self) -> Iterator[tuple[TaskId, CancelledResponse]]:
|
||||
if not self._cancelled_tasks:
|
||||
return iter([])
|
||||
|
||||
cancel_all = CANCEL_ALL_TASKS in self._cancelled_tasks
|
||||
results: list[tuple[TaskId, CancelledResponse]] = []
|
||||
task_ids_to_abort: list[TaskId] = []
|
||||
|
||||
for task_id, (task, _, _) in list(self._active_tasks.items()):
|
||||
if cancel_all or task.task_id in self._cancelled_tasks:
|
||||
task_ids_to_abort.append(task_id)
|
||||
results.append((task.task_id, CancelledResponse()))
|
||||
del self._active_tasks[task_id]
|
||||
self._all_tasks.pop(task.task_id, None)
|
||||
|
||||
if self._queue:
|
||||
kept_queue: deque[TextGeneration] = deque()
|
||||
for task in self._queue:
|
||||
if cancel_all or task.task_id in self._cancelled_tasks:
|
||||
results.append((task.task_id, CancelledResponse()))
|
||||
self._all_tasks.pop(task.task_id, None)
|
||||
else:
|
||||
kept_queue.append(task)
|
||||
self._queue = kept_queue
|
||||
|
||||
if task_ids_to_abort:
|
||||
self._gen.cancel(task_ids_to_abort)
|
||||
|
||||
already_cancelled = {task_id for task_id, _ in results}
|
||||
for task_id in self._cancelled_tasks:
|
||||
if (
|
||||
task_id != CANCEL_ALL_TASKS
|
||||
and task_id in self._all_tasks
|
||||
and task_id not in already_cancelled
|
||||
):
|
||||
results.append((task_id, CancelledResponse()))
|
||||
self._all_tasks.pop(task_id, None)
|
||||
|
||||
self._cancelled_tasks.clear()
|
||||
return iter(results)
|
||||
|
||||
def _send_error(self, task: TextGeneration, e: Exception) -> None:
|
||||
self.event_sender.send(
|
||||
ChunkGenerated(
|
||||
command_id=task.command_id,
|
||||
chunk=ErrorChunk(
|
||||
model=self.model_id,
|
||||
finish_reason="error",
|
||||
error_message=str(e),
|
||||
),
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,490 @@
|
||||
import gc
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Callable, Generator
|
||||
from dataclasses import dataclass, field
|
||||
from typing import cast
|
||||
|
||||
import torch
|
||||
from vllm.config import CompilationConfig
|
||||
from vllm.config.compilation import CompilationMode, CUDAGraphMode
|
||||
from vllm.engine.arg_utils import EngineArgs
|
||||
from vllm.entrypoints.chat_utils import (
|
||||
ChatCompletionMessageParam,
|
||||
CustomChatCompletionMessageParam,
|
||||
)
|
||||
from vllm.outputs import RequestOutput
|
||||
from vllm.sampling_params import SamplingParams
|
||||
from vllm.tokenizers import TokenizerLike
|
||||
from vllm.v1.attention.backends.registry import AttentionBackendEnum
|
||||
from vllm.v1.engine.llm_engine import LLMEngine
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig
|
||||
|
||||
from exo.api.types import (
|
||||
CompletionTokensDetails,
|
||||
GenerationStats,
|
||||
PromptTokensDetails,
|
||||
Usage,
|
||||
)
|
||||
from exo.download.download_utils import build_model_path
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.tasks import TaskId
|
||||
from exo.shared.types.text_generation import TextGenerationTaskParams
|
||||
from exo.shared.types.worker.runner_response import GenerationResponse
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
from exo.worker.runner.llm_inference.tool_parsers import ToolParser, infer_tool_parser
|
||||
|
||||
|
||||
@dataclass
|
||||
class _EngineRequest:
|
||||
request_id: str
|
||||
prompt_token_count: int
|
||||
prefill_done: bool = False
|
||||
prefill_steps: int = 0
|
||||
prev_text: str = ""
|
||||
prev_token_count: int = 0
|
||||
start_time: float = field(default_factory=time.perf_counter)
|
||||
first_token_time: float | None = None
|
||||
on_generation_token: Callable[[], None] | None = None
|
||||
on_prefill_progress: Callable[[int, int], None] | None = None
|
||||
|
||||
|
||||
def _stop_token_ids(tokenizer: TokenizerLike, model_id: ModelId) -> set[int]:
|
||||
from exo.worker.engines.mlx.utils_mlx import get_eos_token_ids_for_model
|
||||
|
||||
ids: set[int] = set()
|
||||
eos_id = getattr(tokenizer, "eos_token_id", None)
|
||||
if eos_id is not None:
|
||||
ids.add(eos_id) # pyright: ignore[reportAny]
|
||||
extra = get_eos_token_ids_for_model(model_id)
|
||||
if extra:
|
||||
ids.update(extra)
|
||||
return ids
|
||||
|
||||
|
||||
def _build_generation_response(
|
||||
tokenizer: TokenizerLike,
|
||||
token_id: int,
|
||||
finish_reason: str | None,
|
||||
prompt_token_count: int,
|
||||
completion_tokens: int,
|
||||
start_time: float,
|
||||
first_token_time: float | None,
|
||||
suppress_text: bool = False,
|
||||
) -> GenerationResponse:
|
||||
token_text: str = "" if suppress_text else tokenizer.decode([token_id])
|
||||
finish_usage: Usage | None = None
|
||||
finish_stats: GenerationStats | None = None
|
||||
mapped_finish_reason: str | None = None
|
||||
if finish_reason:
|
||||
now = time.perf_counter()
|
||||
prefill_elapsed = (first_token_time or now) - start_time
|
||||
decode_elapsed = now - (first_token_time or now)
|
||||
finish_usage = Usage(
|
||||
prompt_tokens=prompt_token_count,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_token_count + completion_tokens,
|
||||
prompt_tokens_details=PromptTokensDetails(),
|
||||
completion_tokens_details=CompletionTokensDetails(),
|
||||
)
|
||||
finish_stats = GenerationStats(
|
||||
prompt_tps=prompt_token_count / prefill_elapsed
|
||||
if prefill_elapsed > 0
|
||||
else 0.0,
|
||||
generation_tps=completion_tokens / decode_elapsed
|
||||
if decode_elapsed > 0
|
||||
else 0.0,
|
||||
prompt_tokens=prompt_token_count,
|
||||
generation_tokens=completion_tokens,
|
||||
peak_memory_usage=Memory.from_bytes(torch.cuda.max_memory_allocated()),
|
||||
)
|
||||
mapped_finish_reason = (
|
||||
finish_reason
|
||||
if finish_reason in ("stop", "length", "content_filter")
|
||||
else "stop"
|
||||
)
|
||||
return GenerationResponse(
|
||||
text=token_text,
|
||||
token=token_id,
|
||||
finish_reason=mapped_finish_reason,
|
||||
usage=finish_usage,
|
||||
stats=finish_stats,
|
||||
)
|
||||
|
||||
|
||||
def warmup_vllm_engine(engine: LLMEngine) -> int:
|
||||
tokenizer = engine.get_tokenizer()
|
||||
messages = [
|
||||
cast(
|
||||
ChatCompletionMessageParam,
|
||||
CustomChatCompletionMessageParam(
|
||||
role="user",
|
||||
content="Prompt to warm up the inference engine. Repeat this.",
|
||||
),
|
||||
)
|
||||
]
|
||||
prompt_text: str | list[int] = tokenizer.apply_chat_template( # pyright: ignore[reportUnknownMemberType]
|
||||
messages, tokenize=False, add_generation_prompt=True
|
||||
)
|
||||
if isinstance(prompt_text, list):
|
||||
token_ids = prompt_text
|
||||
else:
|
||||
token_ids: list[int] = tokenizer.encode(prompt_text, add_special_tokens=False)
|
||||
|
||||
params = SamplingParams(max_tokens=50, detokenize=False)
|
||||
engine.add_request("warmup", {"prompt_token_ids": token_ids}, params)
|
||||
t = time.monotonic()
|
||||
tokens_generated = 0
|
||||
while engine.has_unfinished_requests():
|
||||
engine.step()
|
||||
tokens_generated += 1
|
||||
elapsed = max(time.monotonic() - t, 0.001)
|
||||
check_for_cancel_every = min(math.ceil(tokens_generated / elapsed), 100)
|
||||
logger.info(
|
||||
f"vLLM warmup complete, check_for_cancel_every={check_for_cancel_every}"
|
||||
)
|
||||
return check_for_cancel_every
|
||||
|
||||
|
||||
@dataclass(eq=False)
|
||||
class VllmBatchEngine:
|
||||
engine: LLMEngine
|
||||
model_id: ModelId
|
||||
|
||||
_active: dict[TaskId, _EngineRequest] = field(default_factory=dict, init=False)
|
||||
|
||||
def warmup(self) -> int:
|
||||
return warmup_vllm_engine(self.engine)
|
||||
|
||||
@property
|
||||
def has_work(self) -> bool:
|
||||
return bool(self._active) or self.engine.has_unfinished_requests()
|
||||
|
||||
def submit(
|
||||
self,
|
||||
task_id: TaskId,
|
||||
task_params: TextGenerationTaskParams,
|
||||
token_ids: list[int],
|
||||
on_prefill_progress: Callable[[int, int], None] | None = None,
|
||||
on_generation_token: Callable[[], None] | None = None,
|
||||
) -> TaskId:
|
||||
from exo.worker.engines.vllm.prompt_format import make_vllm_sampling_params
|
||||
|
||||
sampling_params = make_vllm_sampling_params(
|
||||
self.engine, task_params, self.model_id
|
||||
)
|
||||
self.engine.add_request(
|
||||
task_id, {"prompt_token_ids": token_ids}, sampling_params
|
||||
)
|
||||
self._active[task_id] = _EngineRequest(
|
||||
request_id=task_id,
|
||||
prompt_token_count=len(token_ids),
|
||||
on_generation_token=on_generation_token,
|
||||
on_prefill_progress=on_prefill_progress,
|
||||
)
|
||||
return task_id
|
||||
|
||||
def step(self) -> list[tuple[TaskId, GenerationResponse]]:
|
||||
if not self.has_work:
|
||||
return []
|
||||
|
||||
outputs = self.engine.step()
|
||||
tokenizer = self.engine.get_tokenizer()
|
||||
stop_ids = _stop_token_ids(tokenizer, self.model_id)
|
||||
max_batch_tokens: int = (
|
||||
getattr(self.engine.model_config, "max_num_batched_tokens", 2048) or 2048
|
||||
)
|
||||
results: list[tuple[TaskId, GenerationResponse]] = []
|
||||
|
||||
for output in outputs:
|
||||
# todo: PoolingRequestOutputs
|
||||
assert isinstance(output, RequestOutput)
|
||||
task_id = TaskId(output.request_id)
|
||||
if task_id not in self._active:
|
||||
continue
|
||||
req = self._active[task_id]
|
||||
completion = output.outputs[0]
|
||||
new_token_count = len(completion.token_ids)
|
||||
new_tokens = completion.token_ids[req.prev_token_count :]
|
||||
finish_reason = completion.finish_reason
|
||||
req.prev_token_count = new_token_count
|
||||
|
||||
if not req.prefill_done and not new_tokens:
|
||||
req.prefill_steps += 1
|
||||
if req.on_prefill_progress:
|
||||
req.on_prefill_progress(
|
||||
min(
|
||||
req.prefill_steps * max_batch_tokens, req.prompt_token_count
|
||||
),
|
||||
req.prompt_token_count,
|
||||
)
|
||||
continue
|
||||
|
||||
if not req.prefill_done and new_tokens:
|
||||
req.first_token_time = time.perf_counter()
|
||||
req.prefill_done = True
|
||||
|
||||
for i, token_id in enumerate(new_tokens):
|
||||
is_last = i == len(new_tokens) - 1
|
||||
is_final_stop = is_last and finish_reason and token_id in stop_ids
|
||||
if req.on_generation_token:
|
||||
req.on_generation_token()
|
||||
results.append(
|
||||
(
|
||||
task_id,
|
||||
_build_generation_response(
|
||||
tokenizer,
|
||||
token_id,
|
||||
finish_reason if is_last and finish_reason else None,
|
||||
req.prompt_token_count,
|
||||
new_token_count,
|
||||
req.start_time,
|
||||
req.first_token_time,
|
||||
suppress_text=bool(is_final_stop),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
if finish_reason:
|
||||
del self._active[task_id]
|
||||
|
||||
for req in self._active.values():
|
||||
if not req.prefill_done:
|
||||
req.prefill_steps += 1
|
||||
if req.on_prefill_progress:
|
||||
req.on_prefill_progress(
|
||||
min(
|
||||
req.prefill_steps * max_batch_tokens, req.prompt_token_count
|
||||
),
|
||||
req.prompt_token_count,
|
||||
)
|
||||
|
||||
return results
|
||||
|
||||
def cancel(self, task_ids: list[TaskId]) -> None:
|
||||
to_abort = [str(tid) for tid in task_ids if tid in self._active]
|
||||
if to_abort:
|
||||
self.engine.abort_request(to_abort)
|
||||
for tid in task_ids:
|
||||
self._active.pop(tid, None)
|
||||
|
||||
def close(self) -> None:
|
||||
if not hasattr(self, "engine"):
|
||||
return
|
||||
rids = [req.request_id for req in self._active.values()]
|
||||
if rids:
|
||||
self.engine.abort_request(rids)
|
||||
self._active.clear()
|
||||
del self.engine
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
if torch.distributed.is_initialized():
|
||||
torch.distributed.destroy_process_group()
|
||||
|
||||
|
||||
_weight_loading_callback: Callable[[int, int], None] | None = None
|
||||
_weight_loading_patched = False
|
||||
|
||||
|
||||
def get_weight_loading_callback() -> Callable[[int, int], None] | None:
|
||||
return _weight_loading_callback
|
||||
|
||||
|
||||
def set_weight_loading_callback(cb: Callable[[int, int], None] | None) -> None:
|
||||
global _weight_loading_callback
|
||||
_weight_loading_callback = cb
|
||||
|
||||
|
||||
_LAYER_INDEX_PATTERN = re.compile(r"\.layers\.(\d+)\.")
|
||||
_n_layers: int = 1
|
||||
|
||||
|
||||
def get_n_layers() -> int:
|
||||
return _n_layers
|
||||
|
||||
|
||||
def set_n_layers(n: int) -> None:
|
||||
global _n_layers
|
||||
_n_layers = n
|
||||
|
||||
|
||||
def _wrap_weights_iterator(
|
||||
original: Callable[..., Generator[tuple[str, "torch.Tensor"], None, None]],
|
||||
) -> Callable[..., Generator[tuple[str, "torch.Tensor"], None, None]]:
|
||||
def patched(
|
||||
hf_weights_files: list[str], *args: object, **kwargs: object
|
||||
) -> Generator[tuple[str, "torch.Tensor"], None, None]:
|
||||
callback = get_weight_loading_callback()
|
||||
if callback is not None and hf_weights_files:
|
||||
total_layers = get_n_layers()
|
||||
seen_layers: set[int] = set()
|
||||
last_reported = 0
|
||||
for name, tensor in original(hf_weights_files, *args, **kwargs):
|
||||
yield name, tensor
|
||||
match = _LAYER_INDEX_PATTERN.search(name)
|
||||
if match:
|
||||
seen_layers.add(int(match.group(1)))
|
||||
current = len(seen_layers)
|
||||
if current > last_reported:
|
||||
callback(current, total_layers)
|
||||
last_reported = current
|
||||
callback(total_layers, total_layers)
|
||||
else:
|
||||
yield from original(hf_weights_files, *args, **kwargs)
|
||||
|
||||
return patched
|
||||
|
||||
|
||||
def _monkey_patch_iterator(weight_utils: object, attr_name: str) -> None:
|
||||
original = getattr(weight_utils, attr_name, None)
|
||||
if original is None:
|
||||
return
|
||||
patched = _wrap_weights_iterator(original) # pyright: ignore[reportAny]
|
||||
setattr(weight_utils, attr_name, patched)
|
||||
for mod in list(sys.modules.values()):
|
||||
if mod is weight_utils:
|
||||
continue
|
||||
for name in list(vars(mod)):
|
||||
if vars(mod)[name] is original:
|
||||
setattr(mod, name, patched)
|
||||
|
||||
|
||||
def _patch_weight_loading_progress() -> None:
|
||||
global _weight_loading_patched
|
||||
if _weight_loading_patched:
|
||||
return
|
||||
_weight_loading_patched = True
|
||||
|
||||
from vllm.model_executor.model_loader import (
|
||||
weight_utils,
|
||||
)
|
||||
|
||||
_monkey_patch_iterator(weight_utils, "safetensors_weights_iterator")
|
||||
_monkey_patch_iterator(weight_utils, "fastsafetensors_weights_iterator")
|
||||
|
||||
import huggingface_hub
|
||||
|
||||
def _noop_metadata(*_a: object, **_kw: object) -> None:
|
||||
pass
|
||||
|
||||
original_metadata = huggingface_hub.get_safetensors_metadata
|
||||
huggingface_hub.get_safetensors_metadata = _noop_metadata
|
||||
for mod in list(sys.modules.values()):
|
||||
if mod is huggingface_hub:
|
||||
continue
|
||||
for attr in list(vars(mod)):
|
||||
if vars(mod)[attr] is original_metadata:
|
||||
setattr(mod, attr, _noop_metadata)
|
||||
|
||||
|
||||
def build_layer_groups(kv_cache_config: KVCacheConfig) -> list[int]:
|
||||
group_lookup: dict[str, int] = {}
|
||||
for group_idx, group_spec in enumerate(kv_cache_config.kv_cache_groups):
|
||||
for layer_name in group_spec.layer_names:
|
||||
group_lookup[layer_name] = group_idx
|
||||
|
||||
layer_to_group: list[int] = []
|
||||
for tensor_spec in kv_cache_config.kv_cache_tensors:
|
||||
for name in tensor_spec.shared_by:
|
||||
layer_to_group.append(group_lookup[name])
|
||||
return layer_to_group
|
||||
|
||||
|
||||
def load_vllm_engine(
|
||||
model_id: ModelId,
|
||||
trust_remote_code: bool,
|
||||
n_layers: int = 1,
|
||||
on_layer_loaded: Callable[[int, int], None] | None = None,
|
||||
kv_connector_cls: type[object] | None = None,
|
||||
) -> tuple[LLMEngine, ToolParser | None]:
|
||||
model_path = build_model_path(model_id)
|
||||
_patch_weight_loading_progress()
|
||||
|
||||
set_n_layers(n_layers)
|
||||
|
||||
# Use the dict-with-colon form the original branch used. The typed
|
||||
# `KVTransferConfig` object goes through a different vLLM code path
|
||||
# and (with `kv_load_failure_policy="recompute"`) trips the
|
||||
# APC/hybrid/chunked-prefill kv-cache nesting bug.
|
||||
kv_transfer_config: dict[str, str] | None = None
|
||||
if kv_connector_cls is not None:
|
||||
kv_transfer_config = {
|
||||
"kv_connector": (
|
||||
f"{kv_connector_cls.__module__}:{kv_connector_cls.__name__}"
|
||||
),
|
||||
"kv_role": "kv_both",
|
||||
}
|
||||
|
||||
has_mamba = False
|
||||
try:
|
||||
with open(model_path / "config.json") as f:
|
||||
model_config = json.load(f) # pyright: ignore[reportAny]
|
||||
text_config = model_config.get("text_config", model_config) # pyright: ignore[reportAny]
|
||||
has_mamba = "mamba_ssm_dtype" in text_config or "linear_attention" in (
|
||||
text_config.get("layer_types") or [] # pyright: ignore[reportAny]
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
if has_mamba:
|
||||
backends = [AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TRITON_ATTN]
|
||||
else:
|
||||
backends = [
|
||||
AttentionBackendEnum.FLASHINFER,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TRITON_ATTN,
|
||||
]
|
||||
|
||||
engine: LLMEngine | None = None
|
||||
for backend in backends:
|
||||
try:
|
||||
engine_args = EngineArgs(
|
||||
model=str(model_path.expanduser().resolve()),
|
||||
served_model_name=str(model_id),
|
||||
gpu_memory_utilization=0.05,
|
||||
trust_remote_code=trust_remote_code,
|
||||
load_format="fastsafetensors",
|
||||
enable_prefix_caching=True,
|
||||
attention_backend=backend,
|
||||
compilation_config=CompilationConfig(
|
||||
mode=CompilationMode.NONE,
|
||||
cudagraph_mode=CUDAGraphMode.NONE,
|
||||
),
|
||||
disable_log_stats=True,
|
||||
max_num_batched_tokens=4096,
|
||||
kv_transfer_config=kv_transfer_config, # pyright: ignore[reportArgumentType]
|
||||
disable_hybrid_kv_cache_manager=False,
|
||||
kv_cache_dtype="auto",
|
||||
)
|
||||
|
||||
set_weight_loading_callback(on_layer_loaded)
|
||||
engine = LLMEngine.from_engine_args(engine_args)
|
||||
logger.info(f"vLLM engine using attention backend: {backend}")
|
||||
break
|
||||
except (ValueError, RuntimeError, NotImplementedError) as e:
|
||||
logger.warning(f"Attention backend {backend} failed: {e}, trying next")
|
||||
engine = None
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
continue
|
||||
|
||||
if engine is None:
|
||||
raise RuntimeError(f"No attention backend worked for {model_id}")
|
||||
|
||||
tool_parser: ToolParser | None = None
|
||||
tokenizer = engine.get_tokenizer()
|
||||
chat_template = getattr(tokenizer, "chat_template", None)
|
||||
if isinstance(chat_template, str):
|
||||
tool_parser = infer_tool_parser(chat_template)
|
||||
if tool_parser:
|
||||
logger.info(
|
||||
f"inferred tool parser: {tool_parser.start_parsing} / {tool_parser.end_parsing}"
|
||||
)
|
||||
|
||||
logger.info(f"vLLM engine loaded for {model_id}")
|
||||
|
||||
return engine, tool_parser
|
||||
@@ -0,0 +1,480 @@
|
||||
# pyright: reportPrivateUsage=false, reportAttributeAccessIssue=false
|
||||
from collections.abc import Callable
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
import torch
|
||||
from vllm.v1.core.block_pool import BlockPool
|
||||
from vllm.v1.core.kv_cache_metrics import KVCacheMetricsCollector
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig
|
||||
from vllm.v1.request import Request
|
||||
from vllm.v1.worker.gpu_model_runner import GPUModelRunner
|
||||
|
||||
from exo.shared.logging import logger
|
||||
|
||||
INITIAL_FRACTION = 0.05
|
||||
GROWTH_HEADROOM_BYTES = 512 * 1024 * 1024
|
||||
MIN_GROWTH_BLOCKS = 16
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.v1.core.kv_cache_manager import KVCacheManager
|
||||
|
||||
|
||||
_patched = False
|
||||
_model_runner: GPUModelRunner | None = None
|
||||
|
||||
|
||||
def get_model_runner() -> GPUModelRunner | None:
|
||||
return _model_runner
|
||||
|
||||
|
||||
def set_model_runner(runner: GPUModelRunner | None) -> None:
|
||||
global _model_runner
|
||||
_model_runner = runner
|
||||
|
||||
|
||||
def patch_vllm() -> None:
|
||||
global _patched
|
||||
if _patched:
|
||||
return
|
||||
_patched = True
|
||||
|
||||
_patch_nogds()
|
||||
_patch_determine_available_memory()
|
||||
_patch_check_enough_kv_cache_memory()
|
||||
_patch_initialize_kv_cache_tensors()
|
||||
_patch_initialize_from_config()
|
||||
_patch_kv_cache_manager_init()
|
||||
_patch_allocate_slots()
|
||||
_patch_moe_sum()
|
||||
_patch_marlin_w2_thread_config()
|
||||
logger.info("vLLM growable KV cache patch applied")
|
||||
|
||||
|
||||
def _patch_nogds() -> None:
|
||||
from vllm.model_executor.model_loader import weight_utils
|
||||
|
||||
original = weight_utils._init_fastsafetensors_loader
|
||||
|
||||
def patched(
|
||||
pg: torch.distributed.ProcessGroup,
|
||||
device: torch.device,
|
||||
f_list: list[str],
|
||||
*,
|
||||
nogds: bool = False,
|
||||
) -> object:
|
||||
return original(pg, device, f_list, nogds=True)
|
||||
|
||||
weight_utils._init_fastsafetensors_loader = patched
|
||||
|
||||
|
||||
def _patch_determine_available_memory() -> None:
|
||||
from vllm.v1.worker.gpu_worker import Worker
|
||||
|
||||
# original = Worker.determine_available_memory
|
||||
|
||||
@torch.inference_mode()
|
||||
def patched(self: Worker) -> int:
|
||||
import pathlib
|
||||
import shutil
|
||||
|
||||
compile_cache = pathlib.Path.home() / ".cache" / "vllm" / "torch_compile_cache"
|
||||
if compile_cache.exists():
|
||||
shutil.rmtree(compile_cache, ignore_errors=True)
|
||||
|
||||
free_bytes, _ = torch.cuda.mem_get_info()
|
||||
# vLLM's get_kv_cache_configs computes per-group block counts via
|
||||
# `tensor.size // num_blocks_old` and asserts the result divides
|
||||
# evenly. With a small `available_kv_cache_memory_bytes` and
|
||||
# multi-MiB-per-slot Mamba/hybrid groups, num_blocks_old can come
|
||||
# back as 0 → ZeroDivisionError. Floor the initial budget so each
|
||||
# group lands at least one block at init; growth picks up from
|
||||
# there.
|
||||
min_initial = 1024 * 1024 * 1024 # 1 GiB
|
||||
if free_bytes < min_initial:
|
||||
raise RuntimeError(
|
||||
f"Insufficient GPU memory for KV cache initialization: "
|
||||
f"{free_bytes / (1024**3):.2f} GiB free, need at least "
|
||||
f"{min_initial / (1024**3):.2f} GiB. Stop other GPU "
|
||||
f"processes (check `nvidia-smi`)."
|
||||
)
|
||||
initial = max(int(free_bytes * INITIAL_FRACTION), min_initial)
|
||||
self._growable_max_kv_bytes = free_bytes
|
||||
self.available_kv_cache_memory_bytes = initial
|
||||
logger.info(
|
||||
f"Growable KV cache: initial {initial / (1024**3):.2f} GiB "
|
||||
f"(max {free_bytes / (1024**3):.2f} GiB)"
|
||||
)
|
||||
return initial
|
||||
|
||||
Worker.determine_available_memory = patched
|
||||
|
||||
|
||||
def _patch_check_enough_kv_cache_memory() -> None:
|
||||
from vllm.v1.core import kv_cache_utils
|
||||
|
||||
def noop(*_args: object, **_kwargs: object) -> None:
|
||||
pass
|
||||
|
||||
kv_cache_utils._check_enough_kv_cache_memory = noop
|
||||
|
||||
|
||||
def _patch_initialize_kv_cache_tensors() -> None:
|
||||
from vllm.v1.worker.gpu_model_runner import GPUModelRunner
|
||||
|
||||
original_alloc = GPUModelRunner._allocate_kv_cache_tensors
|
||||
|
||||
def patched_alloc(
|
||||
self: GPUModelRunner, kv_cache_config: KVCacheConfig
|
||||
) -> dict[str, torch.Tensor]:
|
||||
raw_tensors = original_alloc(self, kv_cache_config)
|
||||
self._growable_raw_tensors = {name: t for name, t in raw_tensors.items()}
|
||||
return raw_tensors
|
||||
|
||||
GPUModelRunner._allocate_kv_cache_tensors = patched_alloc
|
||||
|
||||
original_init_tensors = GPUModelRunner.initialize_kv_cache_tensors
|
||||
|
||||
def patched_init_tensors(
|
||||
self: GPUModelRunner,
|
||||
kv_cache_config: KVCacheConfig,
|
||||
kernel_block_sizes: list[int],
|
||||
) -> dict[str, torch.Tensor]:
|
||||
self._growable_kv_cache_config = kv_cache_config
|
||||
self._growable_kernel_block_sizes = kernel_block_sizes
|
||||
return original_init_tensors(self, kv_cache_config, kernel_block_sizes)
|
||||
|
||||
GPUModelRunner.initialize_kv_cache_tensors = patched_init_tensors
|
||||
|
||||
|
||||
def _patch_initialize_from_config() -> None:
|
||||
from vllm.v1.worker.gpu_model_runner import GPUModelRunner
|
||||
from vllm.v1.worker.gpu_worker import Worker
|
||||
|
||||
original_init_attn = GPUModelRunner.initialize_attn_backend
|
||||
|
||||
def clear_and_reinit_attn(
|
||||
self: GPUModelRunner,
|
||||
kv_cache_config: KVCacheConfig,
|
||||
) -> None:
|
||||
self.attn_groups.clear()
|
||||
original_init_attn(self, kv_cache_config)
|
||||
|
||||
GPUModelRunner.initialize_attn_backend = clear_and_reinit_attn
|
||||
|
||||
original = Worker.initialize_from_config
|
||||
|
||||
def patched(self: Worker, kv_cache_config: KVCacheConfig) -> None:
|
||||
original(self, kv_cache_config)
|
||||
set_model_runner(self.model_runner)
|
||||
|
||||
Worker.initialize_from_config = patched
|
||||
|
||||
|
||||
def _patch_kv_cache_manager_init() -> None:
|
||||
from vllm.v1.core.kv_cache_manager import KVCacheManager
|
||||
|
||||
original_init = KVCacheManager.__init__
|
||||
|
||||
def patched_init(
|
||||
self: KVCacheManager,
|
||||
kv_cache_config: KVCacheConfig,
|
||||
max_model_len: int,
|
||||
hash_block_size: int,
|
||||
enable_caching: bool = True,
|
||||
use_eagle: bool = False,
|
||||
log_stats: bool = False,
|
||||
enable_kv_cache_events: bool = False,
|
||||
dcp_world_size: int = 1,
|
||||
pcp_world_size: int = 1,
|
||||
metrics_collector: KVCacheMetricsCollector | None = None,
|
||||
) -> None:
|
||||
original_init(
|
||||
self,
|
||||
kv_cache_config,
|
||||
max_model_len,
|
||||
hash_block_size,
|
||||
enable_caching,
|
||||
use_eagle,
|
||||
log_stats,
|
||||
enable_kv_cache_events,
|
||||
dcp_world_size,
|
||||
pcp_world_size,
|
||||
metrics_collector,
|
||||
)
|
||||
self._growable_model_runner = get_model_runner()
|
||||
|
||||
KVCacheManager.__init__ = patched_init
|
||||
|
||||
|
||||
def _patch_allocate_slots() -> None:
|
||||
from vllm.v1.core.kv_cache_manager import KVCacheBlocks, KVCacheManager
|
||||
|
||||
original = KVCacheManager.allocate_slots
|
||||
|
||||
def patched(
|
||||
self: KVCacheManager,
|
||||
request: Request,
|
||||
num_new_tokens: int,
|
||||
num_new_computed_tokens: int = 0,
|
||||
new_computed_blocks: KVCacheBlocks | None = None,
|
||||
num_lookahead_tokens: int = 0,
|
||||
num_external_computed_tokens: int = 0,
|
||||
delay_cache_blocks: bool = False,
|
||||
num_encoder_tokens: int = 0,
|
||||
) -> KVCacheBlocks | None:
|
||||
result = original(
|
||||
self,
|
||||
request,
|
||||
num_new_tokens,
|
||||
num_new_computed_tokens,
|
||||
new_computed_blocks,
|
||||
num_lookahead_tokens,
|
||||
num_external_computed_tokens,
|
||||
delay_cache_blocks,
|
||||
num_encoder_tokens,
|
||||
)
|
||||
while result is None and _try_grow_cache(self):
|
||||
result = original(
|
||||
self,
|
||||
request,
|
||||
num_new_tokens,
|
||||
num_new_computed_tokens,
|
||||
new_computed_blocks,
|
||||
num_lookahead_tokens,
|
||||
num_external_computed_tokens,
|
||||
delay_cache_blocks,
|
||||
num_encoder_tokens,
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
KVCacheManager.allocate_slots = patched
|
||||
|
||||
if hasattr(KVCacheManager, "can_fit_full_sequence"):
|
||||
original_can_fit = cast(
|
||||
Callable[..., bool],
|
||||
KVCacheManager.can_fit_full_sequence,
|
||||
)
|
||||
|
||||
def patched_can_fit(
|
||||
self: KVCacheManager,
|
||||
request: Request,
|
||||
num_new_computed_tokens: int = 0,
|
||||
new_computed_blocks: KVCacheBlocks | None = None,
|
||||
num_external_computed_tokens: int = 0,
|
||||
num_encoder_tokens: int = 0,
|
||||
) -> bool:
|
||||
result: bool = original_can_fit(
|
||||
self,
|
||||
request,
|
||||
num_new_computed_tokens,
|
||||
new_computed_blocks,
|
||||
num_external_computed_tokens,
|
||||
num_encoder_tokens,
|
||||
)
|
||||
while not result and _try_grow_cache(self):
|
||||
result = original_can_fit(
|
||||
self,
|
||||
request,
|
||||
num_new_computed_tokens,
|
||||
new_computed_blocks,
|
||||
num_external_computed_tokens,
|
||||
num_encoder_tokens,
|
||||
)
|
||||
return result
|
||||
|
||||
KVCacheManager.can_fit_full_sequence = patched_can_fit
|
||||
|
||||
|
||||
def _try_grow_cache(kv_cache_manager: "KVCacheManager") -> bool:
|
||||
block_pool = kv_cache_manager.block_pool
|
||||
model_runner = cast(GPUModelRunner | None, kv_cache_manager._growable_model_runner)
|
||||
|
||||
if model_runner is None:
|
||||
return False
|
||||
|
||||
free_bytes, _ = torch.cuda.mem_get_info()
|
||||
if free_bytes < GROWTH_HEADROOM_BYTES:
|
||||
return False
|
||||
|
||||
kv_cache_config = cast(KVCacheConfig, model_runner._growable_kv_cache_config)
|
||||
old_num_blocks: int = kv_cache_config.num_blocks
|
||||
|
||||
total_tensor_bytes = sum(t.size for t in kv_cache_config.kv_cache_tensors)
|
||||
per_block_bytes = total_tensor_bytes // old_num_blocks
|
||||
|
||||
usable_bytes = int(free_bytes * 0.8)
|
||||
growth_blocks = min(usable_bytes // per_block_bytes, old_num_blocks)
|
||||
|
||||
if growth_blocks < MIN_GROWTH_BLOCKS:
|
||||
return False
|
||||
|
||||
new_num_blocks = old_num_blocks + growth_blocks
|
||||
|
||||
logger.info(
|
||||
f"Growing KV cache: {old_num_blocks} → {new_num_blocks} blocks "
|
||||
f"(+{growth_blocks * per_block_bytes / (1024**3):.2f} GiB)"
|
||||
)
|
||||
|
||||
try:
|
||||
kv_cache_config.num_blocks = new_num_blocks
|
||||
for tensor_spec in kv_cache_config.kv_cache_tensors:
|
||||
tensor_spec.size = int(tensor_spec.size * new_num_blocks / old_num_blocks)
|
||||
_grow_tensors(model_runner, kv_cache_config, old_num_blocks, new_num_blocks)
|
||||
_grow_block_pool(block_pool, old_num_blocks, new_num_blocks)
|
||||
logger.info(f"KV cache grown successfully to {new_num_blocks} blocks")
|
||||
return True
|
||||
except Exception:
|
||||
logger.opt(exception=True).error("Failed to grow KV cache")
|
||||
return False
|
||||
|
||||
|
||||
def _grow_tensors(
|
||||
model_runner: GPUModelRunner,
|
||||
kv_cache_config: KVCacheConfig,
|
||||
old_num_blocks: int,
|
||||
new_num_blocks: int,
|
||||
) -> None:
|
||||
raw_tensors: dict[str, torch.Tensor] = cast(
|
||||
dict[str, torch.Tensor], model_runner._growable_raw_tensors
|
||||
)
|
||||
ratio = new_num_blocks / old_num_blocks
|
||||
|
||||
already_grown: dict[int, torch.Tensor] = {}
|
||||
new_raw_tensors: dict[str, torch.Tensor] = {}
|
||||
|
||||
for layer_name, old_raw in raw_tensors.items():
|
||||
storage_id = old_raw.data_ptr()
|
||||
if storage_id in already_grown:
|
||||
new_raw_tensors[layer_name] = already_grown[storage_id]
|
||||
continue
|
||||
|
||||
old_size = old_raw.numel()
|
||||
new_size = int(old_size * ratio)
|
||||
new_raw = torch.zeros(new_size, dtype=torch.int8, device=old_raw.device)
|
||||
new_raw[:old_size] = old_raw
|
||||
already_grown[storage_id] = new_raw
|
||||
new_raw_tensors[layer_name] = new_raw
|
||||
|
||||
model_runner._growable_raw_tensors = new_raw_tensors
|
||||
|
||||
kernel_block_sizes: list[int] = cast(
|
||||
list[int], model_runner._growable_kernel_block_sizes
|
||||
)
|
||||
new_kv_caches: dict[str, torch.Tensor] = model_runner._reshape_kv_cache_tensors(
|
||||
kv_cache_config,
|
||||
new_raw_tensors,
|
||||
kernel_block_sizes,
|
||||
)
|
||||
|
||||
forward_context: dict[str, Any] = (
|
||||
model_runner.compilation_config.static_forward_context
|
||||
)
|
||||
runner_kv_caches: list[torch.Tensor] = model_runner.kv_caches
|
||||
|
||||
from collections import defaultdict
|
||||
|
||||
from vllm.model_executor.models.utils import extract_layer_index
|
||||
|
||||
num_attn_module = 1
|
||||
hf_config = getattr(getattr(model_runner, "model_config", None), "hf_config", None)
|
||||
if getattr(hf_config, "model_type", "") == "longcat_flash":
|
||||
num_attn_module = 2
|
||||
|
||||
index2name: dict[int, list[str]] = defaultdict(list)
|
||||
for ln in new_kv_caches:
|
||||
index2name[extract_layer_index(ln, num_attn_module)].append(ln)
|
||||
|
||||
new_ordered: list[torch.Tensor] = []
|
||||
for layer_index in sorted(index2name.keys()):
|
||||
for ln in index2name[layer_index]:
|
||||
new_ordered.append(new_kv_caches[ln])
|
||||
|
||||
for i, new_kv in enumerate(new_ordered):
|
||||
if i < len(runner_kv_caches):
|
||||
runner_kv_caches[i] = new_kv
|
||||
else:
|
||||
runner_kv_caches.append(new_kv)
|
||||
|
||||
new_kv_typed = cast(dict[str, torch.Tensor | list[torch.Tensor]], new_kv_caches)
|
||||
for layer_name, new_kv in new_kv_typed.items():
|
||||
# vLLM uses different shapes per layer kind (gpu_model_runner.py:5852):
|
||||
# - full / sliding-window attention: `attn.kv_cache: torch.Tensor`
|
||||
# (paged storage with K/V stacked along dim 0; consumers call
|
||||
# `.unbind(0)` so it MUST be a Tensor, not a list)
|
||||
# - Mamba / hybrid: `attn.kv_cache: list[Tensor]`
|
||||
# ([conv_state, ssm_state])
|
||||
# Preserve that distinction here. In-place .set_() keeps the existing
|
||||
# tensor identities valid for any captured refs (torch.compile graph,
|
||||
# layer module attrs); we only fall back to assignment on first
|
||||
# install or a shape mismatch.
|
||||
old_kv = cast(
|
||||
list[Any] | list[torch.Tensor] | torch.Tensor,
|
||||
forward_context[layer_name].kv_cache,
|
||||
)
|
||||
|
||||
if isinstance(new_kv, list):
|
||||
if (
|
||||
isinstance(old_kv, list)
|
||||
and len(old_kv) == len(new_kv)
|
||||
and all(isinstance(t, torch.Tensor) for t in old_kv)
|
||||
):
|
||||
for old_t, new_t in zip(old_kv, new_kv, strict=True):
|
||||
old_t.set_(
|
||||
new_t.storage(),
|
||||
new_t.storage_offset(),
|
||||
new_t.shape,
|
||||
new_t.stride(),
|
||||
)
|
||||
else:
|
||||
forward_context[layer_name].kv_cache = new_kv
|
||||
else:
|
||||
if isinstance(old_kv, torch.Tensor) and old_kv.numel() > 0:
|
||||
old_kv.set_(
|
||||
new_kv.storage(),
|
||||
new_kv.storage_offset(),
|
||||
new_kv.shape,
|
||||
new_kv.stride(),
|
||||
)
|
||||
else:
|
||||
forward_context[layer_name].kv_cache = new_kv
|
||||
|
||||
|
||||
def _grow_block_pool(
|
||||
block_pool: BlockPool, old_num_blocks: int, new_num_blocks: int
|
||||
) -> None:
|
||||
from vllm.v1.core.kv_cache_utils import KVCacheBlock
|
||||
|
||||
new_blocks: list[KVCacheBlock] = []
|
||||
for idx in range(old_num_blocks, new_num_blocks):
|
||||
block = KVCacheBlock(idx)
|
||||
block_pool.blocks.append(block)
|
||||
new_blocks.append(block)
|
||||
|
||||
block_pool.free_block_queue.append_n(new_blocks)
|
||||
block_pool.num_gpu_blocks = new_num_blocks
|
||||
|
||||
|
||||
def _patch_moe_sum() -> None:
|
||||
import vllm._custom_ops as ops
|
||||
|
||||
def moe_sum_f32(x: torch.Tensor, output: torch.Tensor) -> None:
|
||||
output[:] = x.to(torch.float32).sum(dim=1).to(output.dtype)
|
||||
|
||||
ops.moe_sum = moe_sum_f32
|
||||
|
||||
|
||||
def _patch_marlin_w2_thread_config() -> None:
|
||||
try:
|
||||
import vllm._custom_ops as ops
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
original_gemm = cast(Callable[..., object], ops.moe_wna16_marlin_gemm)
|
||||
|
||||
def patched_gemm(*args: object, **kwargs: object) -> object:
|
||||
kwargs["thread_k"] = 64
|
||||
kwargs["thread_n"] = 128
|
||||
return original_gemm(*args, **kwargs)
|
||||
|
||||
ops.moe_wna16_marlin_gemm = patched_gemm
|
||||
@@ -0,0 +1,787 @@
|
||||
# pyright: reportAny = false
|
||||
import contextlib
|
||||
import queue
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
from vllm.config import VllmConfig
|
||||
from vllm.config.kv_transfer import KVTransferConfig
|
||||
from vllm.distributed.kv_transfer.kv_connector.v1.base import (
|
||||
KVConnectorBase_V1,
|
||||
KVConnectorMetadata,
|
||||
KVConnectorRole,
|
||||
SupportsHMA,
|
||||
)
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig
|
||||
from vllm.v1.request import Request
|
||||
|
||||
from exo.worker.engines.vllm.disaggregated.adapter import (
|
||||
extract_kv_via_slot_mapping,
|
||||
to_bf16,
|
||||
)
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
|
||||
_LAYER_RE = re.compile(r"layers\.(\d+)\.")
|
||||
|
||||
|
||||
# Module-level shared state. Populated by the connector's hooks (running inside
|
||||
# vLLM's scheduler/worker, same process since V1 multiprocessing is off);
|
||||
# drained by the producer engine in `serve_prefill` after the request finishes.
|
||||
#
|
||||
# `_kv_queue` is the original streaming-connector path ported into this module.
|
||||
# We defer prefix reuse to vLLM APC and do not keep a separate TorchKVCache.
|
||||
# 4-tuple: (layer_idx, keys_host_pinned, values_host_pinned, copy_done_event)
|
||||
# The writer thread does `event.synchronize()` (CPU-side, doesn't block GPU)
|
||||
# before reading the pinned host bytes.
|
||||
# 5-tuple: (layer_idx, num_tokens, keys_host_pinned, values_host_pinned, copy_done_event)
|
||||
# `num_tokens` is the authoritative token count for this item. The writer uses
|
||||
# it for skip_tokens accounting *and* to slice the keys/values tensors before
|
||||
# writing to wire — never trusts `keys.shape[0]`, since shape can disagree with
|
||||
# token count when the source path packs/reshapes (e.g. NVFP4 layouts).
|
||||
_kv_queue: queue.Queue[
|
||||
tuple[int, int, torch.Tensor, torch.Tensor, torch.cuda.Event] | None
|
||||
] = queue.Queue()
|
||||
# 3-tuple: (layer_idx, arrays_host_or_gpu, copy_done_event_or_none)
|
||||
# - From save_kv_layer hybrid path: tensors are GPU, event=None (writer .cpu()s)
|
||||
# - From GDN capture (after both conv+ssm ready): tensors are pinned host,
|
||||
# event is a CUDA event the writer must synchronize on before reading
|
||||
_arrays_queue: queue.Queue[
|
||||
tuple[int, list[torch.Tensor], torch.cuda.Event | None] | None
|
||||
] = queue.Queue()
|
||||
# Per-layer tracking of which layers' GDN states have been shipped via the
|
||||
# async pipeline. Entries here are excluded from the post-writer fallback drain.
|
||||
_gdn_shipped: set[int] = set()
|
||||
_captured_layers: dict[int, dict[str, torch.Tensor]] = {}
|
||||
_captured_arrays: dict[int, list[torch.Tensor]] = {}
|
||||
# Hybrid-model SSM/conv state captured via causal_conv1d + delta-rule patches.
|
||||
_gdn_states: dict[int, dict[str, torch.Tensor]] = {}
|
||||
_gdn_layer_order: list[int] = []
|
||||
_gdn_call_idx: list[int] = [0]
|
||||
_ssm_call_idx: list[int] = [0]
|
||||
# Per-layer save_kv_layer call diagnostics: list of slot_mapping sizes seen.
|
||||
_save_kv_layer_diag: dict[int, list[int]] = {}
|
||||
# Side CUDA stream for K/V extract + async D2H, so vLLM's compute stream
|
||||
# isn't blocked on D2H/extract during forward.
|
||||
_save_stream: torch.cuda.Stream | None = None
|
||||
# Holds a reference to the set tracked by patched_schedule so
|
||||
# `reset_capture_state` can clear it between requests.
|
||||
_apc_extracted_set_ref: dict[str, set[str]] = {}
|
||||
# request_id → actual APC hit token count (captured at the moment vLLM's
|
||||
# kv_cache_manager.get_computed_blocks runs, before scheduler chunks the
|
||||
# remaining tokens). Used by patched_schedule to pre-extract exactly the
|
||||
# matched portion, not the matched+about-to-forward portion.
|
||||
_apc_hit_tokens: dict[str, int] = {}
|
||||
|
||||
|
||||
def _get_save_stream() -> torch.cuda.Stream:
|
||||
global _save_stream
|
||||
if _save_stream is None:
|
||||
_save_stream = torch.cuda.Stream()
|
||||
return _save_stream
|
||||
|
||||
|
||||
def get_kv_queue() -> queue.Queue[
|
||||
tuple[int, int, torch.Tensor, torch.Tensor, torch.cuda.Event] | None
|
||||
]:
|
||||
return _kv_queue
|
||||
|
||||
|
||||
def get_arrays_queue() -> queue.Queue[
|
||||
tuple[int, list[torch.Tensor], torch.cuda.Event | None] | None
|
||||
]:
|
||||
return _arrays_queue
|
||||
|
||||
|
||||
def get_gdn_states() -> dict[int, dict[str, torch.Tensor]]:
|
||||
return _gdn_states
|
||||
|
||||
|
||||
def get_gdn_shipped() -> set[int]:
|
||||
return _gdn_shipped
|
||||
|
||||
|
||||
def _try_ship_gdn(layer_idx: int) -> None:
|
||||
"""If both conv and ssm have been captured for `layer_idx`, kick off an
|
||||
async pinned D2H on the side stream and enqueue an arrays-state item so
|
||||
the writer thread can ship the bytes during forward instead of after.
|
||||
|
||||
Called from BOTH the conv and ssm capture patches. Conv always fires
|
||||
before ssm in a Mamba layer's forward, so this is a no-op after conv
|
||||
(state lacks ssm) and ships once after ssm. For chunked prefill the
|
||||
pair fires once per chunk: we ship every time, and the consumer's
|
||||
`arrays[layer_idx] = ...` last-write-wins keeps the final-chunk state
|
||||
(Mamba state is cumulative, only the final state matters).
|
||||
"""
|
||||
state = _gdn_states.get(layer_idx)
|
||||
if state is None or "conv" not in state or "ssm" not in state:
|
||||
return
|
||||
conv_gpu = state["conv"]
|
||||
ssm_gpu = state["ssm"]
|
||||
side_stream = _get_save_stream()
|
||||
side_stream.wait_stream(torch.cuda.current_stream()) # pyright: ignore[reportUnknownMemberType]
|
||||
with torch.cuda.stream(side_stream):
|
||||
conv_host = torch.empty(conv_gpu.shape, dtype=conv_gpu.dtype, pin_memory=True)
|
||||
ssm_host = torch.empty(ssm_gpu.shape, dtype=ssm_gpu.dtype, pin_memory=True)
|
||||
conv_host.copy_(conv_gpu, non_blocking=True)
|
||||
ssm_host.copy_(ssm_gpu, non_blocking=True)
|
||||
event = torch.cuda.Event()
|
||||
event.record(side_stream)
|
||||
_arrays_queue.put((layer_idx, [conv_host, ssm_host], event))
|
||||
_gdn_shipped.add(layer_idx)
|
||||
|
||||
|
||||
def get_save_kv_layer_diag() -> dict[int, list[int]]:
|
||||
return _save_kv_layer_diag
|
||||
|
||||
|
||||
def get_captured_layers() -> dict[int, dict[str, torch.Tensor]]:
|
||||
return _captured_layers
|
||||
|
||||
|
||||
def get_captured_arrays() -> dict[int, list[torch.Tensor]]:
|
||||
return _captured_arrays
|
||||
|
||||
|
||||
def reset_capture_state() -> None:
|
||||
while not _kv_queue.empty():
|
||||
try:
|
||||
_kv_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
while not _arrays_queue.empty():
|
||||
try:
|
||||
_arrays_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
_captured_layers.clear()
|
||||
_captured_arrays.clear()
|
||||
_gdn_states.clear()
|
||||
_gdn_shipped.clear()
|
||||
_gdn_call_idx[0] = 0
|
||||
_ssm_call_idx[0] = 0
|
||||
_save_kv_layer_diag.clear()
|
||||
_apc_hit_tokens.clear()
|
||||
apc_set = _apc_extracted_set_ref.get("set")
|
||||
if apc_set is not None:
|
||||
apc_set.clear()
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamingConnectorMetadata(KVConnectorMetadata):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class BatchConnectorMetadata(KVConnectorMetadata):
|
||||
pass
|
||||
|
||||
|
||||
class StreamingConnector(KVConnectorBase_V1, SupportsHMA):
|
||||
"""Original streaming producer connector, kept under the new server abstraction."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vllm_config: VllmConfig,
|
||||
role: KVConnectorRole,
|
||||
kv_cache_config: KVCacheConfig | None = None,
|
||||
) -> None:
|
||||
super().__init__(vllm_config, role, kv_cache_config)
|
||||
self._save_count = 0
|
||||
|
||||
# =========================================================================
|
||||
# Worker-side hooks (the only ones we actually use)
|
||||
# =========================================================================
|
||||
|
||||
def start_load_kv(self, forward_context: Any, **kwargs: Any) -> None:
|
||||
return
|
||||
|
||||
def wait_for_layer_load(self, layer_name: str) -> None:
|
||||
return
|
||||
|
||||
def save_kv_layer(
|
||||
self,
|
||||
layer_name: str,
|
||||
kv_layer: Any,
|
||||
attn_metadata: Any,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
slot_mapping = getattr(attn_metadata, "slot_mapping", None)
|
||||
m = _LAYER_RE.search(layer_name)
|
||||
layer_idx_for_diag = int(m.group(1)) if m else -1
|
||||
slot_size = int(slot_mapping.shape[0]) if slot_mapping is not None else -1
|
||||
is_list_kv = isinstance(kv_layer, (list, tuple))
|
||||
# Tag list/tuple as negative so the diag log distinguishes hybrid from
|
||||
# non-hybrid even when slot_size is the same.
|
||||
_save_kv_layer_diag.setdefault(layer_idx_for_diag, []).append(
|
||||
-slot_size if is_list_kv else slot_size
|
||||
)
|
||||
|
||||
# Skip decode-step saves (small slot mapping); we only want prefill.
|
||||
if slot_mapping is not None and slot_mapping.shape[0] <= 100:
|
||||
return
|
||||
if m is None:
|
||||
return
|
||||
layer_idx = int(m.group(1))
|
||||
|
||||
# Hybrid (Mamba+attention) layers: kv_layer is a list/tuple of state
|
||||
# tensors (conv + ssm). Send them straight to the arrays queue —
|
||||
# they don't live in the paged KV cache. Stay on GPU; the writer
|
||||
# thread does the D2H copy via `tensor_to_wire_bytes`.
|
||||
if isinstance(kv_layer, (list, tuple)):
|
||||
arrays = [
|
||||
to_bf16(t)
|
||||
for t in cast(list[torch.Tensor] | tuple[torch.Tensor, ...], kv_layer)
|
||||
]
|
||||
_arrays_queue.put((layer_idx, arrays, None))
|
||||
return
|
||||
|
||||
# Standard attention layers (full or sliding-window): extract K/V
|
||||
# via slot_mapping, which points to where vLLM is *writing* this
|
||||
# forward step's tokens. Capturing here, before sliding-window
|
||||
# eviction in the block pool, is the only way to ship every prompt
|
||||
# token's K/V regardless of attention type.
|
||||
#
|
||||
# All of this work — gather + bf16 cast + D2H — runs on a side
|
||||
# CUDA stream into pinned host memory. vLLM's compute stream is
|
||||
# never blocked: it only has to record-event for our side stream
|
||||
# to wait on, then it continues into the next layer's forward.
|
||||
# The writer thread later waits on the CUDA event (CPU-side wait,
|
||||
# doesn't block GPU) and ships the already-on-host bytes.
|
||||
if slot_mapping is not None:
|
||||
try:
|
||||
save_stream = _get_save_stream()
|
||||
save_stream.wait_stream(torch.cuda.current_stream()) # pyright: ignore[reportUnknownMemberType] # TODO: stub
|
||||
with torch.cuda.stream(save_stream):
|
||||
keys_gpu, values_gpu = extract_kv_via_slot_mapping(
|
||||
kv_layer, slot_mapping
|
||||
)
|
||||
keys_host = torch.empty(
|
||||
keys_gpu.shape, dtype=keys_gpu.dtype, pin_memory=True
|
||||
)
|
||||
values_host = torch.empty(
|
||||
values_gpu.shape, dtype=values_gpu.dtype, pin_memory=True
|
||||
)
|
||||
keys_host.copy_(keys_gpu, non_blocking=True)
|
||||
values_host.copy_(values_gpu, non_blocking=True)
|
||||
num_tokens = int(keys_gpu.shape[0])
|
||||
event = torch.cuda.Event()
|
||||
event.record(save_stream)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
f"save_kv_layer extract failed layer={layer_idx} "
|
||||
f"kv_layer.shape={getattr(kv_layer, 'shape', None)} "
|
||||
f"slot_mapping.shape={slot_mapping.shape}: {exc!r}"
|
||||
)
|
||||
return
|
||||
_kv_queue.put((layer_idx, num_tokens, keys_host, values_host, event))
|
||||
|
||||
def wait_for_save(self) -> None:
|
||||
return
|
||||
|
||||
# =========================================================================
|
||||
# Scheduler-side hooks (no-ops; we don't load and don't track allocs)
|
||||
# =========================================================================
|
||||
|
||||
def get_num_new_matched_tokens(
|
||||
self, request: Any, num_computed_tokens: int
|
||||
) -> tuple[int, bool]:
|
||||
return 0, False
|
||||
|
||||
def update_state_after_alloc(
|
||||
self, request: Any, blocks: Any, num_external_tokens: int
|
||||
) -> None:
|
||||
return
|
||||
|
||||
def build_connector_meta(self, scheduler_output: Any) -> StreamingConnectorMetadata:
|
||||
return StreamingConnectorMetadata()
|
||||
|
||||
def request_finished(
|
||||
self, request: Any, block_ids: list[int]
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
return False, None
|
||||
|
||||
def request_finished_all_groups(
|
||||
self, request: Any, block_ids: tuple[list[int], ...]
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
return False, None
|
||||
|
||||
|
||||
class BatchConnector(KVConnectorBase_V1, SupportsHMA):
|
||||
"""Original batch producer connector, ported for parity with the old branch."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
vllm_config: VllmConfig,
|
||||
role: KVConnectorRole,
|
||||
kv_cache_config: KVCacheConfig | None = None,
|
||||
) -> None:
|
||||
super().__init__(vllm_config, role, kv_cache_config)
|
||||
|
||||
def start_load_kv(self, forward_context: Any, **kwargs: Any) -> None:
|
||||
return
|
||||
|
||||
def wait_for_layer_load(self, layer_name: str) -> None:
|
||||
return
|
||||
|
||||
def save_kv_layer(
|
||||
self,
|
||||
layer_name: str,
|
||||
kv_layer: Any,
|
||||
attn_metadata: Any,
|
||||
**kwargs: Any,
|
||||
) -> None:
|
||||
slot_mapping = getattr(attn_metadata, "slot_mapping", None)
|
||||
if slot_mapping is not None and slot_mapping.shape[0] <= 100:
|
||||
return
|
||||
|
||||
m = _LAYER_RE.search(layer_name)
|
||||
if m is None:
|
||||
return
|
||||
layer_idx = int(m.group(1))
|
||||
|
||||
if isinstance(kv_layer, (list, tuple)):
|
||||
_captured_arrays[layer_idx] = [
|
||||
to_bf16(t).cpu()
|
||||
for t in cast(list[torch.Tensor] | tuple[torch.Tensor, ...], kv_layer)
|
||||
]
|
||||
return
|
||||
|
||||
if slot_mapping is None:
|
||||
return
|
||||
keys, values = extract_kv_via_slot_mapping(kv_layer, slot_mapping)
|
||||
prev = _captured_layers.get(layer_idx)
|
||||
if prev is None:
|
||||
_captured_layers[layer_idx] = {"keys": keys, "values": values}
|
||||
else:
|
||||
_captured_layers[layer_idx] = {
|
||||
"keys": torch.cat([prev["keys"], keys], dim=0),
|
||||
"values": torch.cat([prev["values"], values], dim=0),
|
||||
}
|
||||
|
||||
def wait_for_save(self) -> None:
|
||||
return
|
||||
|
||||
def request_finished(
|
||||
self, request: Any, block_ids: list[int]
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
return False, None
|
||||
|
||||
def request_finished_all_groups(
|
||||
self, request: Any, block_ids: tuple[list[int], ...]
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
return False, None
|
||||
|
||||
def get_num_new_matched_tokens(
|
||||
self, request: Any, num_computed_tokens: int
|
||||
) -> tuple[int, bool]:
|
||||
return 0, False
|
||||
|
||||
def update_state_after_alloc(
|
||||
self, request: Any, blocks: Any, num_external_tokens: int
|
||||
) -> None:
|
||||
return
|
||||
|
||||
def build_connector_meta(self, scheduler_output: Any) -> BatchConnectorMetadata:
|
||||
return BatchConnectorMetadata()
|
||||
|
||||
|
||||
ExoKVProducerConnector = StreamingConnector
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Bypass patches — necessary to make our connector usable inside vLLM 1.x.
|
||||
# Ported from the original branch's prefill_server.py:_patch_vllm_for_connector.
|
||||
# =============================================================================
|
||||
|
||||
_connector_patched = False
|
||||
|
||||
|
||||
def _patch_vllm_for_connector(connector_class: type[Any]) -> None:
|
||||
"""Three patches that make a custom save-only connector cooperate with vLLM.
|
||||
|
||||
1. Suppress `unify_hybrid_kv_cache_specs` ValueError on hybrid (Mamba +
|
||||
attention) models — the unifier complains about mixed cache specs we
|
||||
don't need to actually unify for save-only operation.
|
||||
2. Override `Scheduler._connector_finished` to short-circuit the
|
||||
async-save state machine. We're synchronous on the producer side.
|
||||
3. Make `KVConnectorFactory._get_connector_class_with_compat` recognize
|
||||
our class name and return our class directly, bypassing vLLM's
|
||||
registry of built-in connectors.
|
||||
"""
|
||||
global _connector_patched
|
||||
if _connector_patched:
|
||||
return
|
||||
_connector_patched = True
|
||||
|
||||
from vllm.v1.core import kv_cache_utils
|
||||
|
||||
original_unify = kv_cache_utils.unify_hybrid_kv_cache_specs
|
||||
|
||||
def patched_unify(kv_cache_spec: Any) -> None:
|
||||
with contextlib.suppress(ValueError):
|
||||
original_unify(kv_cache_spec)
|
||||
|
||||
kv_cache_utils.unify_hybrid_kv_cache_specs = patched_unify
|
||||
|
||||
from vllm.v1.core.sched import scheduler as sched_mod
|
||||
|
||||
def patched_connector_finished(
|
||||
self: sched_mod.Scheduler, request: Request
|
||||
) -> tuple[bool, dict[str, Any] | None]:
|
||||
return False, None
|
||||
|
||||
sched_mod.Scheduler._connector_finished = patched_connector_finished # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
from vllm.distributed.kv_transfer.kv_connector import factory
|
||||
|
||||
original_get = factory.KVConnectorFactory._get_connector_class_with_compat # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
def patched_get(kv_transfer_config: KVTransferConfig) -> tuple[Any, Any]:
|
||||
kv_conn = kv_transfer_config.kv_connector or ""
|
||||
kv_conn_lower = kv_conn.lower()
|
||||
if (
|
||||
kv_conn
|
||||
in {
|
||||
connector_class.__name__,
|
||||
f"{connector_class.__module__}:{connector_class.__name__}",
|
||||
"ExoKVProducerConnector",
|
||||
f"{__name__}:ExoKVProducerConnector",
|
||||
"StreamingConnector",
|
||||
f"{__name__}:StreamingConnector",
|
||||
}
|
||||
or "streaming_connector" in kv_conn_lower
|
||||
):
|
||||
return connector_class, None
|
||||
if "batch_connector" in kv_conn_lower:
|
||||
return BatchConnector, None
|
||||
return original_get(kv_transfer_config)
|
||||
|
||||
factory.KVConnectorFactory._get_connector_class_with_compat = patched_get # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
# Patch KVCacheManager.get_computed_blocks so we capture the actual APC-hit
|
||||
# token count for each request at the moment vLLM looks it up — *before*
|
||||
# the scheduler bumps `req.num_computed_tokens` with the chunked-prefill
|
||||
# first-chunk size. Reading `req.num_computed_tokens` post-schedule yields
|
||||
# `apc_hit + first_chunk` and would cause us to extract bytes from blocks
|
||||
# that haven't been written yet for the first-chunk tail.
|
||||
try:
|
||||
from vllm.v1.core.kv_cache_manager import ( # pyright: ignore[reportMissingImports]
|
||||
KVCacheManager,
|
||||
)
|
||||
except ImportError:
|
||||
KVCacheManager = None # noqa: N806
|
||||
|
||||
if KVCacheManager is not None:
|
||||
original_get_computed_blocks = KVCacheManager.get_computed_blocks
|
||||
|
||||
def patched_get_computed_blocks(self: Any, request: Any) -> Any:
|
||||
result = original_get_computed_blocks(self, request)
|
||||
try:
|
||||
req_id = getattr(request, "request_id", None)
|
||||
if req_id is not None:
|
||||
if isinstance(result, tuple) and len(result) >= 2: # pyright: ignore[reportUnknownArgumentType]
|
||||
num = int(result[1]) # pyright: ignore[reportUnknownArgumentType]
|
||||
else:
|
||||
num = 0
|
||||
total = int(getattr(request, "num_tokens", 0) or 0)
|
||||
logger.info(
|
||||
f"APC get_computed_blocks: req={req_id} hit={num} total={total}"
|
||||
)
|
||||
if num > 0:
|
||||
_apc_hit_tokens[req_id] = num
|
||||
except Exception:
|
||||
logger.opt(exception=True).warning(
|
||||
"patched_get_computed_blocks: capture failed"
|
||||
)
|
||||
return result
|
||||
|
||||
KVCacheManager.get_computed_blocks = patched_get_computed_blocks # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
# Patch Scheduler.schedule so APC-cached prefix blocks are extracted out
|
||||
# of the paged pool and pushed to _kv_queue at scheduling time — BEFORE
|
||||
# forward runs. Forward will only execute the suffix (vLLM's own APC
|
||||
# behavior). save_kv_layer fires for the suffix as usual. The writer
|
||||
# thread sees: prefix items from this hook + suffix items from save_kv_layer
|
||||
# and ships them in arrival order (prefix before suffix per layer).
|
||||
original_schedule = sched_mod.Scheduler.schedule
|
||||
_scheduled_apc_extracted: set[str] = set()
|
||||
|
||||
def patched_schedule(self: sched_mod.Scheduler) -> Any:
|
||||
scheduler_output = original_schedule(self)
|
||||
try:
|
||||
new_reqs = getattr(scheduler_output, "scheduled_new_reqs", None) or []
|
||||
if not new_reqs:
|
||||
return scheduler_output
|
||||
from exo.worker.engines.vllm.disaggregated.adapter import (
|
||||
build_layer_to_group,
|
||||
gather_layer_kv_from_blocks,
|
||||
)
|
||||
from exo.worker.engines.vllm.growable_cache import get_model_runner
|
||||
|
||||
mr = get_model_runner()
|
||||
if mr is None:
|
||||
return scheduler_output
|
||||
cfg = getattr(mr, "_growable_kv_cache_config", None)
|
||||
if cfg is None:
|
||||
return scheduler_output
|
||||
layer_to_group = build_layer_to_group(cfg)
|
||||
n_layers = len(mr.kv_caches)
|
||||
|
||||
for new_req in new_reqs:
|
||||
req_id = getattr(new_req, "req_id", None)
|
||||
if req_id is None or req_id in _scheduled_apc_extracted:
|
||||
continue
|
||||
pre_layers_shipped = 0
|
||||
pre_bytes_shipped = 0
|
||||
req = self.requests.get(req_id)
|
||||
if req is None:
|
||||
continue
|
||||
# Use the count captured by patched_get_computed_blocks (the
|
||||
# actual APC hit), NOT req.num_computed_tokens — that field has
|
||||
# already been bumped by the scheduler with the first chunk's
|
||||
# about-to-forward token count and would over-extract.
|
||||
num_apc = _apc_hit_tokens.get(req_id, 0)
|
||||
req_total = int(getattr(req, "num_tokens", 0) or 0)
|
||||
req_computed = int(getattr(req, "num_computed_tokens", 0) or 0)
|
||||
logger.info(
|
||||
f"APC patched_schedule: req={req_id} apc_hit={num_apc} "
|
||||
f"req.num_computed_tokens={req_computed} req.num_tokens={req_total}"
|
||||
)
|
||||
if num_apc <= 0:
|
||||
_scheduled_apc_extracted.add(req_id)
|
||||
continue
|
||||
# Pull the request's full per-group block list from
|
||||
# scheduler_output.scheduled_new_reqs[i].block_ids — that field
|
||||
# includes APC-cached prefix blocks. The KVCacheManager's
|
||||
# `req_to_blocks` only tracks newly-allocated blocks for this
|
||||
# step's suffix, so reading from there misses the prefix and
|
||||
# makes gather return ~bock_count_suffix tokens of garbage.
|
||||
req_block_ids_per_group: tuple[list[int], ...] | None = getattr(
|
||||
new_req, "block_ids", None
|
||||
)
|
||||
if not req_block_ids_per_group:
|
||||
logger.warning(
|
||||
f"APC pre-extract: new_req.block_ids missing for {req_id}"
|
||||
)
|
||||
_scheduled_apc_extracted.add(req_id)
|
||||
continue
|
||||
save_stream = _get_save_stream()
|
||||
save_stream.wait_stream(torch.cuda.current_stream()) # pyright: ignore[reportUnknownMemberType]
|
||||
first_log_done = False
|
||||
# Run the entire gather + cast + pinned alloc + D2H on the
|
||||
# side stream — scheduler thread only issues kernel launches
|
||||
# and records an event per layer. Compute stream is untouched.
|
||||
with torch.cuda.stream(save_stream):
|
||||
for layer_idx in range(n_layers):
|
||||
kv_layer = mr.kv_caches[layer_idx]
|
||||
if isinstance(kv_layer, (list, tuple)):
|
||||
continue
|
||||
gi = (
|
||||
layer_to_group[layer_idx]
|
||||
if layer_idx < len(layer_to_group)
|
||||
else 0
|
||||
)
|
||||
if gi >= len(req_block_ids_per_group):
|
||||
continue
|
||||
block_ids = list(req_block_ids_per_group[gi])
|
||||
if not block_ids:
|
||||
continue
|
||||
keys_gpu, values_gpu = gather_layer_kv_from_blocks(
|
||||
kv_layer, block_ids, num_apc
|
||||
)
|
||||
if not first_log_done:
|
||||
first_log_done = True
|
||||
logger.info(
|
||||
f"APC pre-extract layer={layer_idx}: "
|
||||
f"kv_layer.shape={tuple(kv_layer.shape)} "
|
||||
f"kv_layer.dtype={kv_layer.dtype} "
|
||||
f"len(block_ids)={len(block_ids)} num_apc={num_apc} "
|
||||
f"keys_gpu.shape={tuple(keys_gpu.shape)} "
|
||||
f"keys_gpu.dtype={keys_gpu.dtype}"
|
||||
)
|
||||
if keys_gpu.numel() == 0:
|
||||
continue
|
||||
keys_host = torch.empty(
|
||||
keys_gpu.shape,
|
||||
dtype=keys_gpu.dtype,
|
||||
pin_memory=True,
|
||||
)
|
||||
values_host = torch.empty(
|
||||
values_gpu.shape,
|
||||
dtype=values_gpu.dtype,
|
||||
pin_memory=True,
|
||||
)
|
||||
keys_host.copy_(keys_gpu, non_blocking=True)
|
||||
values_host.copy_(values_gpu, non_blocking=True)
|
||||
event = torch.cuda.Event()
|
||||
event.record(save_stream)
|
||||
_kv_queue.put(
|
||||
(layer_idx, num_apc, keys_host, values_host, event)
|
||||
)
|
||||
pre_layers_shipped += 1
|
||||
pre_bytes_shipped += (
|
||||
keys_host.numel() * keys_host.element_size()
|
||||
+ values_host.numel() * values_host.element_size()
|
||||
)
|
||||
logger.info(
|
||||
f"APC pre-extract done: req={req_id} layers={pre_layers_shipped} "
|
||||
f"tokens={num_apc} bytes={pre_bytes_shipped}"
|
||||
)
|
||||
_scheduled_apc_extracted.add(req_id)
|
||||
except Exception:
|
||||
logger.opt(exception=True).warning(
|
||||
"patched_schedule: APC pre-extract failed; continuing"
|
||||
)
|
||||
return scheduler_output
|
||||
|
||||
sched_mod.Scheduler.schedule = patched_schedule
|
||||
# Reset the per-request-extracted set when reset_capture_state runs.
|
||||
_apc_extracted_set_ref["set"] = _scheduled_apc_extracted
|
||||
logger.info("Installed vLLM connector bypass patches")
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Hybrid-model GDN state capture (Qwen3.5/3.6 etc.).
|
||||
# Patches the conv1d kernel + delta-rule fns to grab conv/ssm states per layer.
|
||||
# =============================================================================
|
||||
|
||||
_gdn_patched = False
|
||||
|
||||
|
||||
def _patch_gdn_capture() -> None:
|
||||
global _gdn_patched
|
||||
if _gdn_patched:
|
||||
return
|
||||
_gdn_patched = True
|
||||
|
||||
try:
|
||||
import vllm.model_executor.layers.mamba.ops.causal_conv1d as cc_mod
|
||||
from vllm.model_executor.layers.mamba.ops.causal_conv1d import (
|
||||
causal_conv1d_fn as orig_fn,
|
||||
)
|
||||
except ImportError:
|
||||
return
|
||||
|
||||
def patched_fn(
|
||||
*args: Any, conv_states: Any = None, cache_indices: Any = None, **kwargs: Any
|
||||
) -> Any:
|
||||
result = orig_fn(
|
||||
*args, conv_states=conv_states, cache_indices=cache_indices, **kwargs
|
||||
)
|
||||
if conv_states is not None and cache_indices is not None:
|
||||
x = args[0] if args else None
|
||||
if x is not None and x.shape[0] <= 100:
|
||||
return result
|
||||
ci: int = cache_indices[0].item() if cache_indices.numel() > 0 else 0
|
||||
idx = _gdn_call_idx[0]
|
||||
if _gdn_layer_order and idx < len(_gdn_layer_order) * 100:
|
||||
layer_idx = _gdn_layer_order[idx % len(_gdn_layer_order)]
|
||||
# `.contiguous()` decouples the slice from the underlying
|
||||
# buffer; D2H is deferred to the writer thread.
|
||||
conv_at_ci = conv_states[ci : ci + 1].transpose(-1, -2).contiguous()
|
||||
_gdn_states.setdefault(layer_idx, {})["conv"] = conv_at_ci
|
||||
_gdn_states[layer_idx]["ci"] = ci
|
||||
# Don't ship from here: conv fires before ssm in a Mamba
|
||||
# forward, so state["ssm"] is either missing (chunk 1) or
|
||||
# stale from the previous chunk (chunk N>=2). Shipping here
|
||||
# would emit a mismatched (conv_N, ssm_{N-1}) pair that the
|
||||
# ssm patch's later ship would overwrite. Just wait for ssm.
|
||||
_gdn_call_idx[0] += 1
|
||||
return result
|
||||
|
||||
cc_mod.causal_conv1d_fn = patched_fn
|
||||
import sys
|
||||
|
||||
for mod in list(sys.modules.values()):
|
||||
if mod is cc_mod:
|
||||
continue
|
||||
# transformers' image_processing_* shims have a lazy __getattr__
|
||||
# that emits a noisy deprecation warning on every attribute probe.
|
||||
# They never use causal_conv1d_fn, so skip them.
|
||||
mod_name = getattr(mod, "__name__", "") or ""
|
||||
if mod_name.startswith("transformers."):
|
||||
continue
|
||||
if (
|
||||
mod.__dict__.get("causal_conv1d_fn") is orig_fn
|
||||
if hasattr(mod, "__dict__")
|
||||
else False
|
||||
):
|
||||
mod.causal_conv1d_fn = patched_fn
|
||||
logger.info("Patched causal_conv1d_fn for GDN conv-state capture")
|
||||
|
||||
# The GDN delta-rule functions live in `mamba/gdn_linear_attn` (defined or
|
||||
# re-imported there) and may also be re-exported by model modules. Patch
|
||||
# all candidate modules + propagate to anywhere they're imported.
|
||||
candidate_modules = [
|
||||
"vllm.model_executor.layers.mamba.gdn_linear_attn",
|
||||
"vllm.model_executor.models.qwen3_next",
|
||||
"vllm.model_executor.models.qwen3_5",
|
||||
]
|
||||
fn_names = ("fi_chunk_gated_delta_rule", "fla_chunk_gated_delta_rule")
|
||||
patched_targets: list[str] = []
|
||||
|
||||
for mod_path in candidate_modules:
|
||||
try:
|
||||
mod = __import__(mod_path, fromlist=["*"])
|
||||
except ImportError:
|
||||
continue
|
||||
for fn_name in fn_names:
|
||||
orig = getattr(mod, fn_name, None)
|
||||
if orig is None:
|
||||
continue
|
||||
|
||||
def make_patched(orig_fn_inner: Any) -> Any:
|
||||
def patched_chunk(*args: Any, **kwargs: Any) -> Any:
|
||||
result = orig_fn_inner(*args, **kwargs)
|
||||
output_final_state = kwargs.get("output_final_state", False)
|
||||
if (
|
||||
output_final_state
|
||||
and isinstance(result, tuple)
|
||||
and len(result) == 2 # pyright: ignore[reportUnknownArgumentType]
|
||||
):
|
||||
_, ssm_state = result # pyright: ignore[reportUnknownVariableType]
|
||||
idx = _ssm_call_idx[0]
|
||||
if _gdn_layer_order and idx < len(_gdn_layer_order) * 100:
|
||||
layer_idx = _gdn_layer_order[idx % len(_gdn_layer_order)]
|
||||
_gdn_states.setdefault(layer_idx, {})["ssm"] = ssm_state
|
||||
_try_ship_gdn(layer_idx)
|
||||
_ssm_call_idx[0] += 1
|
||||
return result # pyright: ignore[reportUnknownVariableType]
|
||||
|
||||
return patched_chunk
|
||||
|
||||
patched_fn = make_patched(orig)
|
||||
setattr(mod, fn_name, patched_fn)
|
||||
patched_targets.append(f"{mod_path}.{fn_name}")
|
||||
# Propagate to any module that imported the original function.
|
||||
import sys as _sys
|
||||
|
||||
for other in list(_sys.modules.values()):
|
||||
if other is mod:
|
||||
continue
|
||||
other_name = getattr(other, "__name__", "") or ""
|
||||
# Skip transformers — see causal_conv1d_fn loop above.
|
||||
if other_name.startswith("transformers."):
|
||||
continue
|
||||
if other.__dict__.get(fn_name) is orig:
|
||||
setattr(other, fn_name, patched_fn)
|
||||
patched_targets.append(f"{other.__name__}.{fn_name} (propagated)")
|
||||
if patched_targets:
|
||||
logger.info(f"Patched delta-rule fns for SSM capture: {patched_targets}")
|
||||
else:
|
||||
logger.warning(
|
||||
"GDN SSM-capture patch installed no targets — hybrid models may miss ssm state"
|
||||
)
|
||||
|
||||
|
||||
def init_gdn_layer_order(kv_caches: Any) -> None:
|
||||
"""Identify hybrid layers (those with list/tuple kv_cache entries)."""
|
||||
_gdn_layer_order.clear()
|
||||
for li in range(len(kv_caches)):
|
||||
kv = kv_caches[li]
|
||||
if isinstance(kv, (list, tuple)) and len(kv) > 1: # pyright: ignore[reportUnknownArgumentType]
|
||||
_gdn_layer_order.append(li)
|
||||
if _gdn_layer_order:
|
||||
logger.info(f"GDN layer order: {len(_gdn_layer_order)} hybrid layers detected")
|
||||
@@ -0,0 +1,63 @@
|
||||
from mlx_lm.tokenizer_utils import TokenizerWrapper
|
||||
from vllm.sampling_params import SamplingParams
|
||||
from vllm.v1.engine.llm_engine import LLMEngine
|
||||
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.text_generation import TextGenerationTaskParams
|
||||
from exo.worker.engines.mlx.utils_mlx import (
|
||||
apply_chat_template,
|
||||
get_eos_token_ids_for_model,
|
||||
)
|
||||
|
||||
|
||||
def format_vllm_prompt(
|
||||
engine: LLMEngine, params: TextGenerationTaskParams
|
||||
) -> tuple[list[int], str, int]:
|
||||
# we should have our own wrapper
|
||||
# (instead of abusing mlx's TokenizerWrapper, use tokenizers Tokenizer)
|
||||
tokenizer = TokenizerWrapper(engine.get_tokenizer())
|
||||
prompt_text = apply_chat_template(tokenizer, params)
|
||||
token_ids: list[int] = tokenizer.encode(prompt_text, add_special_tokens=False)
|
||||
return token_ids, prompt_text, len(token_ids)
|
||||
|
||||
|
||||
def make_vllm_sampling_params(
|
||||
engine: LLMEngine,
|
||||
params: TextGenerationTaskParams,
|
||||
model_id: ModelId | None = None,
|
||||
) -> SamplingParams:
|
||||
kwargs: SamplingParams = SamplingParams()
|
||||
|
||||
if params.max_output_tokens is not None:
|
||||
kwargs.max_tokens = params.max_output_tokens
|
||||
else:
|
||||
kwargs.max_tokens = min(engine.model_config.max_model_len, 32168)
|
||||
if params.temperature is not None:
|
||||
kwargs.temperature = params.temperature
|
||||
if params.top_p is not None:
|
||||
kwargs.top_p = params.top_p
|
||||
if params.top_k is not None:
|
||||
kwargs.top_k = params.top_k
|
||||
if params.min_p is not None:
|
||||
kwargs.min_p = params.min_p
|
||||
if params.stop is not None:
|
||||
kwargs.stop = params.stop
|
||||
if params.seed is not None:
|
||||
kwargs.seed = params.seed
|
||||
if params.repetition_penalty is not None:
|
||||
kwargs.repetition_penalty = params.repetition_penalty
|
||||
if params.logprobs:
|
||||
kwargs.logprobs = params.top_logprobs or 1
|
||||
|
||||
if model_id is not None:
|
||||
extra_stop = get_eos_token_ids_for_model(model_id)
|
||||
if extra_stop:
|
||||
kwargs.stop_token_ids = extra_stop
|
||||
|
||||
if params.bench:
|
||||
kwargs.ignore_eos = True
|
||||
kwargs.min_tokens = kwargs.max_tokens
|
||||
if not params.use_prefix_cache:
|
||||
kwargs.skip_reading_prefix_cache = True
|
||||
|
||||
return kwargs
|
||||
@@ -5,7 +5,7 @@ import loguru
|
||||
|
||||
from exo.shared.types.events import Event, RunnerStatusUpdated
|
||||
from exo.shared.types.tasks import Task, TaskId
|
||||
from exo.shared.types.worker.instances import BoundInstance
|
||||
from exo.shared.types.worker.instances import BoundInstance, VllmInstance
|
||||
from exo.shared.types.worker.runners import RunnerFailed
|
||||
from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender
|
||||
from exo.worker.engines.base import Builder
|
||||
@@ -46,6 +46,21 @@ def entrypoint(
|
||||
builder = MfluxBuilder(
|
||||
event_sender, cancel_receiver, bound_instance.bound_shard
|
||||
)
|
||||
elif isinstance(bound_instance.instance, VllmInstance):
|
||||
from exo.worker.engines.vllm.builder import VllmBuilder
|
||||
from exo.worker.engines.vllm.growable_cache import patch_vllm
|
||||
|
||||
os.environ["VLLM_ENABLE_V1_MULTIPROCESSING"] = "0"
|
||||
os.environ["VLLM_KV_CACHE_LAYOUT"] = "NHD"
|
||||
os.environ["VLLM_BATCH_INVARIANT"] = "1"
|
||||
os.environ.setdefault("FASTSAFETENSORS_NOGDS", "1")
|
||||
|
||||
patch_vllm()
|
||||
builder = VllmBuilder(
|
||||
bound_instance.bound_shard.model_card.model_id,
|
||||
event_sender,
|
||||
cancel_receiver,
|
||||
)
|
||||
else:
|
||||
from exo.worker.engines.mlx.patches import apply_mlx_patches
|
||||
|
||||
|
||||
@@ -221,7 +221,6 @@ class SequentialGenerator(Engine):
|
||||
apply_chat_template(self.tokenizer, task.task_params),
|
||||
self.tool_parser,
|
||||
self.tokenizer,
|
||||
type(self.model),
|
||||
self.model_id,
|
||||
task.task_params.tools,
|
||||
)
|
||||
@@ -418,7 +417,6 @@ class BatchGenerator(Engine):
|
||||
apply_chat_template(self.tokenizer, task.task_params),
|
||||
self.tool_parser,
|
||||
self.tokenizer,
|
||||
type(self.model),
|
||||
self.model_id,
|
||||
task.task_params.tools,
|
||||
)
|
||||
|
||||
@@ -2,9 +2,6 @@ from collections.abc import Callable, Generator, Iterator
|
||||
from functools import cache
|
||||
from typing import Any
|
||||
|
||||
from mlx_lm.models.deepseek_v4 import Model as DeepseekV4Model
|
||||
from mlx_lm.models.deepseek_v32 import Model as DeepseekV32Model
|
||||
from mlx_lm.models.gpt_oss import Model as GptOssModel
|
||||
from mlx_lm.tokenizer_utils import TokenizerWrapper
|
||||
from openai_harmony import ( # pyright: ignore[reportMissingTypeStubs]
|
||||
HarmonyEncodingName,
|
||||
@@ -23,7 +20,6 @@ from exo.shared.types.chunks import (
|
||||
)
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.worker.runner_response import GenerationResponse, ToolCallResponse
|
||||
from exo.worker.engines.mlx.types import Model
|
||||
from exo.worker.engines.mlx.utils_mlx import (
|
||||
detect_thinking_prompt_suffix,
|
||||
)
|
||||
@@ -69,16 +65,15 @@ def apply_all_parsers(
|
||||
prompt: str,
|
||||
tool_parser: ToolParser | None,
|
||||
tokenizer: TokenizerWrapper,
|
||||
model_type: type[Model],
|
||||
model_id: ModelId,
|
||||
tools: list[dict[str, Any]] | None,
|
||||
) -> Iterator[GenerationChunk | None]:
|
||||
generator = receiver
|
||||
|
||||
normalized_id = model_id.normalize().lower()
|
||||
if issubclass(model_type, GptOssModel):
|
||||
normalized_id = model_id.short().lower()
|
||||
if "gpt-oss" in normalized_id:
|
||||
generator = parse_gpt_oss(generator)
|
||||
elif issubclass(model_type, DeepseekV32Model) and "deepseek" in normalized_id:
|
||||
elif "deepseek-v3.2" in normalized_id:
|
||||
if tokenizer.has_thinking:
|
||||
generator = parse_thinking_models(
|
||||
generator,
|
||||
@@ -87,7 +82,7 @@ def apply_all_parsers(
|
||||
starts_in_thinking=detect_thinking_prompt_suffix(prompt, tokenizer),
|
||||
)
|
||||
generator = parse_deepseek_v32(generator)
|
||||
elif issubclass(model_type, DeepseekV4Model) and "deepseek-v4" in normalized_id:
|
||||
elif "deepseek-v4" in normalized_id:
|
||||
if tokenizer.has_thinking:
|
||||
generator = parse_thinking_models(
|
||||
generator,
|
||||
|
||||
@@ -178,6 +178,10 @@ class Runner:
|
||||
|
||||
def _serve_prefill(self, req: PrefillTask) -> None:
|
||||
req.started.set()
|
||||
nested = isinstance(self.current_status, RunnerRunning)
|
||||
if not nested:
|
||||
self.update_status(RunnerRunning())
|
||||
logger.info("runner running")
|
||||
try:
|
||||
assert isinstance(self.generator, Engine)
|
||||
self.generator.serve_prefill(req.request, req.wfile)
|
||||
@@ -187,6 +191,11 @@ class Runner:
|
||||
)
|
||||
finally:
|
||||
req.done.set()
|
||||
if not nested:
|
||||
self.update_status(
|
||||
RunnerReady(prefill_server_port=self._prefill_server_port)
|
||||
)
|
||||
logger.info("runner ready")
|
||||
|
||||
def update_status(self, status: RunnerStatus):
|
||||
self.current_status = status
|
||||
|
||||
Reference in new issue
Block a user