mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-12 13:27:43 -04:00
Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
13ce4e9052 | ||
|
|
4ea5244e85 |
No files matched your search
@@ -120,6 +120,81 @@ From .cursorrules:
|
||||
|
||||
Tests use pytest-asyncio with `asyncio_mode = "auto"`. Tests are in `tests/` subdirectories alongside the code they test. The `EXO_TESTS=1` env var is set during tests.
|
||||
|
||||
Integration tests live in `tests/` (root) and are opt-in via `--ignore=tests` in the default pytest addopts. They require an `eco`-managed cluster:
|
||||
|
||||
```bash
|
||||
uv run pytest tests/ -v # constraint-driven host pick
|
||||
uv run pytest tests/ -v --hosts s4 # explicit host override
|
||||
```
|
||||
|
||||
## Benchmarking
|
||||
|
||||
Benchmarks live in `bench/`. The framework is a CLI with subcommands; each benchmark is a small library module under `bench/lib/<name>.py` plus a CLI front-end under `bench/cli/<name>.py`.
|
||||
|
||||
```
|
||||
bench/
|
||||
├── lib/ # composable, typed building blocks
|
||||
│ ├── prompt.py # PromptSizer, load_tokenizer_for_bench
|
||||
│ ├── completion.py # run_one_completion + typed payloads
|
||||
│ ├── session.py # BenchSession (cluster + client + instance)
|
||||
│ ├── results.py # RunMetadata, ResultsBundle, JSON schema
|
||||
│ ├── model_meta.py # HF API: total weights size, max context, layers
|
||||
│ ├── cluster.py # managed_cluster + managed_instance ctx-managers
|
||||
│ └── context_scaling.py # prompt-TPS / decode-TPS vs context-size sweep
|
||||
├── cli/ # CLI subcommands
|
||||
│ ├── _common.py # shared argparse args + SharedOptions
|
||||
│ ├── context_scaling.py # `python -m bench.cli context-scaling …`
|
||||
│ └── __main__.py # subcommand dispatcher
|
||||
└── exo_bench.py, prefill_decode_bench.py
|
||||
# legacy CLI scripts; PromptSizer / run_one_completion
|
||||
# / load_tokenizer_for_bench are re-exports of bench.lib.
|
||||
```
|
||||
|
||||
Run a benchmark:
|
||||
|
||||
```bash
|
||||
# Defaults assume a multi-node, Thunderbolt-connected cluster with tensor
|
||||
# parallelism + JACCL: --sharding Tensor --comm MlxJaccl --thunderbolt a2a.
|
||||
# Memory + disk minimums are auto-derived from HF metadata.
|
||||
uv run python -m bench.cli context-scaling \
|
||||
--model mlx-community/Qwen3-30B-A3B-4bit --nodes 2 --num-steps 32
|
||||
|
||||
# Single-node smoke: opt out of TB / tensor / jaccl
|
||||
uv run python -m bench.cli context-scaling --hosts s4 \
|
||||
--model mlx-community/Llama-3.2-1B-Instruct-4bit --num-steps 4 \
|
||||
--sharding Pipeline --comm MlxRing --thunderbolt none
|
||||
|
||||
# From a TOML config (CLI flags override config values)
|
||||
uv run python -m bench.cli context-scaling \
|
||||
--config bench/configs/context_scaling.example.toml --hosts s4,s9
|
||||
```
|
||||
|
||||
Shared CLI flags (every subcommand inherits these via `bench/cli/_common.py`):
|
||||
- `--config <path>.toml` — load run parameters from a TOML file
|
||||
- `--model`, `--sharding {Pipeline,Tensor}` (default Tensor), `--comm {MlxRing,MlxJaccl}` (default MlxJaccl), `--min-nodes` — placement
|
||||
- `--hosts`, `--nodes` (number of cluster hosts; distinct from `--min-nodes`), `--thunderbolt {a2a,ring,none}` (default a2a), `--chip` — host pool
|
||||
- `--min-memory-gb`, `--max-memory-gb`, `--min-disk-gb`, `--max-disk-gb` (minimums auto-derived from HF model size when not supplied)
|
||||
- `--evict-downloads` (default on; auto-evicts smallest-first when disk is short)
|
||||
- `--cleanup-instance` (default on; deletes the instance on exit)
|
||||
- `--output-dir`, `--tag key=value` (repeatable)
|
||||
|
||||
Run a multi-run campaign from a single TOML file (each `[[runs]]` = its own cluster deploy + bench + teardown; `[defaults]` is shared, per-run keys override; `[plot]` triggers a comparison PNG):
|
||||
|
||||
```bash
|
||||
uv run python -m bench.cli campaign bench/configs/llama-family-smoke.toml
|
||||
```
|
||||
|
||||
Plot any results JSON to a PNG (auto-detects benchmark type from `metadata.benchmark`):
|
||||
|
||||
```bash
|
||||
uv run python -m bench.cli plot bench/results/context_scaling/latest.json
|
||||
uv run python -m bench.cli plot a.json b.json --label-tag operator # multi-run comparison
|
||||
```
|
||||
|
||||
Adding a new benchmark = (1) write a `bench/lib/<name>.py` exposing a typed `run(session, params, bundle)` callable; (2) add a `bench/cli/<name>.py` with `add_subparser(...)` + `run(args) -> Path`; (3) register the imports in `bench/cli/__main__.py`. To enable plotting for the new benchmark, add a `render_<name>(inputs)` function in `bench/lib/plotting.py` and a dispatch entry in `bench/cli/plot.py::run`.
|
||||
|
||||
Results land at `bench/results/<benchmark>/<run_id>.json` (with a `latest.json` symlink alongside) containing metadata (exo SHA, hostname, platform, ISO timestamps, methodology version, user tags), full cluster snapshot, the resolved + derived params, per-step rows, optional cold-control rows, and any derived summaries (e.g. `t_cum_seconds[]` for context-scaling).
|
||||
|
||||
## Dashboard UI Testing & Screenshots
|
||||
|
||||
### Building and Running the Dashboard
|
||||
|
||||
@@ -550,6 +550,120 @@ uv run bench/exo_bench.py \
|
||||
|
||||
The tool outputs performance metrics including prompt tokens per second (prompt_tps), generation tokens per second (generation_tps), and peak memory usage for each configuration.
|
||||
|
||||
### Composable benchmarks (CLI)
|
||||
|
||||
For benchmarks that need an `eco`-managed cluster and a stable JSON result format, exo ships a CLI under `bench/cli/`. The CLI handles cluster lifecycle, instance placement, model-metadata resolution (HuggingFace), and result capture; benchmark logic lives in `bench/lib/` so each new benchmark is a small library module + a CLI subcommand.
|
||||
|
||||
**Run the prompt-TPS / decode-TPS vs context-size sweep:**
|
||||
|
||||
The defaults assume a multi-node, Thunderbolt-connected cluster with tensor parallelism + JACCL — the typical exo benchmarking setup:
|
||||
|
||||
```bash
|
||||
# Defaults: --sharding Tensor --comm MlxJaccl --thunderbolt a2a, with
|
||||
# memory/disk minimums auto-derived from the HF model size. eco picks
|
||||
# `--nodes` hosts from its inventory that form a TB clique and satisfy
|
||||
# those constraints.
|
||||
uv run python -m bench.cli context-scaling \
|
||||
--model mlx-community/Qwen3-30B-A3B-4bit --nodes 2 --num-steps 32
|
||||
|
||||
# Pin to specific hosts (defaults still apply for sharding/comm/topology)
|
||||
uv run python -m bench.cli context-scaling --hosts s4,s9 \
|
||||
--model mlx-community/Qwen3-30B-A3B-4bit --num-steps 16
|
||||
|
||||
# Single-node smoke test: explicit single-node placement overrides
|
||||
uv run python -m bench.cli context-scaling --hosts s4 \
|
||||
--model mlx-community/Llama-3.2-1B-Instruct-4bit --num-steps 4 \
|
||||
--sharding Pipeline --comm MlxRing --thunderbolt none
|
||||
|
||||
# Override the auto-derived ramp / cold controls
|
||||
uv run python -m bench.cli context-scaling --hosts s4,s9 --model X \
|
||||
--pp-step 4096 --num-steps 32 --cold-controls 8192,32768,65536,131072
|
||||
|
||||
# Custom output dir + tags
|
||||
uv run python -m bench.cli context-scaling --hosts s4,s9 --model X \
|
||||
--output-dir bench/results/2026-05-10/ --tag operator=$USER --tag run=full
|
||||
|
||||
# Run from a TOML config (CLI flags override values from the file)
|
||||
uv run python -m bench.cli context-scaling \
|
||||
--config bench/configs/context_scaling.example.toml
|
||||
```
|
||||
|
||||
**Shared flags (every benchmark subcommand has these):**
|
||||
|
||||
- `--model` — HuggingFace model id (required)
|
||||
- `--config <path>.toml` — load run parameters from a TOML file
|
||||
- `--sharding {Pipeline,Tensor}` (default **Tensor**) — sharding mode
|
||||
- `--comm {MlxRing,MlxJaccl}` (default **MlxJaccl**) — inter-node comm mode
|
||||
- `--min-nodes N` (default 1) — minimum nodes for the placement
|
||||
- `--hosts s4,s9` — pin to specific hosts; bypasses constraint search
|
||||
- `--nodes N` (default 1) — number of cluster hosts to reserve when `--hosts` is unset (distinct from `--min-nodes`, which controls the model's instance placement)
|
||||
- `--thunderbolt {a2a,ring,none}` (default **a2a**) — required Thunderbolt topology
|
||||
- `--chip "M3 Ultra"` — required chip (substring match; comment to allow any)
|
||||
- `--min-memory-gb`, `--max-memory-gb`, `--min-disk-gb`, `--max-disk-gb` — host RAM / disk constraints. The minimums are auto-derived from the HF model size (×1.30 + 1 GiB for memory, ×1.10 + 1 GiB for disk) when not supplied; explicit values always win.
|
||||
- `--evict-downloads` (default **on**) — auto-evict existing models smallest-first on disk-full; pass `--no-evict-downloads` to keep
|
||||
- `--cleanup-instance` (default **on**) — delete the placed instance after exit; pass `--no-cleanup-instance` to leave it running for debugging
|
||||
- `--output-dir bench/results` — base directory for JSON results (subcommands add their own subfolder)
|
||||
- `--tag key=value` — append to `metadata.tags` (repeatable)
|
||||
|
||||
**Context-scaling-specific flags:**
|
||||
|
||||
- `--num-steps N` — number of equally-spaced ramp points (default 32)
|
||||
- `--pp-step Δ` — explicit Δ in tokens (overrides auto-derivation from `max_position_embeddings`)
|
||||
- `--fraction-of-max F` — when Δ is auto-derived, use `F × max_context` as the upper bound
|
||||
- `--tg` — tokens generated per step (default 64)
|
||||
- `--warmup` — warmup requests at `pp=Δ` (default 1)
|
||||
- `--cold-controls auto` (4 evenly-spaced points across the ramp) **or** `--cold-controls 8192,32768,…` (explicit pp values). Default: no cold controls.
|
||||
|
||||
**Output:** each run writes `bench/results/<benchmark>/<run_id>.json` plus a `latest.json` symlink. The JSON contains metadata (exo SHA, hostname, platform, user tags), the full cluster snapshot at run start, the resolved + derived params, per-step rows, cold-control rows, and derived summaries (`t_cum_seconds`, `control_gaps`).
|
||||
|
||||
**Multi-run campaigns** — `bench campaign` runs a list of bench invocations from a single TOML file. Each `[[runs]]` entry is its own cluster deploy + bench + teardown, with a shared `[defaults]` table for DRY config:
|
||||
|
||||
```toml
|
||||
# bench/configs/llama-family-smoke.toml
|
||||
[defaults]
|
||||
nodes = 4
|
||||
num_steps = 8
|
||||
fraction_of_max = 0.5
|
||||
|
||||
[[runs]]
|
||||
subcommand = "context-scaling"
|
||||
model = "mlx-community/Llama-3.2-3B-Instruct-4bit"
|
||||
[runs.tags]
|
||||
model_short = "llama-3.2-3b-4bit"
|
||||
|
||||
[[runs]]
|
||||
subcommand = "context-scaling"
|
||||
model = "mlx-community/Meta-Llama-3.1-8B-Instruct-4bit"
|
||||
[runs.tags]
|
||||
model_short = "llama-3.1-8b-4bit"
|
||||
|
||||
[plot]
|
||||
label_tag = "model_short"
|
||||
```
|
||||
|
||||
```bash
|
||||
uv run python -m bench.cli campaign bench/configs/llama-family-smoke.toml
|
||||
```
|
||||
|
||||
After all runs finish, an optional `[plot]` table triggers a comparison plot per benchmark group (one PNG per benchmark type with ≥2 runs).
|
||||
|
||||
**Plotting** — `bench plot` renders any results JSON to a 2-panel PNG (prompt_tps + generation_tps vs pp_tokens, cold controls overlaid as 'x' markers):
|
||||
|
||||
```bash
|
||||
# Plot the most recent run next to its JSON
|
||||
uv run python -m bench.cli plot bench/results/context_scaling/latest.json
|
||||
|
||||
# Compare multiple runs (one line per run; legend label = the chosen tag)
|
||||
uv run python -m bench.cli plot run_a.json run_b.json --label-tag operator
|
||||
|
||||
# Custom output path + title
|
||||
uv run python -m bench.cli plot run.json --output /tmp/scaling.png --title "30B 4-node sweep"
|
||||
```
|
||||
|
||||
The benchmark type is detected from each JSON's `metadata.benchmark`, so the same `plot` command will work for future benchmarks once their renderer is registered in `bench/lib/plotting.py`.
|
||||
|
||||
Methodology for the context-scaling benchmark is documented in detail in `bench/lib/context_scaling.py`'s module docstring and in `bench/METHODOLOGY.md`.
|
||||
|
||||
---
|
||||
|
||||
## Hardware Accelerator Support
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
"""CLI front-ends for bench library benchmarks.
|
||||
|
||||
Each benchmark is a sub-package / module with two pieces:
|
||||
|
||||
- a ``run(...)`` callable in ``bench.lib.<name>`` that does the actual
|
||||
measurement (no argparse, no eco, no I/O)
|
||||
- an ``add_subparser(subparsers)`` helper here that wires CLI args to a
|
||||
handler invoking the lib
|
||||
|
||||
The main entry point dispatches to the requested subcommand:
|
||||
|
||||
uv run python -m bench.cli context-scaling --hosts s4 \\
|
||||
--model mlx-community/Qwen3-30B-A3B-4bit
|
||||
"""
|
||||
@@ -0,0 +1,56 @@
|
||||
"""``python -m bench.cli`` — dispatcher for benchmark subcommands.
|
||||
|
||||
To add a new benchmark:
|
||||
|
||||
1. Implement the methodology in ``bench.lib.<name>`` exposing a typed
|
||||
``run(session, params, bundle)`` callable (no argparse, no eco I/O).
|
||||
2. Implement a ``bench.cli.<name>`` module with an ``add_subparser`` and
|
||||
a ``run(args) -> Path`` handler.
|
||||
3. Add an ``import + add_subparser(subparsers)`` line below.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
|
||||
from bench.cli import campaign, context_scaling, plot
|
||||
from bench.cli._common import expand_config_in_argv
|
||||
|
||||
|
||||
def _build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="python -m bench.cli",
|
||||
description=(
|
||||
"Composable, eco-managed benchmarks for exo. "
|
||||
"Pick a subcommand and pass its model / cluster options."
|
||||
),
|
||||
)
|
||||
subparsers = parser.add_subparsers(
|
||||
dest="subcommand",
|
||||
required=True,
|
||||
metavar="SUBCOMMAND",
|
||||
)
|
||||
context_scaling.add_subparser(subparsers)
|
||||
plot.add_subparser(subparsers)
|
||||
campaign.add_subparser(subparsers)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
raw_argv = list(argv if argv is not None else sys.argv[1:])
|
||||
expanded = expand_config_in_argv(raw_argv)
|
||||
args = _build_parser().parse_args(expanded)
|
||||
handler = getattr(args, "handler", None)
|
||||
if not callable(handler):
|
||||
subcommand = getattr(args, "subcommand", "<unknown>")
|
||||
raise SystemExit(f"subcommand {subcommand!r} did not register a handler")
|
||||
cast("Callable[[argparse.Namespace], Path]", handler)(args)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,345 @@
|
||||
"""Shared CLI argument parsing for the bench command-line interface.
|
||||
|
||||
Every benchmark subcommand inherits the same model / cluster / output
|
||||
arguments via :func:`add_shared_args` and consumes them through
|
||||
:class:`SharedOptions`. argparse's ``Namespace.<attr>`` is fundamentally
|
||||
typed ``Any``; the :func:`get_arg` / :func:`get_arg_optional` helpers are
|
||||
the single boundary where we coerce to typed values.
|
||||
|
||||
A ``--config <path>.toml`` flag lets the caller capture a run definition
|
||||
in a TOML file. :func:`expand_config_in_argv` rewrites argv in place,
|
||||
substituting the config's keys as CLI flags placed *before* any explicit
|
||||
user args so that explicit CLI flags always win.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import tomllib
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, TypeVar
|
||||
|
||||
from exo_tools.cluster import Chip, Thunderbolt
|
||||
from exo_tools.harness import Comm, Sharding
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
def get_arg(args: argparse.Namespace, name: str, type_: type[_T]) -> _T:
|
||||
"""Return ``args.<name>``, asserting it's an instance of ``type_``.
|
||||
|
||||
For ``int`` and ``float`` we additionally accept inputs that ``int(.)`` /
|
||||
``float(.)`` would parse, since argparse's ``type=int`` already coerces
|
||||
cleanly on input but post-`set_defaults` callers may pass raw values.
|
||||
"""
|
||||
raw: Any = getattr(args, name) # type: ignore[reportAny]
|
||||
if isinstance(raw, type_):
|
||||
return raw
|
||||
if type_ is int and isinstance(raw, (int, str)):
|
||||
return int(raw) # type: ignore[return-value]
|
||||
if type_ is float and isinstance(raw, (int, float, str)):
|
||||
return float(raw) # type: ignore[return-value]
|
||||
raise TypeError(
|
||||
f"argparse field {name!r} expected {type_.__name__}, got {type(raw).__name__}" # type: ignore[reportUnknownArgumentType]
|
||||
)
|
||||
|
||||
|
||||
def get_arg_optional(args: argparse.Namespace, name: str, type_: type[_T]) -> _T | None:
|
||||
"""Like :func:`get_arg` but allows the field to be missing or None."""
|
||||
raw = getattr(args, name, None)
|
||||
if raw is None:
|
||||
return None
|
||||
if isinstance(raw, type_):
|
||||
return raw
|
||||
if type_ is int and isinstance(raw, (int, str)):
|
||||
return int(raw) # type: ignore[return-value]
|
||||
if type_ is float and isinstance(raw, (int, float, str)):
|
||||
return float(raw) # type: ignore[return-value]
|
||||
raise TypeError(
|
||||
f"argparse field {name!r} expected {type_.__name__} or None, "
|
||||
f"got {type(raw).__name__}" # type: ignore[reportUnknownArgumentType]
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SharedOptions:
|
||||
"""Parsed shared CLI options for any benchmark."""
|
||||
|
||||
model: str
|
||||
hosts: tuple[str, ...]
|
||||
nodes: int
|
||||
thunderbolt: Thunderbolt | None
|
||||
chip: Chip | None
|
||||
min_memory_gb: float | None
|
||||
max_memory_gb: float | None
|
||||
min_disk_gb: float | None
|
||||
max_disk_gb: float | None
|
||||
evict_downloads: bool
|
||||
sharding: Sharding
|
||||
comm: Comm
|
||||
min_nodes: int
|
||||
output_dir: Path
|
||||
tags: dict[str, str]
|
||||
cleanup_instance: bool
|
||||
user_prefix: str
|
||||
|
||||
@classmethod
|
||||
def from_namespace(cls, args: argparse.Namespace) -> SharedOptions:
|
||||
hosts_raw = get_arg_optional(args, "hosts", str)
|
||||
thunderbolt_raw = get_arg_optional(args, "thunderbolt", str)
|
||||
chip_raw = get_arg_optional(args, "chip", str)
|
||||
tag_list_raw: object = getattr(args, "tag", None) or []
|
||||
if isinstance(tag_list_raw, list):
|
||||
tag_list: list[str] = [
|
||||
str(t) # type: ignore[reportUnknownArgumentType]
|
||||
for t in tag_list_raw # type: ignore[reportUnknownVariableType]
|
||||
]
|
||||
else:
|
||||
tag_list = []
|
||||
return cls(
|
||||
model=get_arg(args, "model", str),
|
||||
hosts=tuple(_parse_csv(hosts_raw)) if hosts_raw else (),
|
||||
nodes=get_arg(args, "nodes", int),
|
||||
thunderbolt=Thunderbolt(thunderbolt_raw) if thunderbolt_raw else None,
|
||||
chip=Chip(chip_raw) if chip_raw else None,
|
||||
min_memory_gb=get_arg_optional(args, "min_memory_gb", float),
|
||||
max_memory_gb=get_arg_optional(args, "max_memory_gb", float),
|
||||
min_disk_gb=get_arg_optional(args, "min_disk_gb", float),
|
||||
max_disk_gb=get_arg_optional(args, "max_disk_gb", float),
|
||||
evict_downloads=get_arg(args, "evict_downloads", bool),
|
||||
sharding=Sharding(get_arg(args, "sharding", str)),
|
||||
comm=Comm(get_arg(args, "comm", str)),
|
||||
min_nodes=get_arg(args, "min_nodes", int),
|
||||
output_dir=Path(get_arg(args, "output_dir", str)),
|
||||
tags=_parse_tags(tag_list),
|
||||
cleanup_instance=get_arg(args, "cleanup_instance", bool),
|
||||
user_prefix=get_arg(args, "eco_user_prefix", str),
|
||||
)
|
||||
|
||||
|
||||
def add_shared_args(parser: argparse.ArgumentParser) -> None:
|
||||
"""Register the shared-arg group on ``parser``.
|
||||
|
||||
The bool flags (``--auto-constrain``, ``--evict-downloads``,
|
||||
``--cleanup-instance``) all default to True and use
|
||||
:class:`argparse.BooleanOptionalAction` so callers opt out via the
|
||||
``--no-X`` form (or set ``X = false`` in a TOML config).
|
||||
"""
|
||||
g_config = parser.add_argument_group("config file")
|
||||
g_config.add_argument(
|
||||
"--config",
|
||||
default=None,
|
||||
help="TOML file with run parameters. CLI flags placed after --config "
|
||||
"override values from the file.",
|
||||
)
|
||||
|
||||
g_model = parser.add_argument_group("model")
|
||||
g_model.add_argument(
|
||||
"--model",
|
||||
required=True,
|
||||
help="HuggingFace model id. To run multiple models in one go, use "
|
||||
"the 'campaign' subcommand with a TOML file listing each as a "
|
||||
"separate [[runs]] entry.",
|
||||
)
|
||||
g_model.add_argument(
|
||||
"--sharding",
|
||||
default=Sharding.TENSOR.value,
|
||||
choices=[s.value for s in Sharding],
|
||||
help="Sharding mode for the placed instance. Default 'Tensor' (splits "
|
||||
"layers within nodes; pairs with --comm MlxJaccl for high throughput "
|
||||
"on TB-connected clusters). Use 'Pipeline' for layer-per-node sharding "
|
||||
"(typical for single-node smoke tests).",
|
||||
)
|
||||
g_model.add_argument(
|
||||
"--comm",
|
||||
default=Comm.JACCL.value,
|
||||
choices=[c.value for c in Comm],
|
||||
help="Inter-node communication mode. Default 'MlxJaccl' (RDMA over "
|
||||
"Thunderbolt; pairs with --sharding Tensor and --thunderbolt a2a). "
|
||||
"Use 'MlxRing' for ring all-reduce over the regular network.",
|
||||
)
|
||||
g_model.add_argument("--min-nodes", type=int, default=1)
|
||||
|
||||
g_cluster = parser.add_argument_group("cluster")
|
||||
g_cluster.add_argument(
|
||||
"--hosts",
|
||||
default=None,
|
||||
help="Comma-separated host list (e.g. s4,s9). Bypasses constraint search.",
|
||||
)
|
||||
g_cluster.add_argument(
|
||||
"--nodes",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of cluster nodes (hosts) to deploy on. "
|
||||
"Distinct from --min-nodes which controls the model's instance placement.",
|
||||
)
|
||||
g_cluster.add_argument(
|
||||
"--thunderbolt",
|
||||
default=Thunderbolt.A2A.value,
|
||||
choices=[t.value for t in Thunderbolt],
|
||||
help="Thunderbolt topology required: 'a2a' (clique, default; needed "
|
||||
"for tensor parallelism + JACCL), 'ring' (cycle; for pipeline + JACCL), "
|
||||
"or 'none' (exclude TB-connected hosts; pair with --sharding Pipeline "
|
||||
"--comm MlxRing).",
|
||||
)
|
||||
g_cluster.add_argument(
|
||||
"--chip",
|
||||
default=None,
|
||||
choices=[c.value for c in Chip],
|
||||
help="Chip required (e.g. 'M3 Ultra')",
|
||||
)
|
||||
g_cluster.add_argument(
|
||||
"--min-memory-gb",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Min RAM (GB) on each host. If unset, auto-derived from the HF "
|
||||
"model size (×1.30 + 1 GiB).",
|
||||
)
|
||||
g_cluster.add_argument(
|
||||
"--max-memory-gb",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Max RAM (GB) on each host. Useful to leave bigger machines "
|
||||
"free for other workloads.",
|
||||
)
|
||||
g_cluster.add_argument(
|
||||
"--min-disk-gb",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Min free disk (GB) on each host. If unset, auto-derived from "
|
||||
"the HF model size (×1.10 + 1 GiB).",
|
||||
)
|
||||
g_cluster.add_argument(
|
||||
"--max-disk-gb",
|
||||
type=float,
|
||||
default=None,
|
||||
help="Max disk (GB) on each host.",
|
||||
)
|
||||
|
||||
g_runtime = parser.add_argument_group("runtime")
|
||||
g_runtime.add_argument(
|
||||
"--evict-downloads",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Auto-evict existing models (smallest first) when disk is short to "
|
||||
"make room for the bench model. Default on; pass --no-evict-downloads "
|
||||
"to keep existing downloads.",
|
||||
)
|
||||
g_runtime.add_argument(
|
||||
"--cleanup-instance",
|
||||
action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Clean up the placed instance after the benchmark exits. "
|
||||
"Default on; pass --no-cleanup-instance to leave it running for debugging.",
|
||||
)
|
||||
g_runtime.add_argument(
|
||||
"--eco-user-prefix",
|
||||
default="bench",
|
||||
help="USER prefix for the eco session (default: 'bench').",
|
||||
)
|
||||
|
||||
g_output = parser.add_argument_group("output")
|
||||
g_output.add_argument(
|
||||
"--output-dir",
|
||||
default="bench/results",
|
||||
help="Base directory for JSON results. Subcommands may add a sub-folder.",
|
||||
)
|
||||
g_output.add_argument(
|
||||
"--tag",
|
||||
action="append",
|
||||
default=[],
|
||||
help="Add a 'key=value' tag to metadata.tags (repeatable).",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TOML config expansion
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def expand_config_in_argv(argv: list[str]) -> list[str]:
|
||||
"""If ``--config <path>`` appears in ``argv``, splice the TOML's contents in.
|
||||
|
||||
The TOML file's keys are converted to CLI flags (``foo_bar`` →
|
||||
``--foo-bar``) and inserted *before* the user's other args, so explicit
|
||||
CLI flags always override the config. The ``--config <path>`` pair
|
||||
itself is removed from argv. The first arg (the subcommand name) is
|
||||
preserved at index 0.
|
||||
|
||||
Special handling:
|
||||
- ``[tags]`` table → repeated ``--tag key=value`` occurrences
|
||||
- lists → joined as a comma-separated value (matches the parser's
|
||||
CSV handling for ``--hosts``)
|
||||
- bool true/false → ``--key`` / ``--no-key`` (assumes the underlying
|
||||
flag uses :class:`argparse.BooleanOptionalAction`)
|
||||
"""
|
||||
if "--config" not in argv:
|
||||
return list(argv)
|
||||
|
||||
idx = argv.index("--config")
|
||||
if idx + 1 >= len(argv):
|
||||
raise ValueError("--config requires a path argument")
|
||||
config_path = Path(argv[idx + 1])
|
||||
if not config_path.is_file():
|
||||
raise FileNotFoundError(f"Config file not found: {config_path}")
|
||||
|
||||
with config_path.open("rb") as f:
|
||||
config_data: dict[str, Any] = tomllib.load(f)
|
||||
|
||||
expanded = _config_to_argv(config_data)
|
||||
stripped = list(argv[:idx]) + list(argv[idx + 2 :])
|
||||
if not stripped:
|
||||
return expanded
|
||||
# The subcommand name must come first; insert config-derived args
|
||||
# right after it so that the user's later explicit args override.
|
||||
return [stripped[0]] + expanded + stripped[1:]
|
||||
|
||||
|
||||
def _config_to_argv(data: dict[str, Any]) -> list[str]:
|
||||
"""Convert a TOML-loaded dict to a list of argv-style CLI flags."""
|
||||
out: list[str] = []
|
||||
for key in data:
|
||||
value: Any = data[key] # type: ignore[reportAny]
|
||||
if key == "tags" and isinstance(value, dict):
|
||||
for tag_key, tag_value in value.items(): # type: ignore[reportUnknownVariableType]
|
||||
out.extend(["--tag", f"{tag_key}={tag_value}"])
|
||||
continue
|
||||
if value is None:
|
||||
continue
|
||||
flag = "--" + key.replace("_", "-")
|
||||
if isinstance(value, bool):
|
||||
out.append(flag if value else f"--no-{key.replace('_', '-')}")
|
||||
elif isinstance(value, list):
|
||||
joined = ",".join(
|
||||
str(x) # type: ignore[reportUnknownArgumentType]
|
||||
for x in value # type: ignore[reportUnknownVariableType]
|
||||
)
|
||||
out.extend([flag, joined])
|
||||
else:
|
||||
out.extend([flag, str(value)]) # type: ignore[reportAny]
|
||||
return out
|
||||
|
||||
|
||||
def _parse_csv(raw: str) -> list[str]:
|
||||
return [s.strip() for s in raw.split(",") if s.strip()]
|
||||
|
||||
|
||||
def _parse_tags(raw: list[str]) -> dict[str, str]:
|
||||
out: dict[str, str] = {}
|
||||
for entry in raw:
|
||||
if "=" not in entry:
|
||||
raise argparse.ArgumentTypeError(
|
||||
f"--tag must be 'key=value', got {entry!r}"
|
||||
)
|
||||
k, v = entry.split("=", 1)
|
||||
out[k.strip()] = v.strip()
|
||||
return out
|
||||
|
||||
|
||||
@dataclass
|
||||
class CommandResult:
|
||||
"""Return value from a benchmark CLI handler."""
|
||||
|
||||
output_path: Path | None = None
|
||||
extra: dict[str, str] = field(default_factory=dict)
|
||||
@@ -0,0 +1,220 @@
|
||||
"""Run a campaign of bench invocations from a single TOML file.
|
||||
|
||||
A campaign config has a ``[defaults]`` table (applied to every run) and a
|
||||
list of ``[[runs]]`` entries (each a fully-formed invocation with its own
|
||||
``subcommand``). The campaign runner merges defaults with each run's
|
||||
overrides, dispatches to the matching subcommand handler, and collects
|
||||
the output JSON paths.
|
||||
|
||||
Each run gets its own cluster — the deploy / teardown happens per-run.
|
||||
After all runs finish, an optional ``[plot]`` table triggers a comparison
|
||||
plot per benchmark group.
|
||||
|
||||
Schema::
|
||||
|
||||
[defaults]
|
||||
nodes = 4
|
||||
num_steps = 8
|
||||
|
||||
[[runs]]
|
||||
subcommand = "context-scaling"
|
||||
model = "mlx-community/Llama-3.2-3B-Instruct-4bit"
|
||||
[runs.tags]
|
||||
model_short = "llama-3.2-3b"
|
||||
|
||||
[[runs]]
|
||||
subcommand = "context-scaling"
|
||||
model = "mlx-community/Meta-Llama-3.1-8B-Instruct-4bit"
|
||||
[runs.tags]
|
||||
model_short = "llama-3.1-8b"
|
||||
|
||||
[plot]
|
||||
label_tag = "model_short"
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import tomllib
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from bench.cli import context_scaling
|
||||
from bench.cli._common import (
|
||||
_config_to_argv, # type: ignore[reportPrivateUsage]
|
||||
get_arg,
|
||||
)
|
||||
from bench.lib.plotting import PlotInputs, render_context_scaling
|
||||
|
||||
# Each subcommand exposes its argparse via add_subparser. The campaign
|
||||
# runner builds a one-off parser per run with only the chosen subcommand
|
||||
# registered, parses the run-derived argv, and invokes the handler.
|
||||
_SUBCOMMAND_PARSERS: dict[
|
||||
str,
|
||||
Callable[
|
||||
[Any], None
|
||||
], # subparsers action — argparse private; Any-typed at boundary
|
||||
] = {
|
||||
"context-scaling": context_scaling.add_subparser,
|
||||
}
|
||||
|
||||
|
||||
def add_subparser(
|
||||
subparsers: argparse._SubParsersAction[argparse.ArgumentParser], # type: ignore[type-arg]
|
||||
) -> None:
|
||||
parser = subparsers.add_parser(
|
||||
"campaign",
|
||||
help="Run a list of bench invocations from a single TOML config.",
|
||||
description=__doc__,
|
||||
)
|
||||
parser.add_argument(
|
||||
"config",
|
||||
type=str,
|
||||
help="TOML campaign file (with [defaults] + [[runs]] tables).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-plot",
|
||||
action="store_true",
|
||||
help="Skip the optional comparison plot at the end of the campaign.",
|
||||
)
|
||||
parser.set_defaults(handler=run)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run(args: argparse.Namespace) -> Path | None:
|
||||
config_path = Path(get_arg(args, "config", str))
|
||||
if not config_path.is_file():
|
||||
raise SystemExit(f"campaign: file not found: {config_path}")
|
||||
|
||||
with config_path.open("rb") as f:
|
||||
raw = tomllib.load(f)
|
||||
|
||||
defaults = _table(raw, "defaults")
|
||||
runs_obj = raw.get("runs")
|
||||
if not isinstance(runs_obj, list) or not runs_obj:
|
||||
raise SystemExit(f"campaign: {config_path}: missing or empty [[runs]] list")
|
||||
runs_raw: list[Any] = cast("list[Any]", runs_obj)
|
||||
plot_cfg = _table(raw, "plot")
|
||||
|
||||
n_runs = len(runs_raw)
|
||||
output_paths: dict[str, list[Path]] = {}
|
||||
for i, run_obj in enumerate(runs_raw): # type: ignore[reportAny]
|
||||
if not isinstance(run_obj, dict):
|
||||
raise SystemExit(
|
||||
f"campaign: run #{i + 1}: expected a TOML table, "
|
||||
f"got {type(run_obj).__name__}" # type: ignore[reportUnknownArgumentType]
|
||||
)
|
||||
run_cfg = cast("dict[str, Any]", run_obj)
|
||||
|
||||
merged = _merge(defaults, run_cfg)
|
||||
subcommand_obj: Any = merged.pop("subcommand", None) # type: ignore[reportAny]
|
||||
if not isinstance(subcommand_obj, str):
|
||||
raise SystemExit(
|
||||
f"campaign: run #{i + 1}: 'subcommand' field is required (str)"
|
||||
)
|
||||
if subcommand_obj not in _SUBCOMMAND_PARSERS:
|
||||
raise SystemExit(
|
||||
f"campaign: run #{i + 1}: unknown subcommand "
|
||||
f"{subcommand_obj!r} (have {sorted(_SUBCOMMAND_PARSERS)})"
|
||||
)
|
||||
|
||||
argv_for_run = _config_to_argv(merged)
|
||||
sub_args = _parse_for_subcommand(subcommand_obj, argv_for_run)
|
||||
handler = getattr(sub_args, "handler", None)
|
||||
if not callable(handler):
|
||||
raise SystemExit(f"campaign: subcommand {subcommand_obj!r} has no handler")
|
||||
|
||||
logger.info(
|
||||
f"campaign: starting run {i + 1}/{n_runs} "
|
||||
f"({subcommand_obj}; {len(merged)} flags)"
|
||||
)
|
||||
out = cast("Callable[[argparse.Namespace], Path]", handler)(sub_args)
|
||||
output_paths.setdefault(subcommand_obj, []).append(Path(out))
|
||||
logger.info(f"campaign: finished run {i + 1}/{n_runs} → {out}")
|
||||
|
||||
last_path: Path | None = None
|
||||
for paths in output_paths.values():
|
||||
if paths:
|
||||
last_path = paths[-1]
|
||||
|
||||
if get_arg(args, "no_plot", bool):
|
||||
return last_path
|
||||
|
||||
comparison = _render_comparisons(output_paths, plot_cfg)
|
||||
return comparison or last_path
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _table(data: dict[str, Any], key: str) -> dict[str, Any]:
|
||||
"""Return ``data[key]`` if it's a table, else an empty dict."""
|
||||
val: Any = data.get(key)
|
||||
return cast("dict[str, Any]", val) if isinstance(val, dict) else {}
|
||||
|
||||
|
||||
def _merge(defaults: dict[str, Any], run: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Shallow-merge ``defaults`` with ``run``; ``run`` wins on conflict.
|
||||
|
||||
The ``tags`` table is deep-merged (defaults' tags + run's tags) so a
|
||||
campaign-level operator tag and a per-run model_short tag both survive.
|
||||
"""
|
||||
merged: dict[str, Any] = {**defaults, **run}
|
||||
default_tags = _table(defaults, "tags")
|
||||
run_tags = _table(run, "tags")
|
||||
if default_tags or run_tags:
|
||||
merged["tags"] = {**default_tags, **run_tags}
|
||||
return merged
|
||||
|
||||
|
||||
def _parse_for_subcommand(
|
||||
subcommand: str, argv_for_run: list[str]
|
||||
) -> argparse.Namespace:
|
||||
"""Build a one-off parser with ``subcommand`` registered + parse argv."""
|
||||
parser = argparse.ArgumentParser(prog=f"bench campaign:{subcommand}")
|
||||
subparsers = parser.add_subparsers(dest="subcommand", required=True)
|
||||
_SUBCOMMAND_PARSERS[subcommand](subparsers)
|
||||
return parser.parse_args([subcommand] + argv_for_run)
|
||||
|
||||
|
||||
def _render_comparisons(
|
||||
output_paths: dict[str, list[Path]],
|
||||
plot_cfg: dict[str, Any],
|
||||
) -> Path | None:
|
||||
"""Render one comparison plot per benchmark group with ≥2 outputs."""
|
||||
label_tag = _str_or_none(plot_cfg.get("label_tag"))
|
||||
title = _str_or_none(plot_cfg.get("title"))
|
||||
|
||||
last: Path | None = None
|
||||
for subcommand, paths in output_paths.items():
|
||||
if len(paths) < 2:
|
||||
continue
|
||||
if subcommand != "context-scaling":
|
||||
logger.warning(
|
||||
f"campaign: no comparison renderer registered for {subcommand!r}; "
|
||||
"skipping comparison plot"
|
||||
)
|
||||
continue
|
||||
out = paths[0].with_name(f"campaign_{subcommand}_compare.png")
|
||||
last = render_context_scaling(
|
||||
PlotInputs(
|
||||
results=paths,
|
||||
output=out,
|
||||
label_tag=label_tag,
|
||||
title=title,
|
||||
)
|
||||
)
|
||||
logger.info(f"campaign: wrote comparison plot {last}")
|
||||
return last
|
||||
|
||||
|
||||
def _str_or_none(value: Any) -> str | None: # type: ignore[reportAny]
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
__all__ = ["add_subparser", "run"]
|
||||
@@ -0,0 +1,270 @@
|
||||
"""Context-scaling benchmark — CLI subcommand.
|
||||
|
||||
Wraps :func:`bench.lib.context_scaling.run` with:
|
||||
|
||||
- HF model-metadata resolution
|
||||
- Auto-derived constraints (memory, disk) and context ramp (Δ, K)
|
||||
- eco cluster + instance lifecycle (managed_cluster + managed_instance)
|
||||
- Cold-control isolation (delete sweep instance before controls)
|
||||
- JSON results + ``latest.json`` symlink under ``<output-dir>/context_scaling/``
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from exo_tools.cluster import EcoSession
|
||||
from loguru import logger
|
||||
|
||||
from bench.cli._common import (
|
||||
SharedOptions,
|
||||
add_shared_args,
|
||||
get_arg,
|
||||
get_arg_optional,
|
||||
)
|
||||
from bench.lib import context_scaling
|
||||
from bench.lib.cluster import managed_cluster, managed_instance
|
||||
from bench.lib.context_scaling import (
|
||||
ContextScalingParams,
|
||||
make_cold_control_factory,
|
||||
)
|
||||
from bench.lib.model_meta import (
|
||||
ModelMeta,
|
||||
derive_cold_controls,
|
||||
derive_context_ramp,
|
||||
fetch_model_meta,
|
||||
)
|
||||
from bench.lib.results import ResultsBundle, RunMetadata, find_repo_root
|
||||
|
||||
|
||||
def add_subparser(
|
||||
subparsers: argparse._SubParsersAction[argparse.ArgumentParser], # type: ignore[type-arg]
|
||||
) -> None:
|
||||
parser = subparsers.add_parser(
|
||||
"context-scaling",
|
||||
help="Prompt-TPS / decode-TPS vs context-size sweep",
|
||||
description=__doc__,
|
||||
)
|
||||
add_shared_args(parser)
|
||||
g = parser.add_argument_group("context-scaling")
|
||||
g.add_argument(
|
||||
"--num-steps",
|
||||
type=int,
|
||||
default=32,
|
||||
help="Number of equally-spaced PP points in the ramp (K).",
|
||||
)
|
||||
g.add_argument(
|
||||
"--pp-step",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Δ (token step). If unset, derived from the model's max context.",
|
||||
)
|
||||
g.add_argument(
|
||||
"--fraction-of-max",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="When Δ is auto-derived, use this fraction of the model's "
|
||||
"max_position_embeddings as the ramp's upper bound (0 < f ≤ 1).",
|
||||
)
|
||||
g.add_argument(
|
||||
"--tg",
|
||||
type=int,
|
||||
default=64,
|
||||
help="Tokens to generate per step (decode duration; constant across ramp).",
|
||||
)
|
||||
g.add_argument(
|
||||
"--warmup",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Warmup requests at pp=Δ before the measured ramp. "
|
||||
"First warmup is cache-disabled (kernel JIT only); subsequent "
|
||||
"warmups are cache-enabled (the second is the one that primes "
|
||||
"the cache entry with a hot-kernel rate). Default 2 is the "
|
||||
"sweet spot: warmup=0 leaves JIT cost in step 0; warmup=1 has "
|
||||
"step 0 as a 'none' hit (still hot-kernel cold prefill, just "
|
||||
"classified differently).",
|
||||
)
|
||||
g.add_argument(
|
||||
"--cold-controls",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Cold-control pp values to take after the cached sweep. Either "
|
||||
"'auto' (4 evenly-spaced points across the ramp) or a comma-separated "
|
||||
"list of explicit pp values (e.g. '8192,32768,65536'). "
|
||||
"Default: no cold controls.",
|
||||
)
|
||||
g.add_argument(
|
||||
"--sleep-between-s",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="Seconds to sleep between consecutive sweep requests.",
|
||||
)
|
||||
parser.set_defaults(handler=run)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def run(args: argparse.Namespace) -> Path:
|
||||
"""Execute the context-scaling benchmark per the parsed args.
|
||||
|
||||
Returns the path of the JSON results file.
|
||||
"""
|
||||
shared = SharedOptions.from_namespace(args)
|
||||
repo_root = find_repo_root()
|
||||
|
||||
# 1. Fetch HF metadata up-front; everything else can be derived from it.
|
||||
logger.info(f"fetching HuggingFace metadata for {shared.model}")
|
||||
meta = fetch_model_meta(shared.model)
|
||||
logger.info(
|
||||
f" weights: {meta.total_weight_gb:.1f}GB; "
|
||||
f"max context: {meta.max_position_embeddings} tokens; "
|
||||
f"layers: {meta.num_hidden_layers}"
|
||||
)
|
||||
|
||||
# 2. Derive constraints (user values always win; otherwise fall back to
|
||||
# ModelMeta heuristics for the *minimums*).
|
||||
min_memory_gb = (
|
||||
shared.min_memory_gb
|
||||
if shared.min_memory_gb is not None
|
||||
else meta.memory_constraint_gb
|
||||
)
|
||||
min_disk_gb = (
|
||||
shared.min_disk_gb
|
||||
if shared.min_disk_gb is not None
|
||||
else meta.disk_constraint_gb
|
||||
)
|
||||
logger.info(f" cluster constraint: min memory {min_memory_gb:.1f}GB")
|
||||
logger.info(f" cluster constraint: min disk {min_disk_gb:.1f}GB")
|
||||
if shared.max_memory_gb is not None:
|
||||
logger.info(f" cluster constraint: max memory {shared.max_memory_gb:.1f}GB")
|
||||
if shared.max_disk_gb is not None:
|
||||
logger.info(f" cluster constraint: max disk {shared.max_disk_gb:.1f}GB")
|
||||
|
||||
explicit_pp_step = get_arg_optional(args, "pp_step", int)
|
||||
num_steps = get_arg(args, "num_steps", int)
|
||||
if explicit_pp_step is not None:
|
||||
pp_step = explicit_pp_step
|
||||
else:
|
||||
pp_step, num_steps = derive_context_ramp(
|
||||
meta,
|
||||
num_steps=num_steps,
|
||||
fraction_of_max=get_arg(args, "fraction_of_max", float),
|
||||
)
|
||||
logger.info(
|
||||
f" derived ramp: Δ={pp_step} × K={num_steps} "
|
||||
f"= {pp_step * num_steps} tokens (max {meta.max_position_embeddings})"
|
||||
)
|
||||
|
||||
cold_controls = _resolve_cold_controls(
|
||||
args, meta, pp_step=pp_step, num_steps=num_steps
|
||||
)
|
||||
if cold_controls:
|
||||
logger.info(f" cold controls: {list(cold_controls)}")
|
||||
|
||||
# 3. Spin up cluster + instance + run.
|
||||
eco = EcoSession(user_prefix=shared.user_prefix)
|
||||
output_dir = (shared.output_dir / "context_scaling").resolve()
|
||||
metadata = RunMetadata.new(
|
||||
benchmark="context_scaling",
|
||||
repo_root=repo_root,
|
||||
tags={**shared.tags, "host_pool": ",".join(shared.hosts) or "<auto>"},
|
||||
)
|
||||
bundle = ResultsBundle(metadata=metadata)
|
||||
|
||||
with (
|
||||
managed_cluster(
|
||||
eco,
|
||||
hosts=list(shared.hosts) or None,
|
||||
count=shared.nodes,
|
||||
thunderbolt=shared.thunderbolt,
|
||||
chip=shared.chip,
|
||||
min_memory_gb=min_memory_gb,
|
||||
max_memory_gb=shared.max_memory_gb,
|
||||
min_disk_gb=min_disk_gb,
|
||||
max_disk_gb=shared.max_disk_gb,
|
||||
) as cluster,
|
||||
managed_instance(
|
||||
cluster,
|
||||
eco,
|
||||
shared.model,
|
||||
sharding=shared.sharding,
|
||||
comm=shared.comm,
|
||||
min_nodes=shared.min_nodes,
|
||||
evict_downloads=shared.evict_downloads,
|
||||
cleanup_on_exit=shared.cleanup_instance,
|
||||
) as session,
|
||||
):
|
||||
params = ContextScalingParams(
|
||||
pp_step=pp_step,
|
||||
num_steps=num_steps,
|
||||
tg=get_arg(args, "tg", int),
|
||||
warmup=get_arg(args, "warmup", int),
|
||||
cold_controls=cold_controls,
|
||||
sleep_between_s=get_arg(args, "sleep_between_s", float),
|
||||
)
|
||||
factory = (
|
||||
make_cold_control_factory(
|
||||
session, shared.sharding, shared.comm, shared.min_nodes
|
||||
)
|
||||
if cold_controls
|
||||
else None
|
||||
)
|
||||
context_scaling.run(session, params, bundle, cold_control_factory=factory)
|
||||
|
||||
out_path = bundle.write_json(output_dir)
|
||||
_update_latest_symlink(out_path)
|
||||
logger.info(f"wrote results → {out_path}")
|
||||
|
||||
_validate_partial_hits(bundle)
|
||||
return out_path
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _resolve_cold_controls(
|
||||
args: argparse.Namespace,
|
||||
meta: ModelMeta,
|
||||
*,
|
||||
pp_step: int,
|
||||
num_steps: int,
|
||||
) -> tuple[int, ...]:
|
||||
raw = get_arg_optional(args, "cold_controls", str)
|
||||
if raw is None or not raw.strip():
|
||||
return ()
|
||||
if raw.strip().lower() == "auto":
|
||||
return derive_cold_controls(meta, pp_step=pp_step, num_steps=num_steps, count=4)
|
||||
return tuple(int(s.strip()) for s in raw.split(",") if s.strip())
|
||||
|
||||
|
||||
def _update_latest_symlink(out_path: Path) -> None:
|
||||
"""Update ``<dir>/latest.json`` to point at the newly-written file."""
|
||||
link = out_path.parent / "latest.json"
|
||||
try:
|
||||
if link.is_symlink() or link.exists():
|
||||
link.unlink()
|
||||
os.symlink(out_path.name, link)
|
||||
except OSError as e:
|
||||
logger.warning(f"could not update latest.json symlink: {e}")
|
||||
|
||||
|
||||
def _validate_partial_hits(bundle: ResultsBundle) -> None:
|
||||
"""Hard-fail if the cached sweep didn't see ``partial`` on every step ≥ 1.
|
||||
|
||||
Step 0 is allowed to be ``exact`` (warmup primed the cache at pp=Δ); a
|
||||
later ``exact`` means Δ was effectively absorbed into the cache and the
|
||||
cold-rate measurement is meaningless. ``none`` means the cache was
|
||||
discarded mid-sweep and ``T_cum`` is unreliable.
|
||||
"""
|
||||
cached = [r for r in bundle.runs if r.get("phase") == "cached_sweep"]
|
||||
bad = [r for r in cached[1:] if r.get("prefix_cache_hit") != "partial"]
|
||||
if bad:
|
||||
bad_summary = [(r["step_index"], r["prefix_cache_hit"]) for r in bad]
|
||||
raise RuntimeError(
|
||||
f"{len(bad)} cached-sweep step(s) reported "
|
||||
f"prefix_cache_hit != 'partial': {bad_summary!r}; "
|
||||
"T_cum is unreliable."
|
||||
)
|
||||
@@ -0,0 +1,125 @@
|
||||
"""Plot benchmark results — CLI subcommand.
|
||||
|
||||
uv run python -m bench.cli plot bench/results/context_scaling/latest.json
|
||||
uv run python -m bench.cli plot run_a.json run_b.json --label-tag operator
|
||||
uv run python -m bench.cli plot latest.json --output /tmp/scaling.png
|
||||
|
||||
The benchmark type is detected from each JSON's ``metadata.benchmark`` —
|
||||
all input files must share the same benchmark.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from bench.cli._common import get_arg_optional
|
||||
from bench.lib.plotting import PlotInputs, render_context_scaling
|
||||
|
||||
|
||||
def add_subparser(
|
||||
subparsers: argparse._SubParsersAction[argparse.ArgumentParser], # type: ignore[type-arg]
|
||||
) -> None:
|
||||
parser = subparsers.add_parser(
|
||||
"plot",
|
||||
help="Render benchmark JSON result(s) as a PNG.",
|
||||
description=__doc__,
|
||||
)
|
||||
parser.add_argument(
|
||||
"paths",
|
||||
nargs="+",
|
||||
type=str,
|
||||
help="One or more bench results JSON files. Multiple files are "
|
||||
"rendered as a comparison plot (one line per file).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
type=str,
|
||||
default=None,
|
||||
help="PNG output path. Default: replace the first JSON's '.json' "
|
||||
"suffix with '.png' (or '.compare.png' when multiple inputs).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--label-tag",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Use metadata.tags[<KEY>] as the legend label for each run "
|
||||
"(falls back to run_id if unset or missing).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--title",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Override the auto-generated figure title.",
|
||||
)
|
||||
parser.set_defaults(handler=run)
|
||||
|
||||
|
||||
def run(args: argparse.Namespace) -> Path:
|
||||
paths_raw = getattr(args, "paths", None)
|
||||
if not isinstance(paths_raw, list) or not paths_raw:
|
||||
raise SystemExit("plot: at least one JSON path is required")
|
||||
paths = [
|
||||
Path(str(p)) # type: ignore[reportUnknownArgumentType]
|
||||
for p in cast("list[Any]", paths_raw) # type: ignore[reportAny]
|
||||
]
|
||||
for p in paths:
|
||||
if not p.is_file():
|
||||
raise SystemExit(f"plot: file not found: {p}")
|
||||
|
||||
benchmarks = {_benchmark_for(p) for p in paths}
|
||||
if len(benchmarks) != 1:
|
||||
raise SystemExit(
|
||||
f"plot: all input JSONs must share the same benchmark, got {benchmarks!r}"
|
||||
)
|
||||
benchmark = next(iter(benchmarks))
|
||||
|
||||
output_arg = get_arg_optional(args, "output", str)
|
||||
output = Path(output_arg) if output_arg is not None else _default_output(paths)
|
||||
inputs = PlotInputs(
|
||||
results=paths,
|
||||
output=output,
|
||||
label_tag=get_arg_optional(args, "label_tag", str),
|
||||
title=get_arg_optional(args, "title", str),
|
||||
)
|
||||
|
||||
if benchmark == "context_scaling":
|
||||
out_path = render_context_scaling(inputs)
|
||||
else:
|
||||
raise SystemExit(f"plot: no renderer registered for benchmark {benchmark!r}")
|
||||
|
||||
logger.info(f"plot: wrote {out_path}")
|
||||
return out_path
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _benchmark_for(path: Path) -> str:
|
||||
with path.open() as f:
|
||||
loaded: Any = json.load(f) # type: ignore[reportAny]
|
||||
if not isinstance(loaded, dict):
|
||||
raise SystemExit(f"plot: {path}: expected top-level JSON object")
|
||||
metadata: Any = loaded.get("metadata", {}) # type: ignore[reportAny]
|
||||
if not isinstance(metadata, dict):
|
||||
raise SystemExit(f"plot: {path}: metadata is not an object")
|
||||
benchmark: Any = metadata.get("benchmark") # type: ignore[reportAny]
|
||||
if not isinstance(benchmark, str):
|
||||
raise SystemExit(f"plot: {path}: metadata.benchmark missing or not a string")
|
||||
return benchmark
|
||||
|
||||
|
||||
def _default_output(paths: list[Path]) -> Path:
|
||||
"""Auto-derive a PNG path next to the first JSON.
|
||||
|
||||
Single input → ``<path>.png`` (replaces ``.json``).
|
||||
Multiple inputs → ``<path>.compare.png`` next to the first JSON.
|
||||
"""
|
||||
first = paths[0]
|
||||
if len(paths) == 1:
|
||||
return first.with_suffix(".png")
|
||||
return first.with_name(first.stem + ".compare.png")
|
||||
Whitespace-only changes.
@@ -0,0 +1,109 @@
|
||||
"""Unit tests for ``bench.cli.campaign``.
|
||||
|
||||
The pure helpers (defaults+run merge, table-lookup, str-or-none) are
|
||||
tested here. End-to-end campaign execution requires a real eco cluster
|
||||
and is exercised manually via ``bench campaign <toml>``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from bench.cli.campaign import (
|
||||
_merge, # type: ignore[reportPrivateUsage]
|
||||
_str_or_none, # type: ignore[reportPrivateUsage]
|
||||
_table, # type: ignore[reportPrivateUsage]
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _table
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestTable:
|
||||
def test_present_table(self) -> None:
|
||||
data = {"defaults": {"nodes": 4}}
|
||||
assert _table(data, "defaults") == {"nodes": 4}
|
||||
|
||||
def test_missing_key_returns_empty(self) -> None:
|
||||
assert _table({}, "absent") == {}
|
||||
|
||||
def test_non_table_value_returns_empty(self) -> None:
|
||||
# `nodes = 4` is an int at top level, not a table; treat as empty.
|
||||
assert _table({"nodes": 4}, "nodes") == {}
|
||||
|
||||
def test_list_value_returns_empty(self) -> None:
|
||||
assert _table({"runs": [{"a": 1}]}, "runs") == {}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _merge
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMerge:
|
||||
def test_run_wins_on_conflict(self) -> None:
|
||||
defaults = {"nodes": 4, "tg": 64}
|
||||
run = {"nodes": 2}
|
||||
assert _merge(defaults, run) == {"nodes": 2, "tg": 64}
|
||||
|
||||
def test_disjoint_keys(self) -> None:
|
||||
defaults = {"nodes": 4}
|
||||
run = {"model": "test/foo"}
|
||||
assert _merge(defaults, run) == {"nodes": 4, "model": "test/foo"}
|
||||
|
||||
def test_run_only(self) -> None:
|
||||
assert _merge({}, {"a": 1, "b": 2}) == {"a": 1, "b": 2}
|
||||
|
||||
def test_defaults_only(self) -> None:
|
||||
assert _merge({"a": 1}, {}) == {"a": 1}
|
||||
|
||||
def test_tags_deep_merged_defaults_only(self) -> None:
|
||||
defaults = {"tags": {"operator": "ciaranbor"}}
|
||||
run = {"model": "test/foo"}
|
||||
merged = _merge(defaults, run)
|
||||
assert merged["tags"] == {"operator": "ciaranbor"}
|
||||
|
||||
def test_tags_deep_merged_run_only(self) -> None:
|
||||
defaults = {"nodes": 4}
|
||||
run = {"tags": {"model_short": "llama-3b"}}
|
||||
merged = _merge(defaults, run)
|
||||
assert merged["tags"] == {"model_short": "llama-3b"}
|
||||
|
||||
def test_tags_deep_merged_both(self) -> None:
|
||||
defaults = {"tags": {"operator": "ciaranbor", "campaign": "smoke"}}
|
||||
run = {"tags": {"model_short": "llama-3b"}}
|
||||
merged = _merge(defaults, run)
|
||||
assert merged["tags"] == {
|
||||
"operator": "ciaranbor",
|
||||
"campaign": "smoke",
|
||||
"model_short": "llama-3b",
|
||||
}
|
||||
|
||||
def test_run_tags_override_defaults_tags(self) -> None:
|
||||
defaults = {"tags": {"operator": "ciaranbor"}}
|
||||
run = {"tags": {"operator": "alice"}}
|
||||
merged = _merge(defaults, run)
|
||||
assert merged["tags"] == {"operator": "alice"}
|
||||
|
||||
def test_no_tags_table_means_no_tags_key(self) -> None:
|
||||
# When neither side has tags, we don't synthesise an empty dict.
|
||||
merged = _merge({"nodes": 4}, {"model": "test/foo"})
|
||||
assert "tags" not in merged
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _str_or_none
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestStrOrNone:
|
||||
def test_str_passes_through(self) -> None:
|
||||
assert _str_or_none("hello") == "hello"
|
||||
|
||||
def test_none_returns_none(self) -> None:
|
||||
assert _str_or_none(None) is None
|
||||
|
||||
def test_int_returns_none(self) -> None:
|
||||
assert _str_or_none(42) is None
|
||||
|
||||
def test_list_returns_none(self) -> None:
|
||||
assert _str_or_none(["a", "b"]) is None
|
||||
@@ -0,0 +1,275 @@
|
||||
"""Unit tests for the argparse boundary helpers in ``bench.cli._common``."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from bench.cli._common import (
|
||||
_config_to_argv, # type: ignore[reportPrivateUsage]
|
||||
_parse_csv, # type: ignore[reportPrivateUsage]
|
||||
_parse_tags, # type: ignore[reportPrivateUsage]
|
||||
expand_config_in_argv,
|
||||
get_arg,
|
||||
get_arg_optional,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _parse_csv
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParseCsv:
|
||||
def test_simple_list(self) -> None:
|
||||
assert _parse_csv("a,b,c") == ["a", "b", "c"]
|
||||
|
||||
def test_strips_whitespace(self) -> None:
|
||||
assert _parse_csv(" a , b , c ") == ["a", "b", "c"]
|
||||
|
||||
def test_skips_empty_entries(self) -> None:
|
||||
assert _parse_csv("a,,b,") == ["a", "b"]
|
||||
assert _parse_csv(",,,") == []
|
||||
|
||||
def test_empty_string_returns_empty(self) -> None:
|
||||
assert _parse_csv("") == []
|
||||
|
||||
def test_single_value(self) -> None:
|
||||
assert _parse_csv("only") == ["only"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _parse_tags
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestParseTags:
|
||||
def test_empty_input_returns_empty_dict(self) -> None:
|
||||
assert _parse_tags([]) == {}
|
||||
|
||||
def test_single_tag(self) -> None:
|
||||
assert _parse_tags(["operator=ciaranbor"]) == {"operator": "ciaranbor"}
|
||||
|
||||
def test_multiple_tags(self) -> None:
|
||||
assert _parse_tags(["a=1", "b=2", "c=3"]) == {"a": "1", "b": "2", "c": "3"}
|
||||
|
||||
def test_strips_whitespace_around_key_and_value(self) -> None:
|
||||
assert _parse_tags([" key = value "]) == {"key": "value"}
|
||||
|
||||
def test_value_can_contain_equals(self) -> None:
|
||||
assert _parse_tags(["url=http://example.com/?a=b"]) == {
|
||||
"url": "http://example.com/?a=b"
|
||||
}
|
||||
|
||||
def test_later_duplicate_key_wins(self) -> None:
|
||||
# Standard dict behaviour; explicit so we notice if it changes.
|
||||
assert _parse_tags(["k=v1", "k=v2"]) == {"k": "v2"}
|
||||
|
||||
def test_missing_equals_raises(self) -> None:
|
||||
with pytest.raises(argparse.ArgumentTypeError, match="key=value"):
|
||||
_ = _parse_tags(["malformed"])
|
||||
|
||||
def test_one_malformed_in_list_raises(self) -> None:
|
||||
with pytest.raises(argparse.ArgumentTypeError):
|
||||
_ = _parse_tags(["good=1", "bad", "alsogood=2"])
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_arg / get_arg_optional
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetArg:
|
||||
def test_str_passes_through(self) -> None:
|
||||
ns = argparse.Namespace(name="hello")
|
||||
assert get_arg(ns, "name", str) == "hello"
|
||||
|
||||
def test_int_passes_through(self) -> None:
|
||||
ns = argparse.Namespace(count=42)
|
||||
assert get_arg(ns, "count", int) == 42
|
||||
|
||||
def test_int_coerces_from_string(self) -> None:
|
||||
ns = argparse.Namespace(count="42")
|
||||
assert get_arg(ns, "count", int) == 42
|
||||
|
||||
def test_float_passes_through(self) -> None:
|
||||
ns = argparse.Namespace(rate=3.14)
|
||||
assert get_arg(ns, "rate", float) == 3.14
|
||||
|
||||
def test_float_coerces_from_int(self) -> None:
|
||||
ns = argparse.Namespace(rate=3)
|
||||
assert get_arg(ns, "rate", float) == 3.0
|
||||
|
||||
def test_float_coerces_from_string(self) -> None:
|
||||
ns = argparse.Namespace(rate="3.14")
|
||||
assert get_arg(ns, "rate", float) == 3.14
|
||||
|
||||
def test_bool_passes_through(self) -> None:
|
||||
ns = argparse.Namespace(flag=True)
|
||||
assert get_arg(ns, "flag", bool) is True
|
||||
|
||||
def test_wrong_type_raises(self) -> None:
|
||||
ns = argparse.Namespace(name=42)
|
||||
with pytest.raises(TypeError, match="expected str"):
|
||||
_ = get_arg(ns, "name", str)
|
||||
|
||||
def test_missing_attribute_raises(self) -> None:
|
||||
ns = argparse.Namespace()
|
||||
with pytest.raises(AttributeError):
|
||||
_ = get_arg(ns, "missing", str)
|
||||
|
||||
|
||||
class TestGetArgOptional:
|
||||
def test_missing_returns_none(self) -> None:
|
||||
ns = argparse.Namespace()
|
||||
assert get_arg_optional(ns, "missing", str) is None
|
||||
|
||||
def test_explicit_none_returns_none(self) -> None:
|
||||
ns = argparse.Namespace(value=None)
|
||||
assert get_arg_optional(ns, "value", str) is None
|
||||
|
||||
def test_present_value_returns_typed(self) -> None:
|
||||
ns = argparse.Namespace(value="present")
|
||||
assert get_arg_optional(ns, "value", str) == "present"
|
||||
|
||||
def test_int_coerces_from_string(self) -> None:
|
||||
ns = argparse.Namespace(value="42")
|
||||
assert get_arg_optional(ns, "value", int) == 42
|
||||
|
||||
def test_float_coerces_from_int(self) -> None:
|
||||
ns = argparse.Namespace(value=42)
|
||||
assert get_arg_optional(ns, "value", float) == 42.0
|
||||
|
||||
def test_wrong_type_raises(self) -> None:
|
||||
ns = argparse.Namespace(value=[1, 2, 3])
|
||||
with pytest.raises(TypeError, match="expected str or None"):
|
||||
_ = get_arg_optional(ns, "value", str)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _config_to_argv
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestConfigToArgv:
|
||||
def test_empty(self) -> None:
|
||||
assert _config_to_argv({}) == []
|
||||
|
||||
def test_string_value(self) -> None:
|
||||
assert _config_to_argv({"model": "mlx/foo"}) == ["--model", "mlx/foo"]
|
||||
|
||||
def test_int_and_float_values(self) -> None:
|
||||
out = _config_to_argv({"num_steps": 32, "fraction_of_max": 0.5})
|
||||
assert out == ["--num-steps", "32", "--fraction-of-max", "0.5"]
|
||||
|
||||
def test_underscore_keys_become_hyphenated_flags(self) -> None:
|
||||
out = _config_to_argv({"min_memory_gb": 21.0})
|
||||
assert out == ["--min-memory-gb", "21.0"]
|
||||
|
||||
def test_bool_true_emits_flag(self) -> None:
|
||||
assert _config_to_argv({"auto_constrain": True}) == ["--auto-constrain"]
|
||||
|
||||
def test_bool_false_emits_no_form(self) -> None:
|
||||
assert _config_to_argv({"auto_constrain": False}) == ["--no-auto-constrain"]
|
||||
|
||||
def test_none_value_skipped(self) -> None:
|
||||
assert _config_to_argv({"chip": None, "model": "foo"}) == [
|
||||
"--model",
|
||||
"foo",
|
||||
]
|
||||
|
||||
def test_list_joined_as_csv(self) -> None:
|
||||
out = _config_to_argv({"hosts": ["s4", "s9"], "cold_controls": [1024, 2048]})
|
||||
assert out == [
|
||||
"--hosts",
|
||||
"s4,s9",
|
||||
"--cold-controls",
|
||||
"1024,2048",
|
||||
]
|
||||
|
||||
def test_tags_table_expands_to_repeated_tag_args(self) -> None:
|
||||
out = _config_to_argv({"tags": {"operator": "ciaranbor", "run": "full"}})
|
||||
# Order within a TOML table is preserved by tomllib
|
||||
assert out == [
|
||||
"--tag",
|
||||
"operator=ciaranbor",
|
||||
"--tag",
|
||||
"run=full",
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# expand_config_in_argv
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestExpandConfigInArgv:
|
||||
def test_no_config_flag_passthrough(self) -> None:
|
||||
argv = ["context-scaling", "--model", "foo"]
|
||||
assert expand_config_in_argv(argv) == argv
|
||||
|
||||
def test_config_at_end(self, tmp_path: Path) -> None:
|
||||
cfg = tmp_path / "run.toml"
|
||||
_ = cfg.write_text('model = "from_config"\nnum_steps = 16\n')
|
||||
argv = ["context-scaling", "--config", str(cfg)]
|
||||
# Config flags are inserted right after the subcommand
|
||||
assert expand_config_in_argv(argv) == [
|
||||
"context-scaling",
|
||||
"--model",
|
||||
"from_config",
|
||||
"--num-steps",
|
||||
"16",
|
||||
]
|
||||
|
||||
def test_explicit_cli_overrides_config(self, tmp_path: Path) -> None:
|
||||
cfg = tmp_path / "run.toml"
|
||||
_ = cfg.write_text('model = "from_config"\nnum_steps = 16\n')
|
||||
# User overrides --num-steps explicitly. Argparse takes the last
|
||||
# occurrence for non-append actions, so the user's 32 wins.
|
||||
argv = ["context-scaling", "--config", str(cfg), "--num-steps", "32"]
|
||||
out = expand_config_in_argv(argv)
|
||||
assert out == [
|
||||
"context-scaling",
|
||||
"--model",
|
||||
"from_config",
|
||||
"--num-steps",
|
||||
"16",
|
||||
"--num-steps",
|
||||
"32",
|
||||
]
|
||||
|
||||
def test_missing_path_arg_raises(self) -> None:
|
||||
with pytest.raises(ValueError, match="--config requires a path"):
|
||||
_ = expand_config_in_argv(["context-scaling", "--config"])
|
||||
|
||||
def test_nonexistent_file_raises(self, tmp_path: Path) -> None:
|
||||
with pytest.raises(FileNotFoundError, match="Config file not found"):
|
||||
_ = expand_config_in_argv(
|
||||
["context-scaling", "--config", str(tmp_path / "missing.toml")]
|
||||
)
|
||||
|
||||
def test_bool_false_in_config(self, tmp_path: Path) -> None:
|
||||
cfg = tmp_path / "run.toml"
|
||||
_ = cfg.write_text("auto_constrain = false\n")
|
||||
argv = ["context-scaling", "--config", str(cfg)]
|
||||
assert expand_config_in_argv(argv) == [
|
||||
"context-scaling",
|
||||
"--no-auto-constrain",
|
||||
]
|
||||
|
||||
def test_tags_table(self, tmp_path: Path) -> None:
|
||||
cfg = tmp_path / "run.toml"
|
||||
_ = cfg.write_text(
|
||||
'model = "foo"\n[tags]\noperator = "ciaranbor"\nrun = "full"\n'
|
||||
)
|
||||
argv = ["context-scaling", "--config", str(cfg)]
|
||||
assert expand_config_in_argv(argv) == [
|
||||
"context-scaling",
|
||||
"--model",
|
||||
"foo",
|
||||
"--tag",
|
||||
"operator=ciaranbor",
|
||||
"--tag",
|
||||
"run=full",
|
||||
]
|
||||
@@ -0,0 +1,70 @@
|
||||
# Example context-scaling run configuration.
|
||||
#
|
||||
# Use it like this:
|
||||
#
|
||||
# uv run python -m bench.cli context-scaling --config bench/configs/context_scaling.example.toml
|
||||
#
|
||||
# CLI flags placed after `--config` override individual values.
|
||||
#
|
||||
# All shared and subcommand-specific flags can appear here. Keys mirror the
|
||||
# CLI flag names with hyphens replaced by underscores. Boolean keys map to
|
||||
# `--key` / `--no-key`; lists are joined as CSV; the `[tags]` table maps to
|
||||
# repeated `--tag key=value` flags.
|
||||
#
|
||||
# NOTE: TOML scoping — once a `[table]` header is opened, all subsequent
|
||||
# top-level-looking assignments belong to that table until the next header.
|
||||
# Keep tables (like `[tags]`) at the END of the file.
|
||||
|
||||
# ---- Model + placement ----
|
||||
model = "mlx-community/Qwen3-30B-A3B-4bit"
|
||||
# sharding = "Tensor" # default: "Tensor"; pairs with --comm MlxJaccl + --thunderbolt a2a
|
||||
# comm = "MlxJaccl" # default: "MlxJaccl" (RDMA over Thunderbolt)
|
||||
# min_nodes = 1
|
||||
|
||||
# ---- Cluster ----
|
||||
# Either pin to specific hosts...
|
||||
# hosts = ["s4"]
|
||||
# ...or let eco pick hosts that satisfy the constraints below.
|
||||
# nodes = 1
|
||||
# chip = "M3 Ultra" # eco chip name (case-insensitive substring); comment to allow any
|
||||
# thunderbolt = "a2a" # default: "a2a" (clique, for Tensor+JACCL)
|
||||
# "ring" (cycle; for Pipeline+JACCL)
|
||||
# "none" (exclude TB; pair with sharding=Pipeline + comm=MlxRing for non-TB hosts)
|
||||
|
||||
# Memory + disk minimums are auto-derived from the HF model size
|
||||
# (×1.30 + 1 GiB for memory, ×1.10 + 1 GiB for disk). Set any of these
|
||||
# explicitly to override the auto-derived value.
|
||||
# min_memory_gb = 96.0
|
||||
# max_memory_gb = 256.0 # leave bigger machines free for other workloads
|
||||
# min_disk_gb = 24.0
|
||||
# max_disk_gb = 4000.0
|
||||
|
||||
# ---- Runtime ----
|
||||
# evict_downloads is true by default — frees disk smallest-first to fit
|
||||
# the bench model. Set to false to keep existing downloads.
|
||||
# evict_downloads = false
|
||||
|
||||
# cleanup_instance is true by default — deletes the placed instance on exit.
|
||||
# Set to false to leave it running for debugging.
|
||||
# cleanup_instance = false
|
||||
|
||||
# ---- Output ----
|
||||
output_dir = "bench/results"
|
||||
|
||||
# ---- Context-scaling sweep ----
|
||||
num_steps = 32 # K — number of equally-spaced ramp points
|
||||
# pp_step = 1024 # Δ — explicit override; otherwise auto-derived
|
||||
# fraction_of_max = 1.0 # use this fraction of max_position_embeddings
|
||||
tg = 64 # tokens generated per step
|
||||
# warmup = 2 # default: 2 (1 cache-disabled JIT warmup + 1 cache-priming warmup)
|
||||
# cold_controls = "auto" # 4 evenly-spaced controls across the ramp, or:
|
||||
# cold_controls = "8192,16384,32768,40960" # explicit pp values
|
||||
sleep_between_s = 1.0
|
||||
|
||||
# ---- Tags ----
|
||||
# Survive into metadata.tags in the output JSON; useful for filtering or
|
||||
# grouping runs across SHAs / hosts / configs. `$USER` is NOT expanded
|
||||
# (TOML is literal); pass `--tag operator=$USER` on the CLI for shell expansion.
|
||||
# Must be the LAST table in the file (see TOML scoping note above).
|
||||
[tags]
|
||||
run = "full"
|
||||
@@ -0,0 +1,34 @@
|
||||
# 4-node smoke campaign: two small/medium Llama models, abbreviated ramps,
|
||||
# auto-everything else (TB a2a + tensor + JACCL + auto-derived constraints).
|
||||
#
|
||||
# Run with:
|
||||
# uv run python -m bench.cli campaign bench/configs/llama-family-smoke.toml
|
||||
#
|
||||
# Each [[runs]] gets its own cluster (deploy + bench + teardown). After
|
||||
# both runs finish, a side-by-side comparison plot is written next to the
|
||||
# JSONs.
|
||||
|
||||
[defaults]
|
||||
nodes = 4
|
||||
num_steps = 8
|
||||
fraction_of_max = 0.5
|
||||
|
||||
[defaults.tags]
|
||||
campaign = "llama-family-smoke"
|
||||
|
||||
[[runs]]
|
||||
subcommand = "context-scaling"
|
||||
model = "mlx-community/Llama-3.2-3B-Instruct-4bit"
|
||||
[runs.tags]
|
||||
model_short = "llama-3.2-3b-4bit"
|
||||
|
||||
[[runs]]
|
||||
subcommand = "context-scaling"
|
||||
model = "mlx-community/Meta-Llama-3.1-8B-Instruct-4bit"
|
||||
[runs.tags]
|
||||
model_short = "llama-3.1-8b-4bit"
|
||||
|
||||
# Final comparison plot (one PNG per benchmark group with ≥2 runs).
|
||||
[plot]
|
||||
label_tag = "model_short"
|
||||
title = "Llama 3 family — 4-node tensor + JACCL smoke"
|
||||
+34
-270
@@ -24,9 +24,7 @@ import json
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from statistics import mean
|
||||
from typing import Any
|
||||
|
||||
@@ -45,125 +43,44 @@ from exo_tools.harness import (
|
||||
wait_for_instance_ready,
|
||||
)
|
||||
from loguru import logger
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
# Monkey-patch for transformers 5.x compatibility
|
||||
# Kimi's tokenization_kimi.py imports bytes_to_unicode from the old location
|
||||
# which was moved in transformers 5.0.0rc2
|
||||
try:
|
||||
import transformers.models.gpt2.tokenization_gpt2 as gpt2_tokenization
|
||||
from transformers.convert_slow_tokenizer import bytes_to_unicode
|
||||
# PromptSizer / run_one_completion / load_tokenizer_for_bench are the
|
||||
# canonical, fully-typed implementations under bench/lib/. They are
|
||||
# re-exported here for backwards compatibility with prefill_decode_bench.py
|
||||
# and any other consumers of `from exo_bench import …`.
|
||||
from bench.lib.completion import run_one_completion as _lib_run_one_completion
|
||||
from bench.lib.prompt import (
|
||||
PromptSizer as _LibPromptSizer,
|
||||
)
|
||||
from bench.lib.prompt import (
|
||||
load_tokenizer_for_bench as _lib_load_tokenizer_for_bench,
|
||||
)
|
||||
|
||||
if not hasattr(gpt2_tokenization, "bytes_to_unicode"):
|
||||
gpt2_tokenization.bytes_to_unicode = bytes_to_unicode # type: ignore[attr-defined]
|
||||
except ImportError:
|
||||
pass # transformers < 5.0 or bytes_to_unicode not available
|
||||
PromptSizer = _LibPromptSizer
|
||||
load_tokenizer_for_bench = _lib_load_tokenizer_for_bench
|
||||
|
||||
|
||||
def load_tokenizer_for_bench(model_id: str) -> Any:
|
||||
"""
|
||||
Load tokenizer for benchmarking, with special handling for Kimi models.
|
||||
|
||||
Kimi uses a custom TikTokenTokenizer that transformers 5.x can't load via AutoTokenizer.
|
||||
This function replicates the logic from utils_mlx.py for bench compatibility.
|
||||
"""
|
||||
model_id_lower = model_id.lower()
|
||||
|
||||
if "kimi-k2" in model_id_lower:
|
||||
import importlib.util
|
||||
import types
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
# Download/get the model path
|
||||
model_path = Path(
|
||||
snapshot_download(
|
||||
model_id,
|
||||
allow_patterns=["*.json", "*.py", "*.tiktoken", "*.model", "*.jinja"],
|
||||
)
|
||||
)
|
||||
|
||||
sys.path.insert(0, str(model_path))
|
||||
|
||||
# Load tool_declaration_ts first (tokenization_kimi imports it with relative import)
|
||||
tool_decl_path = model_path / "tool_declaration_ts.py"
|
||||
if tool_decl_path.exists():
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"tool_declaration_ts", tool_decl_path
|
||||
)
|
||||
if spec and spec.loader:
|
||||
tool_decl_module = importlib.util.module_from_spec(spec)
|
||||
sys.modules["tool_declaration_ts"] = tool_decl_module
|
||||
spec.loader.exec_module(tool_decl_module)
|
||||
|
||||
# Load tokenization_kimi with patched source (convert relative to absolute import)
|
||||
tok_path = model_path / "tokenization_kimi.py"
|
||||
source = tok_path.read_text()
|
||||
source = source.replace("from .tool_declaration_ts", "from tool_declaration_ts")
|
||||
spec = importlib.util.spec_from_file_location("tokenization_kimi", tok_path)
|
||||
if spec:
|
||||
tok_module = types.ModuleType("tokenization_kimi")
|
||||
tok_module.__file__ = str(tok_path)
|
||||
sys.modules["tokenization_kimi"] = tok_module
|
||||
exec(compile(source, tok_path, "exec"), tok_module.__dict__) # noqa: S102
|
||||
TikTokenTokenizer = tok_module.TikTokenTokenizer # noqa: N806
|
||||
else:
|
||||
from tokenization_kimi import TikTokenTokenizer # type: ignore[import-not-found] # noqa: I001
|
||||
|
||||
hf_tokenizer: Any = TikTokenTokenizer.from_pretrained(model_path)
|
||||
|
||||
# Patch encode to use internal tiktoken model directly
|
||||
# transformers 5.x has a bug in the encode->pad path for slow tokenizers
|
||||
def _patched_encode(text: str, **kwargs: object) -> list[int]:
|
||||
# Pass allowed_special="all" to handle special tokens like <|im_user|>
|
||||
return list(hf_tokenizer.model.encode(text, allowed_special="all"))
|
||||
|
||||
hf_tokenizer.encode = _patched_encode
|
||||
|
||||
return hf_tokenizer
|
||||
|
||||
# TODO: Change back to using only transformers
|
||||
try:
|
||||
return AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
|
||||
except (AttributeError, ValueError):
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
model_path = Path(
|
||||
snapshot_download(
|
||||
model_id,
|
||||
allow_patterns=[
|
||||
"*.json",
|
||||
"*.py",
|
||||
"tokenizer.model",
|
||||
"*.tiktoken",
|
||||
"tiktoken.model",
|
||||
"*.txt",
|
||||
"*.jsonl",
|
||||
"*.jinja",
|
||||
],
|
||||
)
|
||||
)
|
||||
stub_kwargs: dict[str, Any] = {}
|
||||
config_file = model_path / "config.json"
|
||||
if config_file.exists():
|
||||
with open(config_file) as f:
|
||||
raw = json.load(f)
|
||||
for key in (
|
||||
"model_type",
|
||||
"max_position_embeddings",
|
||||
"vocab_size",
|
||||
"bos_token_id",
|
||||
"eos_token_id",
|
||||
"pad_token_id",
|
||||
):
|
||||
if key in raw:
|
||||
stub_kwargs[key] = raw[key]
|
||||
return AutoTokenizer.from_pretrained(
|
||||
str(model_path),
|
||||
config=PretrainedConfig(**stub_kwargs),
|
||||
trust_remote_code=True,
|
||||
)
|
||||
def run_one_completion(
|
||||
client: ExoClient,
|
||||
model_id: str,
|
||||
pp_hint: int,
|
||||
tg: int,
|
||||
prompt_sizer: PromptSizer,
|
||||
*,
|
||||
use_prefix_cache: bool = False,
|
||||
stream: bool = False,
|
||||
) -> tuple[dict[str, Any], int]:
|
||||
"""Backwards-compatible shim returning a plain ``dict`` row."""
|
||||
row, pp_tokens = _lib_run_one_completion(
|
||||
client,
|
||||
model_id,
|
||||
pp_hint,
|
||||
tg,
|
||||
prompt_sizer,
|
||||
use_prefix_cache=use_prefix_cache,
|
||||
stream=stream,
|
||||
)
|
||||
return dict(row), pp_tokens
|
||||
|
||||
|
||||
def format_peak_memory(b: float) -> str:
|
||||
@@ -269,159 +186,6 @@ def parse_int_list(values: list[str]) -> list[int]:
|
||||
return items
|
||||
|
||||
|
||||
def run_one_completion(
|
||||
client: ExoClient,
|
||||
model_id: str,
|
||||
pp_hint: int,
|
||||
tg: int,
|
||||
prompt_sizer: PromptSizer,
|
||||
*,
|
||||
use_prefix_cache: bool = False,
|
||||
stream: bool = False,
|
||||
) -> tuple[dict[str, Any], int]:
|
||||
content, pp_tokens = prompt_sizer.build(pp_hint)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model_id,
|
||||
"messages": [{"role": "user", "content": content}],
|
||||
"max_tokens": tg,
|
||||
"logprobs": False,
|
||||
"use_prefix_cache": use_prefix_cache,
|
||||
}
|
||||
|
||||
if not stream:
|
||||
payload["stream"] = False
|
||||
t0 = time.perf_counter()
|
||||
out = client.post_bench_chat_completions(payload)
|
||||
elapsed = time.perf_counter() - t0
|
||||
|
||||
stats = out.get("generation_stats")
|
||||
choices = out.get("choices") or [{}]
|
||||
message = choices[0].get("message", {}) if choices else {}
|
||||
content = message.get("content") or ""
|
||||
preview = content[:200] if content else ""
|
||||
else:
|
||||
tokens = 0
|
||||
first_token_time = None
|
||||
t0 = time.perf_counter()
|
||||
text_parts: list[str] = []
|
||||
stats = None
|
||||
|
||||
for raw_line in client.stream_bench_chat_completions(payload):
|
||||
line = raw_line.strip()
|
||||
if line.startswith(": generation_stats "):
|
||||
with contextlib.suppress(json.JSONDecodeError):
|
||||
stats = json.loads(line[len(": generation_stats ") :])
|
||||
continue
|
||||
if not line.startswith("data: "):
|
||||
continue
|
||||
data = line[6:]
|
||||
if data == "[DONE]":
|
||||
break
|
||||
try:
|
||||
chunk = json.loads(data)
|
||||
delta = chunk.get("choices", [{}])[0].get("delta", {})
|
||||
if delta.get("content"):
|
||||
if first_token_time is None:
|
||||
first_token_time = time.perf_counter()
|
||||
tokens += 1
|
||||
text_parts.append(delta["content"])
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
elapsed = time.perf_counter() - t0
|
||||
preview = "".join(text_parts)[:200]
|
||||
|
||||
if not stats:
|
||||
ttft = (first_token_time - t0) if first_token_time else elapsed
|
||||
gen_time = elapsed - ttft if tokens > 1 else elapsed
|
||||
gen_tps = (tokens - 1) / gen_time if tokens > 1 and gen_time > 0 else 0.0
|
||||
prompt_tps = pp_tokens / ttft if ttft > 0 else 0.0
|
||||
stats = {
|
||||
"prompt_tokens": pp_tokens,
|
||||
"generation_tokens": tokens,
|
||||
"prompt_tps": round(prompt_tps, 2),
|
||||
"generation_tps": round(gen_tps, 2),
|
||||
"peak_memory_usage": {"inBytes": 0},
|
||||
}
|
||||
|
||||
return {
|
||||
"elapsed_s": elapsed,
|
||||
"output_text_preview": preview,
|
||||
"stats": stats,
|
||||
}, pp_tokens
|
||||
|
||||
|
||||
class PromptSizer:
|
||||
def __init__(self, tokenizer: Any, atom: str = "a "):
|
||||
self.tokenizer = tokenizer
|
||||
self.atom = atom
|
||||
self.count_fn = PromptSizer._make_counter(tokenizer)
|
||||
self.base_tokens = self.count_fn("")
|
||||
|
||||
@staticmethod
|
||||
def _make_counter(tokenizer: Any) -> Callable[[str], int]:
|
||||
def count_fn(user_content: str) -> int:
|
||||
messages = [{"role": "user", "content": user_content}]
|
||||
try:
|
||||
ids = tokenizer.apply_chat_template(
|
||||
messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
except ValueError:
|
||||
# Models without a Jinja chat template (e.g. DeepSeek V4 which
|
||||
# ships its own Python encoder). Use the exo-side V4 encoder.
|
||||
from exo.worker.engines.mlx.deepseek_v4_encoding import (
|
||||
encode_messages as encode_v4,
|
||||
)
|
||||
|
||||
prompt = encode_v4(messages, thinking_mode="thinking")
|
||||
ids = tokenizer.encode(prompt, add_special_tokens=False)
|
||||
# Fix for transformers 5.x
|
||||
if hasattr(ids, "input_ids"):
|
||||
ids = ids.input_ids
|
||||
return int(len(ids))
|
||||
|
||||
return count_fn
|
||||
|
||||
def build(self, target_prompt_tokens: int) -> tuple[str, int]:
|
||||
target = int(target_prompt_tokens)
|
||||
if target < self.base_tokens:
|
||||
raise RuntimeError(
|
||||
f"Target ({target}) is smaller than template overhead ({self.base_tokens})."
|
||||
)
|
||||
|
||||
# Estimate tokens per atom using a sample
|
||||
sample_count = 100
|
||||
sample_content = self.atom * sample_count
|
||||
sample_tokens = self.count_fn(sample_content) - self.base_tokens
|
||||
tokens_per_atom = sample_tokens / sample_count
|
||||
|
||||
# Estimate starting point
|
||||
needed_tokens = target - self.base_tokens
|
||||
estimated_atoms = int(needed_tokens / tokens_per_atom)
|
||||
|
||||
# Binary search to find exact atom count
|
||||
low, high = 0, estimated_atoms * 2 + 100
|
||||
while low < high:
|
||||
mid = (low + high) // 2
|
||||
tok = self.count_fn(self.atom * mid)
|
||||
if tok < target:
|
||||
low = mid + 1
|
||||
else:
|
||||
high = mid
|
||||
|
||||
content = self.atom * low
|
||||
tok = self.count_fn(content)
|
||||
logger.info(f"{tok=}")
|
||||
|
||||
if tok != target:
|
||||
raise RuntimeError(
|
||||
f"Overshot: got {tok} tokens (target {target}). "
|
||||
f"Pick a different atom (try ' a' or '\\n' or '0 ')."
|
||||
)
|
||||
|
||||
return content, tok
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser(
|
||||
prog="exo-bench",
|
||||
|
||||
@@ -0,0 +1,18 @@
|
||||
"""Composable bench library for exo.
|
||||
|
||||
Provides reusable building blocks for benchmarks:
|
||||
|
||||
- :class:`bench.lib.session.BenchSession` — cluster + instance + client wrapper
|
||||
- :class:`bench.lib.results.ResultsBundle` — structured results + JSON writer
|
||||
- :func:`bench.lib.cluster.managed_cluster` /
|
||||
:func:`bench.lib.cluster.managed_instance` — eco-managed lifecycle ctx-managers
|
||||
- :func:`bench.lib.model_meta.fetch_model_meta` — HF metadata fetcher driving
|
||||
cluster constraints + auto-derived context ramps
|
||||
- :mod:`bench.lib.context_scaling` — prompt-TPS / decode-TPS vs context-size sweep
|
||||
|
||||
CLI entrypoints under ``bench/cli/`` consume this library via
|
||||
``python -m bench.cli <subcommand>``. Adding a new benchmark = (i) write
|
||||
``bench/lib/<name>.py`` exposing a typed ``run(session, params, bundle)``
|
||||
callable, (ii) write ``bench/cli/<name>.py`` with an ``add_subparser`` and
|
||||
a handler, (iii) register it in ``_REGISTRY`` in ``bench/cli/__main__.py``.
|
||||
"""
|
||||
@@ -0,0 +1,215 @@
|
||||
"""Eco-managed cluster + instance lifecycle helpers for the bench CLI.
|
||||
|
||||
Two context managers:
|
||||
|
||||
- :func:`managed_cluster` deploys exo on the requested hosts (or via
|
||||
constraint-based reservation) and tears it down on exit.
|
||||
- :func:`managed_instance` resolves the model on the cluster, optionally
|
||||
frees disk via ``--danger-delete-downloads`` (default on for benches),
|
||||
places the instance, and deletes it on exit.
|
||||
|
||||
The library never reaches for global state — every call takes an
|
||||
explicit :class:`EcoSession`. Callers are expected to instantiate one
|
||||
session per CLI invocation and use it across both context managers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, cast
|
||||
|
||||
from exo_tools.client import ExoClient
|
||||
from exo_tools.cluster import Chip, ClusterInfo, EcoSession, Thunderbolt
|
||||
from exo_tools.harness import (
|
||||
Comm,
|
||||
Sharding,
|
||||
cleanup_all_instances,
|
||||
place_instance,
|
||||
resolve_model_short_id,
|
||||
run_planning_phase,
|
||||
)
|
||||
from loguru import logger
|
||||
|
||||
from .session import BenchSession
|
||||
|
||||
|
||||
@contextmanager
|
||||
def managed_cluster(
|
||||
eco: EcoSession,
|
||||
*,
|
||||
hosts: list[str] | None = None,
|
||||
count: int = 1,
|
||||
thunderbolt: Thunderbolt | None = None,
|
||||
chip: Chip | None = None,
|
||||
min_memory_gb: float | None = None,
|
||||
max_memory_gb: float | None = None,
|
||||
min_disk_gb: float | None = None,
|
||||
max_disk_gb: float | None = None,
|
||||
deploy_timeout_s: int = 600,
|
||||
) -> Iterator[ClusterInfo]:
|
||||
"""Deploy exo for the duration of the ``with`` block, then ``eco stop``.
|
||||
|
||||
If ``hosts`` is given, deploys on exactly those hosts (constraint flags
|
||||
are ignored — eco doesn't re-validate the explicit list). Otherwise eco
|
||||
reserves any matching hosts that satisfy all of:
|
||||
|
||||
- ``count`` (number of hosts)
|
||||
- ``thunderbolt`` topology (``A2A``, ``RING``, or ``NONE`` to
|
||||
exclude TB-connected hosts)
|
||||
- ``chip`` (substring match against eco's chip names)
|
||||
- memory bounds (``min_memory_gb`` / ``max_memory_gb``)
|
||||
- disk bounds (``min_disk_gb`` / ``max_disk_gb``)
|
||||
"""
|
||||
if hosts:
|
||||
cluster = eco.start_deploy(
|
||||
hosts=hosts[:count],
|
||||
wait=True,
|
||||
timeout=deploy_timeout_s,
|
||||
)
|
||||
else:
|
||||
cluster = eco.start_deploy(
|
||||
count=count,
|
||||
thunderbolt=thunderbolt,
|
||||
chip=chip,
|
||||
min_memory_gb=min_memory_gb,
|
||||
max_memory_gb=max_memory_gb,
|
||||
min_disk_gb=min_disk_gb,
|
||||
max_disk_gb=max_disk_gb,
|
||||
wait=True,
|
||||
timeout=deploy_timeout_s,
|
||||
)
|
||||
logger.info(
|
||||
f"cluster deployed: {len(cluster.hosts)} host(s) "
|
||||
f"({', '.join(cluster.hosts)}); namespace={cluster.namespace}"
|
||||
)
|
||||
try:
|
||||
yield cluster
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
eco.stop(cluster.hosts)
|
||||
logger.info("cluster stopped")
|
||||
|
||||
|
||||
@contextmanager
|
||||
def managed_instance(
|
||||
cluster: ClusterInfo,
|
||||
eco: EcoSession,
|
||||
model_id: str,
|
||||
*,
|
||||
sharding: Sharding = Sharding.PIPELINE,
|
||||
comm: Comm = Comm.RING,
|
||||
min_nodes: int = 1,
|
||||
evict_downloads: bool = True,
|
||||
cleanup_on_exit: bool = True,
|
||||
instance_timeout_s: float = 7200.0,
|
||||
settle_timeout_s: float = 60.0,
|
||||
) -> Iterator[BenchSession]:
|
||||
"""Resolve the model on the cluster, place an instance, yield a session.
|
||||
|
||||
Steps on entry:
|
||||
1. Resolve ``model_id`` to ``(short_id, full_id)`` against the cluster's
|
||||
``/models`` endpoint (auto-adds from HuggingFace if missing).
|
||||
2. Run the harness's planning phase: validates each node has enough
|
||||
disk for the model and starts the download (or reuses an existing
|
||||
download). When ``evict_downloads=True`` (the default for benches),
|
||||
this also evicts smaller existing models if disk is short.
|
||||
3. Place the instance, wait for it to be ``RunnerReady``.
|
||||
4. Yield a :class:`BenchSession` pointing at the cluster's primary API.
|
||||
|
||||
On exit: deletes the placed instance (and any other lingering
|
||||
instances) so the cluster is clean for the next benchmark.
|
||||
"""
|
||||
client = cluster.make_client(timeout_s=instance_timeout_s)
|
||||
|
||||
short_id, full_id = resolve_model_short_id(client, model_id, force_download=True)
|
||||
logger.info(f"resolved model: short_id={short_id} full_id={full_id}")
|
||||
|
||||
# The planning phase needs a concrete preview (instance + runner-to-shard
|
||||
# mapping) to know which nodes to download to. Pull the placements API
|
||||
# directly and take the first valid one — bench cares about disk +
|
||||
# download, not the specific shard mapping.
|
||||
preview = _first_valid_preview(client, full_id, settle_timeout_s)
|
||||
if preview is None:
|
||||
raise RuntimeError(
|
||||
f"No placement available for {full_id} on cluster {cluster.hosts}"
|
||||
)
|
||||
|
||||
duration = run_planning_phase(
|
||||
client,
|
||||
full_id,
|
||||
preview,
|
||||
danger_delete=evict_downloads,
|
||||
timeout=instance_timeout_s,
|
||||
settle_deadline=None,
|
||||
)
|
||||
if duration is not None:
|
||||
logger.info(f"download: {duration:.1f}s (freshly downloaded)")
|
||||
else:
|
||||
logger.info("download: model already cached on all nodes")
|
||||
|
||||
instance_id = place_instance(
|
||||
client,
|
||||
model_id,
|
||||
sharding=sharding,
|
||||
comm=comm,
|
||||
min_nodes=min_nodes,
|
||||
timeout=instance_timeout_s,
|
||||
)
|
||||
logger.info(f"placed instance {instance_id} ({sharding.value}/{comm.value})")
|
||||
|
||||
sess = BenchSession(
|
||||
cluster=cluster,
|
||||
eco=eco,
|
||||
instance_id=instance_id,
|
||||
model_id=short_id,
|
||||
full_model_id=full_id,
|
||||
)
|
||||
try:
|
||||
yield sess
|
||||
finally:
|
||||
if cleanup_on_exit:
|
||||
with contextlib.suppress(Exception):
|
||||
cleanup_all_instances(sess.client)
|
||||
else:
|
||||
logger.info(
|
||||
f"cleanup_on_exit=False: leaving instance(s) on {cluster.hosts}"
|
||||
)
|
||||
|
||||
|
||||
def _first_valid_preview(
|
||||
client: ExoClient, full_model_id: str, settle_timeout_s: float
|
||||
) -> dict[str, Any] | None:
|
||||
"""Poll ``/instance/previews`` until at least one valid preview comes back."""
|
||||
deadline = time.monotonic() + settle_timeout_s
|
||||
backoff_s = 1.0
|
||||
while True:
|
||||
resp_obj: Any = client.request_json( # type: ignore[reportAny]
|
||||
"GET", "/instance/previews", params={"model_id": full_model_id}
|
||||
)
|
||||
resp: dict[str, Any] = (
|
||||
cast("dict[str, Any]", resp_obj) if isinstance(resp_obj, dict) else {}
|
||||
)
|
||||
previews_raw: object = resp.get("previews") or []
|
||||
previews: list[Any] = (
|
||||
cast("list[Any]", previews_raw) if isinstance(previews_raw, list) else []
|
||||
)
|
||||
for raw in previews: # type: ignore[reportAny]
|
||||
if not isinstance(raw, dict):
|
||||
continue
|
||||
entry = cast("dict[str, Any]", raw)
|
||||
if entry.get("error") is not None:
|
||||
continue
|
||||
instance = entry.get("instance")
|
||||
if isinstance(instance, dict):
|
||||
return entry
|
||||
if time.monotonic() >= deadline:
|
||||
return None
|
||||
logger.info(
|
||||
f"waiting for placement to appear for {full_model_id} "
|
||||
f"({deadline - time.monotonic():.0f}s remaining)..."
|
||||
)
|
||||
time.sleep(min(backoff_s, max(0.0, deadline - time.monotonic())))
|
||||
backoff_s = min(backoff_s * 2, 30.0)
|
||||
@@ -0,0 +1,194 @@
|
||||
"""Typed wrapper around ``/bench/chat/completions`` for benchmarks.
|
||||
|
||||
The bench endpoint disables EOS suppression and KV prefix caching by
|
||||
default (see ``bench/METHODOLOGY.md``). This module exposes a single
|
||||
function :func:`run_one_completion` that:
|
||||
|
||||
1. Builds an exact-token-length prompt via :class:`PromptSizer`.
|
||||
2. POSTs to ``/bench/chat/completions``.
|
||||
3. Returns a ``(BenchRow, prompt_tokens)`` pair where ``BenchRow`` is a
|
||||
:class:`typing.TypedDict` with the fields the caller needs.
|
||||
|
||||
Streaming is supported but rarely needed for context-scaling — the
|
||||
non-streaming path is the default.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import time
|
||||
from typing import Any, Literal, NotRequired, TypedDict, cast
|
||||
|
||||
from exo_tools.client import ExoClient
|
||||
|
||||
from .prompt import PromptSizer
|
||||
|
||||
PrefixCacheHit = Literal["none", "partial", "exact"]
|
||||
|
||||
|
||||
class GenerationStats(TypedDict, total=False):
|
||||
"""Server-reported per-task timing stats."""
|
||||
|
||||
prompt_tps: float
|
||||
generation_tps: float
|
||||
prompt_tokens: int
|
||||
generation_tokens: int
|
||||
peak_memory_usage: dict[str, int]
|
||||
prefix_cache_hit: PrefixCacheHit
|
||||
|
||||
|
||||
class BenchRow(TypedDict):
|
||||
"""Per-request result row returned to callers."""
|
||||
|
||||
elapsed_s: float
|
||||
output_text_preview: str
|
||||
stats: GenerationStats
|
||||
error: NotRequired[str]
|
||||
|
||||
|
||||
def _as_dict(value: Any) -> dict[str, Any]: # type: ignore[reportAny]
|
||||
"""Narrow an arbitrary JSON value to a typed ``dict[str, Any]``."""
|
||||
if isinstance(value, dict):
|
||||
return cast("dict[str, Any]", value)
|
||||
return {}
|
||||
|
||||
|
||||
def _as_list(value: Any) -> list[Any]: # type: ignore[reportAny]
|
||||
if isinstance(value, list):
|
||||
return cast("list[Any]", value)
|
||||
return []
|
||||
|
||||
|
||||
def _extract_stats(raw_response: dict[str, Any]) -> GenerationStats:
|
||||
stats_obj = raw_response.get("generation_stats")
|
||||
if not isinstance(stats_obj, dict):
|
||||
return {}
|
||||
return cast("GenerationStats", cast("object", stats_obj))
|
||||
|
||||
|
||||
def _extract_preview(raw_response: dict[str, Any], limit: int = 200) -> str:
|
||||
choices = _as_list(raw_response.get("choices"))
|
||||
if not choices:
|
||||
return ""
|
||||
first = _as_dict(choices[0])
|
||||
message = _as_dict(first.get("message"))
|
||||
content_obj = message.get("content")
|
||||
if isinstance(content_obj, str):
|
||||
return content_obj[:limit]
|
||||
return ""
|
||||
|
||||
|
||||
def run_one_completion(
|
||||
client: ExoClient,
|
||||
model_id: str,
|
||||
pp_hint: int,
|
||||
tg: int,
|
||||
prompt_sizer: PromptSizer,
|
||||
*,
|
||||
use_prefix_cache: bool = False,
|
||||
stream: bool = False,
|
||||
) -> tuple[BenchRow, int]:
|
||||
"""Send one request to ``/bench/chat/completions`` and return its row.
|
||||
|
||||
``pp_hint`` is the *target* prompt-token count; the actual prompt is
|
||||
sized via :class:`PromptSizer` and the verified value is returned as
|
||||
the second element of the tuple.
|
||||
"""
|
||||
content, pp_tokens = prompt_sizer.build(pp_hint)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model_id,
|
||||
"messages": [{"role": "user", "content": content}],
|
||||
"max_tokens": tg,
|
||||
"logprobs": False,
|
||||
"use_prefix_cache": use_prefix_cache,
|
||||
}
|
||||
|
||||
if not stream:
|
||||
payload["stream"] = False
|
||||
t0 = time.perf_counter()
|
||||
raw_obj = client.post_bench_chat_completions(payload)
|
||||
elapsed = time.perf_counter() - t0
|
||||
raw = _as_dict(raw_obj)
|
||||
return (
|
||||
BenchRow(
|
||||
elapsed_s=elapsed,
|
||||
output_text_preview=_extract_preview(raw),
|
||||
stats=_extract_stats(raw),
|
||||
),
|
||||
pp_tokens,
|
||||
)
|
||||
|
||||
return _run_streaming(client, payload, pp_tokens)
|
||||
|
||||
|
||||
def _run_streaming(
|
||||
client: ExoClient,
|
||||
payload: dict[str, Any],
|
||||
pp_tokens: int,
|
||||
) -> tuple[BenchRow, int]:
|
||||
"""Streaming variant: parse SSE lines, recover ``GenerationStats``."""
|
||||
payload = {**payload, "stream": True}
|
||||
|
||||
tokens = 0
|
||||
first_token_time: float | None = None
|
||||
t0 = time.perf_counter()
|
||||
text_parts: list[str] = []
|
||||
stats: GenerationStats = {}
|
||||
|
||||
for raw_line in client.stream_bench_chat_completions(payload):
|
||||
line = raw_line.strip()
|
||||
if line.startswith(": generation_stats "):
|
||||
with contextlib.suppress(json.JSONDecodeError):
|
||||
parsed_obj: Any = json.loads( # type: ignore[reportAny]
|
||||
line[len(": generation_stats ") :]
|
||||
)
|
||||
if isinstance(parsed_obj, dict):
|
||||
stats = cast("GenerationStats", cast("object", parsed_obj))
|
||||
continue
|
||||
if not line.startswith("data: "):
|
||||
continue
|
||||
data = line[6:]
|
||||
if data == "[DONE]":
|
||||
break
|
||||
try:
|
||||
chunk_obj: Any = json.loads(data) # type: ignore[reportAny]
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
chunk = _as_dict(chunk_obj)
|
||||
choices = _as_list(chunk.get("choices"))
|
||||
if not choices:
|
||||
continue
|
||||
first = _as_dict(choices[0])
|
||||
delta = _as_dict(first.get("delta"))
|
||||
delta_content_obj = delta.get("content")
|
||||
if isinstance(delta_content_obj, str) and delta_content_obj:
|
||||
if first_token_time is None:
|
||||
first_token_time = time.perf_counter()
|
||||
tokens += 1
|
||||
text_parts.append(delta_content_obj)
|
||||
|
||||
elapsed = time.perf_counter() - t0
|
||||
preview = "".join(text_parts)[:200]
|
||||
|
||||
if not stats:
|
||||
ttft = (first_token_time - t0) if first_token_time is not None else elapsed
|
||||
gen_time = elapsed - ttft if tokens > 1 else elapsed
|
||||
gen_tps = (tokens - 1) / gen_time if tokens > 1 and gen_time > 0 else 0.0
|
||||
prompt_tps = pp_tokens / ttft if ttft > 0 else 0.0
|
||||
stats = GenerationStats(
|
||||
prompt_tokens=pp_tokens,
|
||||
generation_tokens=tokens,
|
||||
prompt_tps=round(prompt_tps, 2),
|
||||
generation_tps=round(gen_tps, 2),
|
||||
peak_memory_usage={"inBytes": 0},
|
||||
)
|
||||
|
||||
return (
|
||||
BenchRow(
|
||||
elapsed_s=elapsed,
|
||||
output_text_preview=preview,
|
||||
stats=stats,
|
||||
),
|
||||
pp_tokens,
|
||||
)
|
||||
@@ -0,0 +1,428 @@
|
||||
"""Prompt-TPS / decode-TPS vs context-size sweep.
|
||||
|
||||
Methodology (see also ``bench/METHODOLOGY.md``):
|
||||
|
||||
Run a single ascending ramp of equally-spaced prompt lengths
|
||||
``pp ∈ {Δ, 2Δ, …, K·Δ}`` with ``prefix_cache=enabled``, ``repeat=1``,
|
||||
``concurrency=1`` and one warmup at ``pp=Δ``.
|
||||
|
||||
Because each step's prefix is exactly what the previous step left in
|
||||
the cache, every step beyond the first is a *partial* hit and the
|
||||
server-reported ``prompt_tps`` reflects the true cold rate over the
|
||||
fresh ``Δ``-token suffix. We accept the warmup's reported rate as the
|
||||
cold equivalent for ``pp=Δ`` (the warmup itself is the cold prefill).
|
||||
|
||||
``decode TPS`` is independent of prefill mechanics — every step's
|
||||
``generation_tps`` is a real decode-rate-at-N data point.
|
||||
|
||||
Cumulative cold-prefill upper bound:
|
||||
``T_cum(pp_k) = Σ_{i=1..k} (Δ_i / prompt_tps_i)``
|
||||
|
||||
Optional cold-control points (``prefix_cache=disabled``) validate the
|
||||
approximation; the gap quantifies per-task overhead. To preserve the
|
||||
``none`` cache-hit classification AND ensure the request actually
|
||||
hits a freshly-placed runner (the master picks the instance with the
|
||||
lowest in-flight task count, which is non-deterministic when multiple
|
||||
same-model instances exist), :func:`run` deletes the sweep instance
|
||||
*before* invoking the cold-control factory. The factory itself places
|
||||
a fresh instance per control and deletes it on exit; the
|
||||
:func:`bench.lib.cluster.managed_instance` ctx-manager calls
|
||||
``cleanup_all_instances`` on exit as a final safety net.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import time
|
||||
from collections.abc import Callable, Iterator
|
||||
from contextlib import AbstractContextManager, contextmanager
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Any
|
||||
|
||||
from exo_tools.client import ExoClient, ExoHttpError
|
||||
from exo_tools.harness import (
|
||||
Comm,
|
||||
Sharding,
|
||||
place_instance,
|
||||
wait_for_instance_gone,
|
||||
)
|
||||
from loguru import logger
|
||||
|
||||
from .completion import GenerationStats, PrefixCacheHit, run_one_completion
|
||||
from .prompt import PromptSizer
|
||||
from .results import ResultsBundle
|
||||
from .session import BenchSession
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ContextScalingParams:
|
||||
"""Inputs for a single context-scaling sweep."""
|
||||
|
||||
pp_step: int
|
||||
num_steps: int
|
||||
tg: int
|
||||
warmup: int = 1
|
||||
cold_controls: tuple[int, ...] = ()
|
||||
sleep_between_s: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class StepResult:
|
||||
pp_tokens: int
|
||||
delta_tokens: int
|
||||
prompt_tps: float
|
||||
generation_tps: float
|
||||
prefix_cache_hit: PrefixCacheHit | str
|
||||
prompt_tokens: int
|
||||
generation_tokens: int
|
||||
elapsed_s: float
|
||||
peak_memory_bytes: int = 0
|
||||
output_text_preview: str = ""
|
||||
|
||||
|
||||
def _peak_bytes(stats: GenerationStats) -> int:
|
||||
pm = stats.get("peak_memory_usage") or {}
|
||||
return int(pm.get("inBytes") or pm.get("in_bytes") or 0)
|
||||
|
||||
|
||||
def _build_step_result(
|
||||
pp_tokens: int,
|
||||
delta_tokens: int,
|
||||
elapsed_s: float,
|
||||
output_text_preview: str,
|
||||
stats: GenerationStats,
|
||||
) -> StepResult:
|
||||
return StepResult(
|
||||
pp_tokens=pp_tokens,
|
||||
delta_tokens=delta_tokens,
|
||||
prompt_tps=float(stats.get("prompt_tps") or 0.0),
|
||||
generation_tps=float(stats.get("generation_tps") or 0.0),
|
||||
prefix_cache_hit=stats.get("prefix_cache_hit") or "unknown",
|
||||
prompt_tokens=int(stats.get("prompt_tokens") or pp_tokens),
|
||||
generation_tokens=int(stats.get("generation_tokens") or 0),
|
||||
elapsed_s=elapsed_s,
|
||||
peak_memory_bytes=_peak_bytes(stats),
|
||||
output_text_preview=output_text_preview[:200],
|
||||
)
|
||||
|
||||
|
||||
def _run_request(
|
||||
client: ExoClient,
|
||||
full_model_id: str,
|
||||
pp: int,
|
||||
tg: int,
|
||||
sizer: PromptSizer,
|
||||
*,
|
||||
use_prefix_cache: bool,
|
||||
) -> tuple[StepResult, int]:
|
||||
"""Send one request and return ``(StepResult, actual_pp_tokens)``."""
|
||||
row, actual_pp = run_one_completion(
|
||||
client,
|
||||
full_model_id,
|
||||
pp,
|
||||
tg,
|
||||
sizer,
|
||||
use_prefix_cache=use_prefix_cache,
|
||||
stream=False,
|
||||
)
|
||||
step = _build_step_result(
|
||||
pp_tokens=actual_pp,
|
||||
delta_tokens=actual_pp, # caller overrides for cached sweep
|
||||
elapsed_s=row["elapsed_s"],
|
||||
output_text_preview=row["output_text_preview"],
|
||||
stats=row["stats"],
|
||||
)
|
||||
return step, actual_pp
|
||||
|
||||
|
||||
def _compute_t_cum(steps: list[StepResult]) -> list[float]:
|
||||
t_cum = 0.0
|
||||
out: list[float] = []
|
||||
for s in steps:
|
||||
if s.prompt_tps > 0 and s.delta_tokens > 0:
|
||||
t_cum += s.delta_tokens / s.prompt_tps
|
||||
out.append(round(t_cum, 6))
|
||||
return out
|
||||
|
||||
|
||||
def run_cached_sweep(
|
||||
session: BenchSession,
|
||||
params: ContextScalingParams,
|
||||
bundle: ResultsBundle,
|
||||
) -> list[StepResult]:
|
||||
"""Run the ascending PP sweep with ``prefix_cache=enabled``.
|
||||
|
||||
Mutates ``bundle.runs`` in place and returns the typed step list.
|
||||
"""
|
||||
if session.full_model_id is None:
|
||||
raise RuntimeError(
|
||||
"BenchSession.full_model_id must be set for context-scaling."
|
||||
)
|
||||
|
||||
sizer = session.get_prompt_sizer()
|
||||
client = session.client
|
||||
pp_targets = [params.pp_step * i for i in range(1, params.num_steps + 1)]
|
||||
logger.info(
|
||||
f"context-scaling: K={params.num_steps} steps, Δ={params.pp_step} tokens, "
|
||||
f"tg={params.tg}, warmup={params.warmup}, cached"
|
||||
)
|
||||
|
||||
# Warmup discipline:
|
||||
# - First warmup runs with the prefix cache DISABLED. This triggers
|
||||
# the MLX kernel JIT compile + KV-buffer alloc for this exact
|
||||
# (Δ, dtype, batch) shape, but does NOT write a cache entry — so
|
||||
# the cold-with-JIT rate isn't fossilised.
|
||||
# - Subsequent warmups run with the prefix cache ENABLED. The
|
||||
# second one finds an empty cache, does a real cold prefill with
|
||||
# a HOT kernel, and writes the resulting rate into the cache
|
||||
# entry at pp=Δ.
|
||||
# - Step 0 (also cache-enabled) is then an exact hit on that entry
|
||||
# and reports the hot rate.
|
||||
# Default warmup=2 gives both effects; warmup=1 still does the JIT
|
||||
# warmup but leaves step 0 as a "none" hit (cold prefill at the hot
|
||||
# kernel, creates the cache entry on the way through).
|
||||
for w in range(params.warmup):
|
||||
is_jit_warmup = w == 0
|
||||
kind = "JIT warmup" if is_jit_warmup else "cache-prime warmup"
|
||||
logger.info(
|
||||
f" warmup {w + 1}/{params.warmup} ({kind}, pp={params.pp_step})"
|
||||
)
|
||||
_run_request(
|
||||
client,
|
||||
session.full_model_id,
|
||||
params.pp_step,
|
||||
params.tg,
|
||||
sizer,
|
||||
use_prefix_cache=not is_jit_warmup,
|
||||
)
|
||||
|
||||
steps: list[StepResult] = []
|
||||
prev_pp = 0
|
||||
for i, pp in enumerate(pp_targets):
|
||||
time.sleep(params.sleep_between_s)
|
||||
try:
|
||||
step, actual_pp = _run_request(
|
||||
client,
|
||||
session.full_model_id,
|
||||
pp,
|
||||
params.tg,
|
||||
sizer,
|
||||
use_prefix_cache=True,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"step {i + 1}/{params.num_steps} (pp={pp}) failed: {e}")
|
||||
raise
|
||||
|
||||
step.delta_tokens = actual_pp - prev_pp
|
||||
steps.append(step)
|
||||
bundle.runs.append({"step_index": i, "phase": "cached_sweep", **asdict(step)})
|
||||
logger.info(
|
||||
f" step {i + 1}/{params.num_steps} pp={actual_pp} Δ={step.delta_tokens} "
|
||||
f"prompt_tps={step.prompt_tps:.1f} gen_tps={step.generation_tps:.2f} "
|
||||
f"hit={step.prefix_cache_hit}"
|
||||
)
|
||||
prev_pp = actual_pp
|
||||
|
||||
return steps
|
||||
|
||||
|
||||
def run_cold_controls(
|
||||
factory: Callable[[], AbstractContextManager[ExoClient]],
|
||||
session: BenchSession,
|
||||
params: ContextScalingParams,
|
||||
bundle: ResultsBundle,
|
||||
) -> list[StepResult]:
|
||||
"""Run cold-control points on a fresh instance to preserve ``none`` hits.
|
||||
|
||||
A cold control is a single request at ``pp=N`` with
|
||||
``prefix_cache=disabled``, executed against a freshly-placed instance
|
||||
(and with no other same-model instance live, so the master's task
|
||||
routing is deterministic). The caller is expected to delete the
|
||||
sweep instance before invoking this — see :func:`run`.
|
||||
"""
|
||||
if not params.cold_controls:
|
||||
return []
|
||||
if session.full_model_id is None:
|
||||
raise RuntimeError("BenchSession.full_model_id must be set for cold controls.")
|
||||
|
||||
sizer = session.get_prompt_sizer()
|
||||
out: list[StepResult] = []
|
||||
for control_pp in params.cold_controls:
|
||||
logger.info(f"cold control: pp={control_pp} (fresh instance, cache disabled)")
|
||||
with factory() as fresh_client:
|
||||
step, actual_pp = _run_request(
|
||||
fresh_client,
|
||||
session.full_model_id,
|
||||
control_pp,
|
||||
params.tg,
|
||||
sizer,
|
||||
use_prefix_cache=False,
|
||||
)
|
||||
step.delta_tokens = actual_pp
|
||||
out.append(step)
|
||||
bundle.cold_controls.append({"phase": "cold_control", **asdict(step)})
|
||||
logger.info(
|
||||
f" cold pp={actual_pp} prompt_tps={step.prompt_tps:.1f} "
|
||||
f"gen_tps={step.generation_tps:.2f} hit={step.prefix_cache_hit}"
|
||||
)
|
||||
if step.prefix_cache_hit != "none":
|
||||
logger.warning(
|
||||
f"cold control at pp={actual_pp} reported "
|
||||
f"prefix_cache_hit={step.prefix_cache_hit!r}; "
|
||||
f"control may not be cold."
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def derive_summary(
|
||||
steps: list[StepResult],
|
||||
cold_controls: list[StepResult],
|
||||
) -> dict[str, Any]:
|
||||
"""Compute the cumulative cold-prefill upper bound + control gaps."""
|
||||
t_cum = _compute_t_cum(steps)
|
||||
bracketed = sorted(
|
||||
((s.pp_tokens, t) for s, t in zip(steps, t_cum, strict=True)),
|
||||
key=lambda x: x[0],
|
||||
)
|
||||
control_gaps: list[dict[str, float]] = []
|
||||
for ctrl in cold_controls:
|
||||
cold_t = ctrl.pp_tokens / ctrl.prompt_tps if ctrl.prompt_tps > 0 else 0.0
|
||||
cum_t = _interp(bracketed, ctrl.pp_tokens)
|
||||
gap = cum_t - cold_t
|
||||
control_gaps.append(
|
||||
{
|
||||
"pp_tokens": ctrl.pp_tokens,
|
||||
"cold_t_seconds": round(cold_t, 4),
|
||||
"t_cum_seconds_at_pp": round(cum_t, 4),
|
||||
"gap_seconds": round(gap, 4),
|
||||
"gap_fraction": round(gap / cold_t, 4) if cold_t > 0 else 0.0,
|
||||
}
|
||||
)
|
||||
return {
|
||||
"t_cum_seconds": t_cum,
|
||||
"control_gaps": control_gaps,
|
||||
}
|
||||
|
||||
|
||||
def _interp(points: list[tuple[int, float]], x: int) -> float:
|
||||
"""Linear interpolate y at x, given sorted ``(x, y)`` points."""
|
||||
if not points:
|
||||
return 0.0
|
||||
if x <= points[0][0]:
|
||||
return points[0][1]
|
||||
if x >= points[-1][0]:
|
||||
return points[-1][1]
|
||||
for i in range(1, len(points)):
|
||||
x0, y0 = points[i - 1]
|
||||
x1, y1 = points[i]
|
||||
if x0 <= x <= x1 and x1 != x0:
|
||||
return y0 + (y1 - y0) * (x - x0) / (x1 - x0)
|
||||
return points[-1][1]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cold-control instance factory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def make_cold_control_factory(
|
||||
session: BenchSession,
|
||||
sharding: Sharding,
|
||||
comm: Comm,
|
||||
min_nodes: int,
|
||||
instance_timeout_s: float = 1800.0,
|
||||
) -> Callable[[], AbstractContextManager[ExoClient]]:
|
||||
"""Return a callable yielding a context manager that places a fresh instance.
|
||||
|
||||
Each ``with factory() as client:`` block places a brand-new instance,
|
||||
yields its client, then deletes the instance on exit. Used to isolate
|
||||
cold-control runs.
|
||||
|
||||
The caller is responsible for ensuring no other same-model instance is
|
||||
live during the ``with`` block — otherwise master routing is
|
||||
non-deterministic and the cold control may be served by a stale runner.
|
||||
See :func:`run` for the orchestration.
|
||||
"""
|
||||
|
||||
@contextmanager
|
||||
def factory() -> Iterator[ExoClient]:
|
||||
if session.full_model_id is None:
|
||||
raise RuntimeError("session.full_model_id is unset")
|
||||
client = session.client
|
||||
instance_id = place_instance(
|
||||
client,
|
||||
session.full_model_id,
|
||||
sharding=sharding,
|
||||
comm=comm,
|
||||
min_nodes=min_nodes,
|
||||
timeout=instance_timeout_s,
|
||||
)
|
||||
try:
|
||||
yield client
|
||||
finally:
|
||||
with contextlib.suppress(ExoHttpError):
|
||||
client.request_json("DELETE", f"/instance/{instance_id}")
|
||||
with contextlib.suppress(Exception):
|
||||
wait_for_instance_gone(client, instance_id, timeout=60.0)
|
||||
|
||||
return factory
|
||||
|
||||
|
||||
def _delete_instance(client: ExoClient, instance_id: str) -> None:
|
||||
"""Best-effort delete of a placed instance."""
|
||||
with contextlib.suppress(ExoHttpError):
|
||||
client.request_json("DELETE", f"/instance/{instance_id}")
|
||||
with contextlib.suppress(Exception):
|
||||
wait_for_instance_gone(client, instance_id, timeout=60.0)
|
||||
|
||||
|
||||
def run(
|
||||
session: BenchSession,
|
||||
params: ContextScalingParams,
|
||||
bundle: ResultsBundle,
|
||||
*,
|
||||
cold_control_factory: Callable[[], AbstractContextManager[ExoClient]] | None = None,
|
||||
) -> ResultsBundle:
|
||||
"""End-to-end: cached sweep + optional cold controls + derived summary.
|
||||
|
||||
To make cold controls truly isolated from the sweep instance, we
|
||||
delete the sweep instance *before* running the controls (otherwise
|
||||
the master might route a control's request to the stale sweep
|
||||
instance, since both match the same ``model_id``). The controls then
|
||||
each place their own fresh instance via the factory.
|
||||
"""
|
||||
bundle.params.update(
|
||||
{
|
||||
"pp_step": params.pp_step,
|
||||
"num_steps": params.num_steps,
|
||||
"tg": params.tg,
|
||||
"warmup": params.warmup,
|
||||
"cold_controls": list(params.cold_controls),
|
||||
"sleep_between_s": params.sleep_between_s,
|
||||
"model_id": session.model_id,
|
||||
"full_model_id": session.full_model_id,
|
||||
}
|
||||
)
|
||||
bundle.capture_cluster(session.client)
|
||||
|
||||
cached_steps = run_cached_sweep(session, params, bundle)
|
||||
|
||||
cold_steps: list[StepResult] = []
|
||||
if params.cold_controls and cold_control_factory is not None:
|
||||
# Delete the sweep instance so the cold-control fresh instance is
|
||||
# the only same-model instance live for the duration of the controls.
|
||||
if session.instance_id is not None:
|
||||
logger.info(
|
||||
f"cold controls: deleting sweep instance {session.instance_id} "
|
||||
"to isolate fresh instance routing"
|
||||
)
|
||||
_delete_instance(session.client, session.instance_id)
|
||||
session.instance_id = None
|
||||
cold_steps = run_cold_controls(cold_control_factory, session, params, bundle)
|
||||
elif params.cold_controls and cold_control_factory is None:
|
||||
logger.warning(
|
||||
"Cold controls requested but no cold_control_factory supplied; skipping."
|
||||
)
|
||||
|
||||
bundle.derived.update(derive_summary(cached_steps, cold_steps))
|
||||
return bundle
|
||||
@@ -0,0 +1,183 @@
|
||||
"""Fetch HuggingFace model metadata for benchmark planning.
|
||||
|
||||
Two pieces of metadata drive every benchmark we run:
|
||||
|
||||
1. **Total weight size** — used to derive ``min-memory`` and ``min-disk``
|
||||
constraints when picking a host. We sum the sizes of all
|
||||
``.safetensors`` (or ``.bin``) shards from the repo's file listing.
|
||||
2. **Max position embeddings** — the model's training context length.
|
||||
Used to bound a context-scaling sweep at the model's max context, and
|
||||
to derive a sensible Δ given a target step count.
|
||||
|
||||
The fetcher uses the ``huggingface_hub`` python API, which talks to the
|
||||
public HF Hub HTTPS endpoints — no exo cluster required, no download
|
||||
of weights.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, cast
|
||||
|
||||
# Files that count toward the on-disk weight footprint.
|
||||
_WEIGHT_SUFFIXES = (".safetensors", ".bin", ".gguf", ".pt", ".npz")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelMeta:
|
||||
"""Subset of HF metadata that a benchmark needs."""
|
||||
|
||||
model_id: str
|
||||
total_weight_bytes: int
|
||||
max_position_embeddings: int
|
||||
num_hidden_layers: int
|
||||
raw_config: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@property
|
||||
def total_weight_gb(self) -> float:
|
||||
return self.total_weight_bytes / (1024**3)
|
||||
|
||||
@property
|
||||
def memory_constraint_gb(self) -> float:
|
||||
"""Estimated minimum host memory to hold weights + overhead.
|
||||
|
||||
Picks the model size + 30 % headroom (KV cache, activations,
|
||||
framework bookkeeping). Rounded up to the next whole GiB.
|
||||
"""
|
||||
return float(int(self.total_weight_gb * 1.30) + 1)
|
||||
|
||||
@property
|
||||
def disk_constraint_gb(self) -> float:
|
||||
"""Disk space the host must have free for the download."""
|
||||
return float(int(self.total_weight_gb * 1.10) + 1)
|
||||
|
||||
|
||||
def _read_config_json(model_id: str) -> dict[str, Any]:
|
||||
from huggingface_hub import (
|
||||
hf_hub_download, # type: ignore[reportUnknownVariableType]
|
||||
)
|
||||
|
||||
raw_path = hf_hub_download(repo_id=model_id, filename="config.json", dry_run=False)
|
||||
with open(raw_path) as f:
|
||||
loaded: Any = json.load(f) # type: ignore[reportAny]
|
||||
return cast("dict[str, Any]", loaded) if isinstance(loaded, dict) else {}
|
||||
|
||||
|
||||
def _sum_weight_sizes(model_id: str) -> int:
|
||||
"""Sum sizes of all weight-shard files in the repo's file listing."""
|
||||
from huggingface_hub import HfApi
|
||||
|
||||
api = HfApi()
|
||||
info = api.model_info(repo_id=model_id, files_metadata=True)
|
||||
siblings = info.siblings or []
|
||||
total = 0
|
||||
for sib in siblings:
|
||||
rfilename = getattr(sib, "rfilename", None)
|
||||
size = getattr(sib, "size", None)
|
||||
if not isinstance(rfilename, str) or not isinstance(size, int):
|
||||
continue
|
||||
if any(rfilename.endswith(suf) for suf in _WEIGHT_SUFFIXES):
|
||||
total += size
|
||||
return total
|
||||
|
||||
|
||||
def _first_int(config: dict[str, Any], *keys: str) -> int:
|
||||
"""Return the first key from ``config`` that holds a usable positive int."""
|
||||
for key in keys:
|
||||
value = config.get(key)
|
||||
if isinstance(value, int) and value > 0:
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
parsed = int(value)
|
||||
except ValueError:
|
||||
continue
|
||||
if parsed > 0:
|
||||
return parsed
|
||||
return 0
|
||||
|
||||
|
||||
def fetch_model_meta(model_id: str) -> ModelMeta:
|
||||
"""Fetch the metadata our benchmarks care about for ``model_id``.
|
||||
|
||||
Args:
|
||||
model_id: HuggingFace repo id, e.g. ``mlx-community/Qwen3-30B-A3B-4bit``.
|
||||
|
||||
Returns:
|
||||
Populated :class:`ModelMeta`.
|
||||
|
||||
Raises:
|
||||
Exception: any HTTP / parse error from ``huggingface_hub`` propagates.
|
||||
"""
|
||||
config = _read_config_json(model_id)
|
||||
return ModelMeta(
|
||||
model_id=model_id,
|
||||
total_weight_bytes=_sum_weight_sizes(model_id),
|
||||
max_position_embeddings=_first_int(
|
||||
config,
|
||||
"max_position_embeddings",
|
||||
"max_seq_len",
|
||||
"model_max_length",
|
||||
"n_positions",
|
||||
),
|
||||
num_hidden_layers=_first_int(
|
||||
config,
|
||||
"num_hidden_layers",
|
||||
"num_layers",
|
||||
"n_layer",
|
||||
"n_layers",
|
||||
"num_decoder_layers",
|
||||
),
|
||||
raw_config=config,
|
||||
)
|
||||
|
||||
|
||||
def derive_context_ramp(
|
||||
meta: ModelMeta,
|
||||
*,
|
||||
num_steps: int,
|
||||
fraction_of_max: float = 1.0,
|
||||
min_pp_step: int = 256,
|
||||
round_to: int = 256,
|
||||
) -> tuple[int, int]:
|
||||
"""Pick ``(pp_step, num_steps)`` covering ``fraction_of_max`` of the context.
|
||||
|
||||
Δ is rounded down to the nearest ``round_to`` so the per-step prompt is a
|
||||
clean number, and clamped to ``min_pp_step`` for tiny-context models.
|
||||
"""
|
||||
if meta.max_position_embeddings <= 0:
|
||||
raise ValueError(
|
||||
f"{meta.model_id} reports max_position_embeddings=0 in config.json"
|
||||
)
|
||||
if not (0.0 < fraction_of_max <= 1.0):
|
||||
raise ValueError(f"fraction_of_max must be in (0, 1], got {fraction_of_max}")
|
||||
if num_steps <= 0:
|
||||
raise ValueError(f"num_steps must be >0, got {num_steps}")
|
||||
|
||||
target_max = int(meta.max_position_embeddings * fraction_of_max)
|
||||
raw_step = max(min_pp_step, target_max // num_steps)
|
||||
pp_step = (raw_step // round_to) * round_to or round_to
|
||||
return pp_step, num_steps
|
||||
|
||||
|
||||
def derive_cold_controls(
|
||||
meta: ModelMeta,
|
||||
*,
|
||||
pp_step: int,
|
||||
num_steps: int,
|
||||
count: int = 4,
|
||||
) -> tuple[int, ...]:
|
||||
"""Pick ``count`` evenly-spaced cold-control points across the ramp.
|
||||
|
||||
Always includes the largest ramp point (``pp_step * num_steps``).
|
||||
Returns control pp values in ascending order, deduped.
|
||||
"""
|
||||
if count <= 0:
|
||||
return ()
|
||||
max_pp = pp_step * num_steps
|
||||
if count == 1:
|
||||
return (max_pp,)
|
||||
spaced = sorted({(max_pp * (i + 1)) // count for i in range(count)})
|
||||
# Filter out anything below pp_step (a control at <Δ is meaningless).
|
||||
return tuple(p for p in spaced if p >= pp_step)
|
||||
@@ -0,0 +1,308 @@
|
||||
"""Typed matplotlib renderers for benchmark JSON results.
|
||||
|
||||
This module owns the *visualisation* of bench results, mirroring how
|
||||
``bench/lib/<name>.py`` owns the methodology and ``bench/cli/<name>.py``
|
||||
owns the orchestration. Adding plotting for a new benchmark = a new
|
||||
``render_<name>`` function here + a dispatch entry in ``bench/cli/plot.py``.
|
||||
|
||||
Functions take typed inputs (``Path`` lists, options) and write a PNG.
|
||||
They never touch argparse or stdout — that's the CLI's job.
|
||||
|
||||
matplotlib's type stubs are thin (most return values are ``Any``), so all
|
||||
calls into ``pyplot`` are concentrated at the bottom of this file with
|
||||
targeted ``# type: ignore[reportUnknownMemberType, reportAny]`` per line.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
# Tab10 cycle from matplotlib's default; we pick colours by index ourselves
|
||||
# instead of fishing them out of `Line2D.get_color()` so the strict-type
|
||||
# fallout stays small and predictable.
|
||||
_COLOR_CYCLE: tuple[str, ...] = (
|
||||
"C0",
|
||||
"C1",
|
||||
"C2",
|
||||
"C3",
|
||||
"C4",
|
||||
"C5",
|
||||
"C6",
|
||||
"C7",
|
||||
"C8",
|
||||
"C9",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class PlotInputs:
|
||||
"""Inputs for any benchmark renderer.
|
||||
|
||||
Attributes:
|
||||
results: One or more bench JSON files. The first is used to
|
||||
auto-derive the title when ``title`` is unset.
|
||||
output: Path to write the PNG to.
|
||||
label_tag: When set, use ``metadata.tags[label_tag]`` as the
|
||||
legend label for each run; otherwise use the run id.
|
||||
title: Override for the figure title.
|
||||
"""
|
||||
|
||||
results: list[Path]
|
||||
output: Path
|
||||
label_tag: str | None = None
|
||||
title: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _RunSeries:
|
||||
"""Pre-extracted plot data for one results JSON.
|
||||
|
||||
``cached_prefill_seconds`` is the cumulative cold-prefill estimate
|
||||
(``T_cum`` from the methodology — read from ``derived.t_cum_seconds``).
|
||||
``control_prefill_seconds`` is the actual cold prefill time per
|
||||
control (``pp_tokens / prompt_tps`` from the cold-control row).
|
||||
"""
|
||||
|
||||
label: str
|
||||
cached_pp: list[int]
|
||||
cached_prefill_seconds: list[float]
|
||||
cached_gen_tps: list[float]
|
||||
control_pp: list[int]
|
||||
control_prefill_seconds: list[float]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pure data extraction (strict-typed, no matplotlib)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _load(path: Path) -> dict[str, Any]:
|
||||
"""Read a bench JSON file and assert top-level shape."""
|
||||
with path.open() as f:
|
||||
loaded: Any = json.load(f) # type: ignore[reportAny]
|
||||
if not isinstance(loaded, dict):
|
||||
raise ValueError(f"{path}: expected top-level JSON object")
|
||||
return cast("dict[str, Any]", loaded)
|
||||
|
||||
|
||||
# dict[str, Any].get(...) returns Any. The five _get_* helpers below
|
||||
# concentrate the Any boundary so the rest of the module can be strict.
|
||||
|
||||
|
||||
def _get_dict(d: dict[str, Any], key: str) -> dict[str, Any]:
|
||||
val: Any = d.get(key)
|
||||
return cast("dict[str, Any]", val) if isinstance(val, dict) else {}
|
||||
|
||||
|
||||
def _get_list(d: dict[str, Any], key: str) -> list[Any]:
|
||||
val: Any = d.get(key)
|
||||
return cast("list[Any]", val) if isinstance(val, list) else []
|
||||
|
||||
|
||||
def _get_str(d: dict[str, Any], key: str, default: str = "") -> str:
|
||||
val: Any = d.get(key, default) # type: ignore[reportAny]
|
||||
return val if isinstance(val, str) else default
|
||||
|
||||
|
||||
def _get_int(row: dict[str, Any], key: str) -> int:
|
||||
val: Any = row.get(key, 0) # type: ignore[reportAny]
|
||||
if isinstance(val, bool): # bool is int; reject explicitly
|
||||
return 0
|
||||
if isinstance(val, (int, float)):
|
||||
return int(val)
|
||||
if isinstance(val, str):
|
||||
try:
|
||||
return int(float(val))
|
||||
except ValueError:
|
||||
return 0
|
||||
return 0
|
||||
|
||||
|
||||
def _get_float(row: dict[str, Any], key: str) -> float:
|
||||
val: Any = row.get(key, 0.0) # type: ignore[reportAny]
|
||||
if isinstance(val, bool):
|
||||
return 0.0
|
||||
if isinstance(val, (int, float)):
|
||||
return float(val)
|
||||
if isinstance(val, str):
|
||||
try:
|
||||
return float(val)
|
||||
except ValueError:
|
||||
return 0.0
|
||||
return 0.0
|
||||
|
||||
|
||||
def _label_for(data: dict[str, Any], label_tag: str | None) -> str:
|
||||
if label_tag is not None:
|
||||
tags = _get_dict(_get_dict(data, "metadata"), "tags")
|
||||
if label_tag in tags:
|
||||
return _get_str(tags, label_tag, "(unnamed)")
|
||||
return _get_str(_get_dict(data, "metadata"), "run_id", "(unnamed)")
|
||||
|
||||
|
||||
def _extract_series(data: dict[str, Any], label: str) -> _RunSeries:
|
||||
"""Pre-extract typed lists from a context-scaling bench JSON."""
|
||||
cached_pp: list[int] = []
|
||||
cached_gen_tps: list[float] = []
|
||||
for raw in _get_list(data, "runs"): # type: ignore[reportAny]
|
||||
if not isinstance(raw, dict):
|
||||
continue
|
||||
row = cast("dict[str, Any]", raw)
|
||||
if _get_str(row, "phase") != "cached_sweep":
|
||||
continue
|
||||
cached_pp.append(_get_int(row, "pp_tokens"))
|
||||
cached_gen_tps.append(_get_float(row, "generation_tps"))
|
||||
|
||||
# Cumulative cold-prefill estimate is computed in derive_summary and
|
||||
# written to derived.t_cum_seconds (parallel to the cached steps).
|
||||
derived = _get_dict(data, "derived")
|
||||
t_cum_raw = _get_list(derived, "t_cum_seconds")
|
||||
cached_prefill_seconds: list[float] = []
|
||||
for raw in t_cum_raw: # type: ignore[reportAny]
|
||||
if isinstance(raw, (int, float)) and not isinstance(raw, bool):
|
||||
cached_prefill_seconds.append(float(raw))
|
||||
|
||||
# Cold controls give us the actual cold prefill time directly:
|
||||
# pp_tokens / prompt_tps. Skip rows with zero/missing prompt_tps.
|
||||
control_pp: list[int] = []
|
||||
control_prefill_seconds: list[float] = []
|
||||
for raw in _get_list(data, "cold_controls"): # type: ignore[reportAny]
|
||||
if not isinstance(raw, dict):
|
||||
continue
|
||||
row = cast("dict[str, Any]", raw)
|
||||
pp = _get_int(row, "pp_tokens")
|
||||
tps = _get_float(row, "prompt_tps")
|
||||
if pp > 0 and tps > 0:
|
||||
control_pp.append(pp)
|
||||
control_prefill_seconds.append(pp / tps)
|
||||
|
||||
return _RunSeries(
|
||||
label=label,
|
||||
cached_pp=cached_pp,
|
||||
cached_prefill_seconds=cached_prefill_seconds,
|
||||
cached_gen_tps=cached_gen_tps,
|
||||
control_pp=control_pp,
|
||||
control_prefill_seconds=control_prefill_seconds,
|
||||
)
|
||||
|
||||
|
||||
def _auto_title(data: dict[str, Any]) -> str:
|
||||
metadata = _get_dict(data, "metadata")
|
||||
params = _get_dict(data, "params")
|
||||
model = (
|
||||
_get_str(params, "full_model_id") or _get_str(params, "model_id") or "(unknown)"
|
||||
)
|
||||
sha = _get_str(metadata, "exo_sha") or "(no-sha)"
|
||||
host = _get_str(metadata, "hostname") or "(no-host)"
|
||||
return f"{model}\n{sha} on {host}"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Matplotlib boundary — each call site has a narrow, justified ignore.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def render_context_scaling(inputs: PlotInputs) -> Path:
|
||||
"""Render a 2-panel context-scaling plot.
|
||||
|
||||
Top: pp_tokens vs prompt_tps (line per run; cold controls as 'x' scatter)
|
||||
Bottom: pp_tokens vs generation_tps (line per run)
|
||||
|
||||
Each line is a separate result file. Multi-file mode is for comparing
|
||||
runs across exo SHAs / hosts / configs; the title is taken from the
|
||||
first file's metadata unless ``inputs.title`` is set.
|
||||
"""
|
||||
if not inputs.results:
|
||||
raise ValueError("at least one results JSON path is required")
|
||||
|
||||
# Validate + extract first so any data-shape error surfaces before we
|
||||
# even import matplotlib.
|
||||
first_data: dict[str, Any] | None = None
|
||||
series: list[_RunSeries] = []
|
||||
for path in inputs.results:
|
||||
data = _load(path)
|
||||
if first_data is None:
|
||||
first_data = data
|
||||
benchmark = _get_str(_get_dict(data, "metadata"), "benchmark")
|
||||
if benchmark != "context_scaling":
|
||||
raise ValueError(
|
||||
f"{path}: expected benchmark=='context_scaling', got {benchmark!r}"
|
||||
)
|
||||
series.append(_extract_series(data, _label_for(data, inputs.label_tag)))
|
||||
|
||||
title = inputs.title
|
||||
if title is None and first_data is not None:
|
||||
title = _auto_title(first_data)
|
||||
if len(inputs.results) > 1:
|
||||
title = f"{title}\n(comparison of {len(inputs.results)} runs)"
|
||||
|
||||
inputs.output.parent.mkdir(parents=True, exist_ok=True)
|
||||
_draw(series, inputs.output, title=title)
|
||||
return inputs.output
|
||||
|
||||
|
||||
def _draw(series: list[_RunSeries], output: Path, *, title: str | None) -> None:
|
||||
"""Concentrated matplotlib boundary."""
|
||||
import matplotlib
|
||||
|
||||
matplotlib.use("Agg")
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
fig, axes = plt.subplots( # type: ignore[reportUnknownMemberType]
|
||||
2, 1, figsize=(10, 8), sharex=True
|
||||
)
|
||||
top: Any = axes[0] # type: ignore[reportAny]
|
||||
bottom: Any = axes[1] # type: ignore[reportAny]
|
||||
|
||||
for i, run in enumerate(series):
|
||||
color = _COLOR_CYCLE[i % len(_COLOR_CYCLE)]
|
||||
# Cumulative cold-prefill estimate (T_cum). Only plot points where
|
||||
# we have a t_cum value — skip if derived was empty for this run.
|
||||
n = min(len(run.cached_pp), len(run.cached_prefill_seconds))
|
||||
if n > 0:
|
||||
top.plot( # type: ignore[reportAny, reportUnknownMemberType]
|
||||
run.cached_pp[:n],
|
||||
run.cached_prefill_seconds[:n],
|
||||
"-o",
|
||||
color=color,
|
||||
label=run.label,
|
||||
)
|
||||
if run.control_pp:
|
||||
top.scatter( # type: ignore[reportAny, reportUnknownMemberType]
|
||||
run.control_pp,
|
||||
run.control_prefill_seconds,
|
||||
marker="x",
|
||||
s=80,
|
||||
color=color,
|
||||
label=f"{run.label} (cold one-shot)",
|
||||
)
|
||||
bottom.plot( # type: ignore[reportAny, reportUnknownMemberType]
|
||||
run.cached_pp,
|
||||
run.cached_gen_tps,
|
||||
"-o",
|
||||
color=color,
|
||||
label=run.label,
|
||||
)
|
||||
|
||||
top.set_ylabel("prefill time (s)") # type: ignore[reportAny, reportUnknownMemberType]
|
||||
top.set_title( # type: ignore[reportAny, reportUnknownMemberType]
|
||||
"cumulative cold-prefill time vs context size "
|
||||
"(line: T_cum estimate; ✕: cold one-shot control)"
|
||||
)
|
||||
top.grid(True, alpha=0.3) # type: ignore[reportAny, reportUnknownMemberType]
|
||||
top.legend(loc="best", fontsize=8) # type: ignore[reportAny, reportUnknownMemberType]
|
||||
|
||||
bottom.set_xlabel("pp_tokens") # type: ignore[reportAny, reportUnknownMemberType]
|
||||
bottom.set_ylabel("generation_tps (tok/s)") # type: ignore[reportAny, reportUnknownMemberType]
|
||||
bottom.set_title("decode throughput vs context size") # type: ignore[reportAny, reportUnknownMemberType]
|
||||
bottom.grid(True, alpha=0.3) # type: ignore[reportAny, reportUnknownMemberType]
|
||||
|
||||
if title is not None:
|
||||
fig.suptitle(title, fontsize=10) # type: ignore[reportUnknownMemberType]
|
||||
|
||||
fig.tight_layout()
|
||||
fig.savefig(output, dpi=120, bbox_inches="tight") # type: ignore[reportUnknownMemberType]
|
||||
plt.close(fig)
|
||||
@@ -0,0 +1,269 @@
|
||||
"""Typed prompt-sizing utilities for benchmarks.
|
||||
|
||||
Wraps the HuggingFace ``transformers`` tokenizer (a fundamentally dynamic
|
||||
object — different models return different types from
|
||||
``apply_chat_template``) behind a small typed API so the rest of the bench
|
||||
library can stay strict-typed.
|
||||
|
||||
``PromptSizer.build(target)`` returns a ``(content, exact_token_count)``
|
||||
pair. Internally it:
|
||||
1. Tokenises the empty user message to learn the chat-template overhead
|
||||
(``base_tokens``).
|
||||
2. Estimates tokens-per-atom from a 100-atom sample.
|
||||
3. Binary-searches over the atom count so the resulting message
|
||||
tokenises to *exactly* ``target`` tokens.
|
||||
|
||||
Callers downstream (``run_one_completion`` etc.) receive the verified
|
||||
token count, so analysis can confirm the prompt hit its target.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Any, Final, cast
|
||||
|
||||
|
||||
def _coerce_token_ids(raw: object) -> list[int]:
|
||||
"""Normalise ``apply_chat_template`` output to a flat list of token ids.
|
||||
|
||||
transformers' ``apply_chat_template`` may return:
|
||||
- ``list[int]`` (slow tokenizers, ``tokenize=True``)
|
||||
- a ``BatchEncoding`` with ``.input_ids`` (fast tokenizers)
|
||||
- a tensor wrapped object (some models)
|
||||
|
||||
We only need ``len(.)`` of the result, so we just need to flatten to a
|
||||
list and return it.
|
||||
"""
|
||||
if isinstance(raw, list):
|
||||
return cast("list[int]", raw)
|
||||
input_ids = getattr(raw, "input_ids", None)
|
||||
if isinstance(input_ids, list):
|
||||
return cast("list[int]", input_ids)
|
||||
raise TypeError(
|
||||
f"Unsupported tokenizer output type {type(raw).__name__}; "
|
||||
"expected list[int] or BatchEncoding-like with .input_ids."
|
||||
)
|
||||
|
||||
|
||||
def _build_token_counter(tokenizer: object) -> Callable[[str], int]:
|
||||
"""Return a closure that counts tokens for a user message.
|
||||
|
||||
Tries ``apply_chat_template`` first; falls back to the DeepSeek-V4
|
||||
Python encoder for models that don't ship a Jinja chat template.
|
||||
"""
|
||||
apply_chat_template = cast(
|
||||
Callable[..., object],
|
||||
tokenizer.apply_chat_template, # type: ignore[reportAttributeAccessIssue, reportUnknownMemberType]
|
||||
)
|
||||
encode = cast(
|
||||
Callable[..., list[int]],
|
||||
tokenizer.encode, # type: ignore[reportAttributeAccessIssue, reportUnknownMemberType]
|
||||
)
|
||||
|
||||
def count_fn(user_content: str) -> int:
|
||||
messages = [{"role": "user", "content": user_content}]
|
||||
try:
|
||||
raw = apply_chat_template(
|
||||
messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
except ValueError:
|
||||
# Models without a Jinja chat template (e.g. DeepSeek V4 which
|
||||
# ships its own Python encoder). Use the exo-side V4 encoder.
|
||||
from exo.worker.engines.mlx.vendor.deepseek_v4_encoding import ( # type: ignore[reportMissingTypeStubs]
|
||||
encode_messages as encode_v4,
|
||||
)
|
||||
|
||||
prompt = cast(str, encode_v4(messages, thinking_mode="thinking")) # type: ignore[reportUnknownArgumentType]
|
||||
raw = encode(prompt, add_special_tokens=False)
|
||||
return len(_coerce_token_ids(raw))
|
||||
|
||||
return count_fn
|
||||
|
||||
|
||||
class PromptSizer:
|
||||
"""Build a chat-completion content string of an exact token length."""
|
||||
|
||||
DEFAULT_ATOM: Final[str] = "a "
|
||||
|
||||
def __init__(self, tokenizer: object, atom: str = DEFAULT_ATOM):
|
||||
self._tokenizer = tokenizer
|
||||
self.atom = atom
|
||||
self._count_fn = _build_token_counter(tokenizer)
|
||||
self.base_tokens = self._count_fn("")
|
||||
|
||||
def count(self, content: str) -> int:
|
||||
"""Return the token count for ``content`` after chat-template expansion."""
|
||||
return self._count_fn(content)
|
||||
|
||||
def build(self, target_prompt_tokens: int) -> tuple[str, int]:
|
||||
"""Return ``(content, exact_token_count)`` summing to ``target``.
|
||||
|
||||
Raises ``RuntimeError`` if the chosen ``atom`` overshoots the target
|
||||
(try a different atom — see ``DEFAULT_ATOM``).
|
||||
"""
|
||||
target = int(target_prompt_tokens)
|
||||
if target < self.base_tokens:
|
||||
raise RuntimeError(
|
||||
f"Target ({target}) is smaller than template overhead "
|
||||
f"({self.base_tokens})."
|
||||
)
|
||||
|
||||
# Estimate tokens per atom using a sample.
|
||||
sample_count = 100
|
||||
sample_tokens = self._count_fn(self.atom * sample_count) - self.base_tokens
|
||||
tokens_per_atom = sample_tokens / sample_count
|
||||
needed_tokens = target - self.base_tokens
|
||||
estimated_atoms = int(needed_tokens / tokens_per_atom)
|
||||
|
||||
# Binary search to find exact atom count.
|
||||
low, high = 0, estimated_atoms * 2 + 100
|
||||
while low < high:
|
||||
mid = (low + high) // 2
|
||||
if self._count_fn(self.atom * mid) < target:
|
||||
low = mid + 1
|
||||
else:
|
||||
high = mid
|
||||
|
||||
content = self.atom * low
|
||||
actual = self._count_fn(content)
|
||||
if actual != target:
|
||||
raise RuntimeError(
|
||||
f"Overshot: got {actual} tokens (target {target}). "
|
||||
f"Pick a different atom (try ' a' or '\\n' or '0 ')."
|
||||
)
|
||||
return content, actual
|
||||
|
||||
|
||||
def _load_kimi_tokenizer(model_id: str) -> object:
|
||||
"""Special-case Kimi K2's custom TikTokenTokenizer (transformers 5.x quirk)."""
|
||||
from huggingface_hub import (
|
||||
snapshot_download, # type: ignore[reportUnknownVariableType]
|
||||
)
|
||||
|
||||
raw_path = snapshot_download(
|
||||
model_id,
|
||||
allow_patterns=[
|
||||
"*.json",
|
||||
"*.py",
|
||||
"*.tiktoken",
|
||||
"*.model",
|
||||
"*.jinja",
|
||||
],
|
||||
dry_run=False,
|
||||
)
|
||||
model_path = Path(raw_path)
|
||||
sys.path.insert(0, str(model_path))
|
||||
|
||||
tool_decl_path = model_path / "tool_declaration_ts.py"
|
||||
if tool_decl_path.exists():
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
"tool_declaration_ts", tool_decl_path
|
||||
)
|
||||
if spec is not None and spec.loader is not None:
|
||||
tool_decl_module = importlib.util.module_from_spec(spec)
|
||||
sys.modules["tool_declaration_ts"] = tool_decl_module
|
||||
spec.loader.exec_module(tool_decl_module)
|
||||
|
||||
tok_path = model_path / "tokenization_kimi.py"
|
||||
source = tok_path.read_text().replace(
|
||||
"from .tool_declaration_ts", "from tool_declaration_ts"
|
||||
)
|
||||
tok_module = types.ModuleType("tokenization_kimi")
|
||||
tok_module.__file__ = str(tok_path)
|
||||
sys.modules["tokenization_kimi"] = tok_module
|
||||
exec(compile(source, str(tok_path), "exec"), tok_module.__dict__) # noqa: S102
|
||||
|
||||
tik_token_cls = cast(Any, tok_module).TikTokenTokenizer # type: ignore[reportAny]
|
||||
hf_tokenizer = cast(Any, tik_token_cls.from_pretrained(model_path)) # type: ignore[reportAny]
|
||||
|
||||
# Patch encode to use internal tiktoken model directly (transformers 5.x
|
||||
# bug in the encode→pad path for slow tokenizers).
|
||||
def _patched_encode(text: str, **_kwargs: object) -> list[int]:
|
||||
return list(
|
||||
hf_tokenizer.model.encode(text, allowed_special="all") # type: ignore[reportAny, reportUnknownMemberType]
|
||||
)
|
||||
|
||||
hf_tokenizer.encode = _patched_encode
|
||||
return cast(object, hf_tokenizer)
|
||||
|
||||
|
||||
def load_tokenizer_for_bench(model_id: str) -> object:
|
||||
"""Load a HuggingFace tokenizer with bench-specific compatibility shims.
|
||||
|
||||
Returns the tokenizer as ``object`` because transformers' types are
|
||||
fundamentally dynamic (concrete class depends on the model). Callers
|
||||
should pass the result straight to :class:`PromptSizer`.
|
||||
"""
|
||||
# Monkey-patch for transformers 5.x: Kimi's tokenization_kimi.py imports
|
||||
# bytes_to_unicode from gpt2_tokenization which moved.
|
||||
try:
|
||||
import transformers.models.gpt2.tokenization_gpt2 as gpt2_tokenization
|
||||
from transformers.convert_slow_tokenizer import bytes_to_unicode
|
||||
|
||||
if not hasattr(gpt2_tokenization, "bytes_to_unicode"):
|
||||
gpt2_tokenization.bytes_to_unicode = bytes_to_unicode # type: ignore[reportAttributeAccessIssue]
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
if "kimi-k2" in model_id.lower():
|
||||
return _load_kimi_tokenizer(model_id)
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
try:
|
||||
return cast(
|
||||
object,
|
||||
AutoTokenizer.from_pretrained(model_id, trust_remote_code=True), # type: ignore[reportUnknownMemberType]
|
||||
)
|
||||
except (AttributeError, ValueError):
|
||||
# Some models ship a Jinja template / encoder that AutoTokenizer
|
||||
# can't introspect from HF directly — download artefacts and load
|
||||
# from the local snapshot path.
|
||||
from huggingface_hub import (
|
||||
snapshot_download, # type: ignore[reportUnknownVariableType]
|
||||
)
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
raw_full_path = snapshot_download(
|
||||
model_id,
|
||||
allow_patterns=[
|
||||
"*.json",
|
||||
"*.py",
|
||||
"tokenizer.model",
|
||||
"*.tiktoken",
|
||||
"tiktoken.model",
|
||||
"*.txt",
|
||||
"*.jsonl",
|
||||
"*.jinja",
|
||||
],
|
||||
dry_run=False,
|
||||
)
|
||||
model_path = Path(raw_full_path)
|
||||
stub_kwargs: dict[str, Any] = {}
|
||||
config_file = model_path / "config.json"
|
||||
if config_file.exists():
|
||||
with config_file.open() as f:
|
||||
raw_config: dict[str, Any] = json.load(f) # type: ignore[reportAny]
|
||||
for key in (
|
||||
"model_type",
|
||||
"max_position_embeddings",
|
||||
"vocab_size",
|
||||
"bos_token_id",
|
||||
"eos_token_id",
|
||||
"pad_token_id",
|
||||
):
|
||||
if key in raw_config:
|
||||
stub_kwargs[key] = raw_config[key]
|
||||
return cast(
|
||||
object,
|
||||
AutoTokenizer.from_pretrained( # type: ignore[reportUnknownMemberType]
|
||||
str(model_path),
|
||||
config=PretrainedConfig(**stub_kwargs), # type: ignore[reportArgumentType, reportAny]
|
||||
trust_remote_code=True,
|
||||
),
|
||||
)
|
||||
@@ -0,0 +1,139 @@
|
||||
"""Structured benchmark results — metadata capture + JSON output.
|
||||
|
||||
Every benchmark run produces a single JSON file with a stable schema:
|
||||
|
||||
- ``metadata``: exo SHA, ISO timestamps, hostnames, and any user-supplied
|
||||
tags identifying the run.
|
||||
- ``cluster``: snapshot from the API (node identities, topology, memory).
|
||||
- ``params``: the benchmark's input parameters (sweep config, etc).
|
||||
- ``runs``: per-request result rows.
|
||||
- ``derived``: any computed summaries (``t_cum_seconds`` for context scaling).
|
||||
|
||||
The format is intentionally additive so downstream tooling (plot scripts,
|
||||
dashboards) can rely on optional fields being absent rather than malformed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import platform
|
||||
import socket
|
||||
import subprocess
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from exo_tools.client import ExoClient
|
||||
from exo_tools.harness import capture_cluster_snapshot
|
||||
|
||||
|
||||
def _git_describe(repo_root: Path) -> str | None:
|
||||
"""Return ``<short-sha>[-dirty]`` for the repo at ``repo_root`` or None."""
|
||||
try:
|
||||
sha = subprocess.run(
|
||||
["git", "rev-parse", "--short=12", "HEAD"],
|
||||
cwd=str(repo_root),
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
check=True,
|
||||
).stdout.strip()
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
FileNotFoundError,
|
||||
subprocess.TimeoutExpired,
|
||||
):
|
||||
return None
|
||||
try:
|
||||
dirty = subprocess.run(
|
||||
["git", "status", "--porcelain"],
|
||||
cwd=str(repo_root),
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
check=True,
|
||||
).stdout.strip()
|
||||
return f"{sha}-dirty" if dirty else sha
|
||||
except (
|
||||
subprocess.CalledProcessError,
|
||||
FileNotFoundError,
|
||||
subprocess.TimeoutExpired,
|
||||
):
|
||||
return sha
|
||||
|
||||
|
||||
@dataclass
|
||||
class RunMetadata:
|
||||
"""Identifies a single bench run."""
|
||||
|
||||
run_id: str
|
||||
benchmark: str
|
||||
started_at: str
|
||||
finished_at: str | None = None
|
||||
exo_sha: str | None = None
|
||||
hostname: str = ""
|
||||
platform: str = ""
|
||||
tags: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def new(
|
||||
cls,
|
||||
benchmark: str,
|
||||
repo_root: Path,
|
||||
*,
|
||||
tags: dict[str, str] | None = None,
|
||||
) -> RunMetadata:
|
||||
now = datetime.now(timezone.utc)
|
||||
run_id = f"{benchmark}_{now.strftime('%Y%m%dT%H%M%SZ')}_{os.getpid()}"
|
||||
return cls(
|
||||
run_id=run_id,
|
||||
benchmark=benchmark,
|
||||
started_at=now.isoformat(),
|
||||
exo_sha=_git_describe(repo_root),
|
||||
hostname=socket.gethostname(),
|
||||
platform=f"{platform.system()} {platform.release()} ({platform.machine()})",
|
||||
tags=dict(tags or {}),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResultsBundle:
|
||||
"""Container for a single benchmark's results, before being written."""
|
||||
|
||||
metadata: RunMetadata
|
||||
params: dict[str, Any] = field(default_factory=dict)
|
||||
cluster: dict[str, Any] = field(default_factory=dict)
|
||||
runs: list[dict[str, Any]] = field(default_factory=list)
|
||||
cold_controls: list[dict[str, Any]] = field(default_factory=list)
|
||||
derived: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def capture_cluster(self, client: ExoClient) -> None:
|
||||
"""Snapshot the cluster state into ``self.cluster``."""
|
||||
try:
|
||||
snapshot = capture_cluster_snapshot(client)
|
||||
if snapshot:
|
||||
self.cluster.update(snapshot)
|
||||
except Exception:
|
||||
# Non-fatal: a benchmark without cluster snapshot is still valid
|
||||
pass
|
||||
|
||||
def write_json(self, output_dir: Path) -> Path:
|
||||
"""Write the bundle as ``<output_dir>/<run_id>.json`` and return the path."""
|
||||
if self.metadata.finished_at is None:
|
||||
self.metadata.finished_at = datetime.now(timezone.utc).isoformat()
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
path = output_dir / f"{self.metadata.run_id}.json"
|
||||
with path.open("w", encoding="utf-8") as f:
|
||||
json.dump(asdict(self), f, indent=2, ensure_ascii=False)
|
||||
return path
|
||||
|
||||
|
||||
def find_repo_root(start: Path | None = None) -> Path:
|
||||
"""Walk upwards from ``start`` (or this file) until a ``.git`` dir is found."""
|
||||
cur = (start or Path(__file__)).resolve()
|
||||
for parent in (cur, *cur.parents):
|
||||
if (parent / ".git").is_dir() or (parent / ".git").is_file():
|
||||
return parent
|
||||
raise RuntimeError(f"Could not locate repo root above {cur}")
|
||||
@@ -0,0 +1,65 @@
|
||||
"""BenchSession — wires together cluster + client + instance + tokenizer.
|
||||
|
||||
Holds the ``EcoSession``, a deployed ``ClusterInfo``, an ``ExoClient`` for
|
||||
the cluster's primary endpoint, and (for benchmarks that need exact-token
|
||||
prompts) a lazily-constructed :class:`PromptSizer`.
|
||||
|
||||
Benchmarks consume this via :func:`bench.lib.cluster.managed_instance`,
|
||||
which yields a populated ``BenchSession``. Library helpers (e.g.
|
||||
``context_scaling.run``) take a ``BenchSession`` and never reach for
|
||||
global state.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, cast
|
||||
|
||||
from exo_tools.client import ExoClient
|
||||
from exo_tools.cluster import ClusterInfo, EcoSession, make_client_from_url
|
||||
|
||||
from .prompt import PromptSizer, load_tokenizer_for_bench
|
||||
|
||||
|
||||
@dataclass
|
||||
class BenchSession:
|
||||
"""Bundle of cluster + client + (optional) instance for benchmarks."""
|
||||
|
||||
cluster: ClusterInfo
|
||||
eco: EcoSession
|
||||
instance_id: str | None = None
|
||||
model_id: str | None = None
|
||||
full_model_id: str | None = None
|
||||
_prompt_sizer: PromptSizer | None = field(default=None, repr=False)
|
||||
|
||||
@property
|
||||
def client(self) -> ExoClient:
|
||||
return make_client_from_url(self.cluster.api_url)
|
||||
|
||||
def state(self) -> dict[str, Any]:
|
||||
raw: Any = self.client.request_json("GET", "/state") # type: ignore[reportAny]
|
||||
if isinstance(raw, dict):
|
||||
return cast("dict[str, Any]", raw)
|
||||
return {}
|
||||
|
||||
def instances(self) -> dict[str, Any]:
|
||||
result: Any = self.state().get("instances", {}) # type: ignore[reportAny]
|
||||
if isinstance(result, dict):
|
||||
return cast("dict[str, Any]", result)
|
||||
return {}
|
||||
|
||||
def get_prompt_sizer(self) -> PromptSizer:
|
||||
"""Return a cached :class:`PromptSizer` for ``self.full_model_id``.
|
||||
|
||||
Loaded lazily because tokenizer load is expensive and not every
|
||||
benchmark needs prompt sizing.
|
||||
"""
|
||||
if self._prompt_sizer is not None:
|
||||
return self._prompt_sizer
|
||||
if self.full_model_id is None:
|
||||
raise RuntimeError(
|
||||
"BenchSession.full_model_id is not set; cannot build a PromptSizer."
|
||||
)
|
||||
tokenizer = load_tokenizer_for_bench(self.full_model_id)
|
||||
self._prompt_sizer = PromptSizer(tokenizer)
|
||||
return self._prompt_sizer
|
||||
Whitespace-only changes.
@@ -0,0 +1,186 @@
|
||||
"""Unit tests for the pure helpers in ``bench.lib.context_scaling``.
|
||||
|
||||
The orchestration entry points (``run``, ``run_cached_sweep``,
|
||||
``run_cold_controls``, ``make_cold_control_factory``) need a real
|
||||
``BenchSession`` and exo cluster, so they're exercised end-to-end via
|
||||
``python -m bench.cli context-scaling``. This module covers the
|
||||
underscore-prefixed pure helpers via direct private-symbol access (the
|
||||
private prefix discourages library users; tests for those helpers are
|
||||
the explicit exception).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import cast
|
||||
|
||||
from bench.lib.context_scaling import (
|
||||
StepResult,
|
||||
_compute_t_cum, # type: ignore[reportPrivateUsage]
|
||||
_interp, # type: ignore[reportPrivateUsage]
|
||||
derive_summary,
|
||||
)
|
||||
|
||||
|
||||
def _step(
|
||||
*,
|
||||
pp: int,
|
||||
delta: int,
|
||||
prompt_tps: float,
|
||||
generation_tps: float = 100.0,
|
||||
hit: str = "partial",
|
||||
) -> StepResult:
|
||||
return StepResult(
|
||||
pp_tokens=pp,
|
||||
delta_tokens=delta,
|
||||
prompt_tps=prompt_tps,
|
||||
generation_tps=generation_tps,
|
||||
prefix_cache_hit=hit,
|
||||
prompt_tokens=pp,
|
||||
generation_tokens=32,
|
||||
elapsed_s=delta / prompt_tps if prompt_tps else 0.0,
|
||||
)
|
||||
|
||||
|
||||
def _close(actual: float, expected: float, abs_tol: float = 1e-3) -> bool:
|
||||
return math.isclose(actual, expected, abs_tol=abs_tol)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _compute_t_cum
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestComputeTCum:
|
||||
def test_empty_returns_empty(self) -> None:
|
||||
assert _compute_t_cum([]) == []
|
||||
|
||||
def test_single_step(self) -> None:
|
||||
# 256 tokens at 1024 tps -> 0.25s
|
||||
out = _compute_t_cum([_step(pp=256, delta=256, prompt_tps=1024.0)])
|
||||
assert len(out) == 1
|
||||
assert _close(out[0], 0.25)
|
||||
|
||||
def test_cumulative_sum_across_three_steps(self) -> None:
|
||||
steps = [
|
||||
_step(pp=256, delta=256, prompt_tps=1000.0), # 0.256s
|
||||
_step(pp=512, delta=256, prompt_tps=2000.0), # +0.128s = 0.384s
|
||||
_step(pp=768, delta=256, prompt_tps=512.0), # +0.500s = 0.884s
|
||||
]
|
||||
out = _compute_t_cum(steps)
|
||||
assert _close(out[0], 0.256)
|
||||
assert _close(out[1], 0.384)
|
||||
assert _close(out[2], 0.884)
|
||||
# Monotonically non-decreasing
|
||||
assert out == sorted(out)
|
||||
|
||||
def test_zero_tps_step_skipped(self) -> None:
|
||||
# A row with prompt_tps == 0 contributes nothing to the cumulative sum
|
||||
steps = [
|
||||
_step(pp=256, delta=256, prompt_tps=1024.0), # +0.25s
|
||||
_step(pp=512, delta=256, prompt_tps=0.0), # +0
|
||||
_step(pp=768, delta=256, prompt_tps=512.0), # +0.5s
|
||||
]
|
||||
out = _compute_t_cum(steps)
|
||||
assert _close(out[0], 0.25)
|
||||
assert _close(out[1], 0.25) # unchanged
|
||||
assert _close(out[2], 0.75)
|
||||
|
||||
def test_zero_delta_step_skipped(self) -> None:
|
||||
# Defensive: a Δ=0 row would otherwise add zero anyway, but we
|
||||
# explicitly guard against negative delta + 0/0.
|
||||
steps = [
|
||||
_step(pp=256, delta=256, prompt_tps=1000.0),
|
||||
_step(pp=256, delta=0, prompt_tps=1000.0), # explicit Δ=0
|
||||
]
|
||||
out = _compute_t_cum(steps)
|
||||
assert _close(out[0], 0.256)
|
||||
assert _close(out[1], 0.256)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _interp
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestInterp:
|
||||
def test_empty_points_returns_zero(self) -> None:
|
||||
assert _interp([], 100) == 0.0
|
||||
|
||||
def test_single_point_returns_y(self) -> None:
|
||||
assert _interp([(100, 1.5)], 50) == 1.5
|
||||
assert _interp([(100, 1.5)], 100) == 1.5
|
||||
assert _interp([(100, 1.5)], 200) == 1.5
|
||||
|
||||
def test_clamps_below_first(self) -> None:
|
||||
points = [(100, 0.1), (200, 0.3), (300, 0.6)]
|
||||
assert _interp(points, 0) == 0.1
|
||||
assert _interp(points, 50) == 0.1
|
||||
assert _interp(points, 100) == 0.1
|
||||
|
||||
def test_clamps_above_last(self) -> None:
|
||||
points = [(100, 0.1), (200, 0.3), (300, 0.6)]
|
||||
assert _interp(points, 300) == 0.6
|
||||
assert _interp(points, 500) == 0.6
|
||||
assert _interp(points, 1_000_000) == 0.6
|
||||
|
||||
def test_mid_bracket_linear_interpolation(self) -> None:
|
||||
points = [(100, 0.0), (200, 1.0)]
|
||||
assert _close(_interp(points, 150), 0.5)
|
||||
assert _close(_interp(points, 175), 0.75)
|
||||
|
||||
def test_multi_segment_linear_interpolation(self) -> None:
|
||||
# Two adjacent segments, x=250 falls in the second one
|
||||
points = [(100, 0.1), (200, 0.3), (300, 0.6)]
|
||||
# 200..300: 0.3 + (0.6-0.3) * (250-200)/(300-200) = 0.3 + 0.15 = 0.45
|
||||
assert _close(_interp(points, 250), 0.45)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# derive_summary
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _gap_at(summary: dict[str, object], index: int) -> dict[str, float]:
|
||||
"""Cast ``summary['control_gaps'][index]`` into the typed shape we expect."""
|
||||
raw = summary["control_gaps"]
|
||||
assert isinstance(raw, list)
|
||||
entry = cast("dict[str, float]", raw[index])
|
||||
return entry
|
||||
|
||||
|
||||
class TestDeriveSummary:
|
||||
def test_no_controls_only_t_cum(self) -> None:
|
||||
steps = [
|
||||
_step(pp=256, delta=256, prompt_tps=1024.0),
|
||||
_step(pp=512, delta=256, prompt_tps=1024.0),
|
||||
]
|
||||
summary = derive_summary(steps, [])
|
||||
t_cum = cast("list[float]", summary["t_cum_seconds"])
|
||||
assert _close(t_cum[0], 0.25)
|
||||
assert _close(t_cum[1], 0.5)
|
||||
assert summary["control_gaps"] == []
|
||||
|
||||
def test_control_gap_at_known_pp(self) -> None:
|
||||
# Sweep: 0.25s @ pp=256, 0.5s @ pp=512
|
||||
steps = [
|
||||
_step(pp=256, delta=256, prompt_tps=1024.0),
|
||||
_step(pp=512, delta=256, prompt_tps=1024.0),
|
||||
]
|
||||
# Cold control at pp=512, 2x faster than the per-step rate -> 0.25s
|
||||
controls = [_step(pp=512, delta=512, prompt_tps=2048.0, hit="none")]
|
||||
summary = derive_summary(steps, controls)
|
||||
gap = _gap_at(summary, 0)
|
||||
assert gap["pp_tokens"] == 512
|
||||
assert _close(gap["cold_t_seconds"], 0.25, abs_tol=0.01)
|
||||
assert _close(gap["t_cum_seconds_at_pp"], 0.5, abs_tol=0.01)
|
||||
assert _close(gap["gap_seconds"], 0.25, abs_tol=0.01)
|
||||
# gap_fraction = 0.25 / 0.25 = 1.0
|
||||
assert _close(gap["gap_fraction"], 1.0, abs_tol=0.01)
|
||||
|
||||
def test_control_gap_zero_cold_tps_yields_zero_fraction(self) -> None:
|
||||
steps = [_step(pp=256, delta=256, prompt_tps=1000.0)]
|
||||
controls = [_step(pp=256, delta=256, prompt_tps=0.0, hit="none")]
|
||||
gap = _gap_at(derive_summary(steps, controls), 0)
|
||||
assert gap["cold_t_seconds"] == 0.0
|
||||
assert gap["gap_fraction"] == 0.0
|
||||
@@ -0,0 +1,171 @@
|
||||
"""Unit tests for ``bench.lib.model_meta``.
|
||||
|
||||
These exercise the pure derivation helpers (no HF round-trip). The HTTP
|
||||
fetchers (``fetch_model_meta``, ``_read_config_json``, ``_sum_weight_sizes``)
|
||||
hit the public hub and aren't covered here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
|
||||
from bench.lib.model_meta import (
|
||||
ModelMeta,
|
||||
derive_cold_controls,
|
||||
derive_context_ramp,
|
||||
)
|
||||
|
||||
|
||||
def _meta(
|
||||
*,
|
||||
weight_bytes: int = 0,
|
||||
max_pos: int = 4096,
|
||||
layers: int = 32,
|
||||
) -> ModelMeta:
|
||||
return ModelMeta(
|
||||
model_id="test/model",
|
||||
total_weight_bytes=weight_bytes,
|
||||
max_position_embeddings=max_pos,
|
||||
num_hidden_layers=layers,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ModelMeta properties
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestModelMetaConstraints:
|
||||
def test_zero_weight_yields_one_gib_floor(self) -> None:
|
||||
meta = _meta(weight_bytes=0)
|
||||
# int(0 * 1.30) + 1 == 1; int(0 * 1.10) + 1 == 1
|
||||
assert meta.memory_constraint_gb == 1.0
|
||||
assert meta.disk_constraint_gb == 1.0
|
||||
|
||||
def test_one_gib_weight_rounds_up(self) -> None:
|
||||
meta = _meta(weight_bytes=1 * (1024**3))
|
||||
# int(1.0 * 1.30) + 1 = 2; int(1.0 * 1.10) + 1 = 2
|
||||
assert meta.memory_constraint_gb == 2.0
|
||||
assert meta.disk_constraint_gb == 2.0
|
||||
|
||||
def test_sixteen_gib_weight_uses_30pct_memory_10pct_disk(self) -> None:
|
||||
meta = _meta(weight_bytes=16 * (1024**3))
|
||||
# memory: int(16 * 1.30) + 1 = 21; disk: int(16 * 1.10) + 1 = 18
|
||||
assert meta.memory_constraint_gb == 21.0
|
||||
assert meta.disk_constraint_gb == 18.0
|
||||
|
||||
def test_total_weight_gb_property(self) -> None:
|
||||
meta = _meta(weight_bytes=2_147_483_648) # 2 GiB exactly
|
||||
assert math.isclose(meta.total_weight_gb, 2.0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# derive_context_ramp
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDeriveContextRamp:
|
||||
def test_full_max_evenly_divides_round_to(self) -> None:
|
||||
meta = _meta(max_pos=131072) # 128k
|
||||
pp_step, num_steps = derive_context_ramp(meta, num_steps=32)
|
||||
# 131072 // 32 = 4096; rounded down to multiple of 256 = 4096
|
||||
assert pp_step == 4096
|
||||
assert num_steps == 32
|
||||
# Top of ramp == max
|
||||
assert pp_step * num_steps == 131072
|
||||
|
||||
def test_qwen30b_a3b_ramp(self) -> None:
|
||||
meta = _meta(max_pos=40960) # Qwen3-30B-A3B
|
||||
pp_step, num_steps = derive_context_ramp(meta, num_steps=32)
|
||||
# 40960 // 32 = 1280; multiple of 256
|
||||
assert pp_step == 1280
|
||||
assert pp_step * num_steps == 40960
|
||||
|
||||
def test_fraction_of_max_half(self) -> None:
|
||||
meta = _meta(max_pos=131072)
|
||||
pp_step, num_steps = derive_context_ramp(meta, num_steps=8, fraction_of_max=0.5)
|
||||
# half = 65536; 65536 // 8 = 8192
|
||||
assert pp_step == 8192
|
||||
assert num_steps == 8
|
||||
|
||||
def test_min_pp_step_floor(self) -> None:
|
||||
meta = _meta(max_pos=512)
|
||||
# 512 // 32 = 16, but min_pp_step=256 floors it; rounded to 256
|
||||
pp_step, num_steps = derive_context_ramp(meta, num_steps=32)
|
||||
assert pp_step == 256
|
||||
assert num_steps == 32
|
||||
|
||||
def test_round_to_truncates_down(self) -> None:
|
||||
meta = _meta(max_pos=10000)
|
||||
pp_step, _ = derive_context_ramp(meta, num_steps=32, round_to=256)
|
||||
# 10000 // 32 = 312; (312 // 256) * 256 = 256
|
||||
assert pp_step == 256
|
||||
|
||||
def test_round_to_zero_step_falls_back_to_round_to(self) -> None:
|
||||
# Pathological: huge round_to relative to per-step size
|
||||
meta = _meta(max_pos=1024)
|
||||
pp_step, _ = derive_context_ramp(meta, num_steps=8, round_to=1024)
|
||||
# 1024 // 8 = 128, but min_pp_step=256 → 256; (256 // 1024) * 1024 = 0;
|
||||
# `or round_to` rescues to 1024.
|
||||
assert pp_step == 1024
|
||||
|
||||
def test_max_pos_zero_raises(self) -> None:
|
||||
meta = _meta(max_pos=0)
|
||||
with pytest.raises(ValueError, match="max_position_embeddings=0"):
|
||||
_ = derive_context_ramp(meta, num_steps=32)
|
||||
|
||||
@pytest.mark.parametrize("fraction", [0.0, -0.1, 1.5, 2.0])
|
||||
def test_fraction_outside_unit_interval_raises(self, fraction: float) -> None:
|
||||
meta = _meta(max_pos=4096)
|
||||
with pytest.raises(ValueError, match="fraction_of_max"):
|
||||
_ = derive_context_ramp(meta, num_steps=4, fraction_of_max=fraction)
|
||||
|
||||
@pytest.mark.parametrize("steps", [0, -1, -100])
|
||||
def test_num_steps_must_be_positive(self, steps: int) -> None:
|
||||
meta = _meta(max_pos=4096)
|
||||
with pytest.raises(ValueError, match="num_steps"):
|
||||
_ = derive_context_ramp(meta, num_steps=steps)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# derive_cold_controls
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDeriveColdControls:
|
||||
def test_count_zero_returns_empty_tuple(self) -> None:
|
||||
meta = _meta()
|
||||
assert derive_cold_controls(meta, pp_step=4096, num_steps=32, count=0) == ()
|
||||
|
||||
def test_count_one_returns_top_only(self) -> None:
|
||||
meta = _meta()
|
||||
assert derive_cold_controls(meta, pp_step=4096, num_steps=32, count=1) == (
|
||||
131072,
|
||||
)
|
||||
|
||||
def test_evenly_spaced_four(self) -> None:
|
||||
meta = _meta()
|
||||
out = derive_cold_controls(meta, pp_step=4096, num_steps=32, count=4)
|
||||
# max_pp = 131072; (131072 * (i+1)) // 4 for i in {0,1,2,3}
|
||||
# = {32768, 65536, 98304, 131072}
|
||||
assert out == (32768, 65536, 98304, 131072)
|
||||
|
||||
def test_filters_below_pp_step(self) -> None:
|
||||
meta = _meta()
|
||||
out = derive_cold_controls(meta, pp_step=8192, num_steps=2, count=4)
|
||||
# max_pp = 16384; spaced points = {4096, 8192, 12288, 16384};
|
||||
# 4096 < pp_step=8192 → dropped.
|
||||
assert out == (8192, 12288, 16384)
|
||||
|
||||
def test_dedups_at_low_count_high_step(self) -> None:
|
||||
meta = _meta()
|
||||
# max_pp = 1024; count=2 → spaced = {512, 1024}; 512 < pp_step? No (=).
|
||||
out = derive_cold_controls(meta, pp_step=512, num_steps=2, count=2)
|
||||
assert out == (512, 1024)
|
||||
|
||||
def test_returned_in_ascending_order(self) -> None:
|
||||
meta = _meta()
|
||||
out = derive_cold_controls(meta, pp_step=1024, num_steps=8, count=4)
|
||||
assert list(out) == sorted(out)
|
||||
@@ -0,0 +1,189 @@
|
||||
"""Smoke tests for ``bench.lib.plotting``.
|
||||
|
||||
Renders a synthetic benchmark JSON to a tmp PNG and verifies the file is
|
||||
non-empty. We deliberately don't assert on pixel values — matplotlib
|
||||
output isn't byte-stable across versions — but a non-empty PNG with a
|
||||
valid header is a strong signal the renderer didn't throw.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from bench.lib.plotting import PlotInputs, render_context_scaling
|
||||
|
||||
|
||||
def _write_synthetic_run(path: Path, *, run_id: str, model: str = "test/model") -> None:
|
||||
"""Write a minimal context-scaling-shaped JSON for plotting tests."""
|
||||
payload = {
|
||||
"metadata": {
|
||||
"run_id": run_id,
|
||||
"benchmark": "context_scaling",
|
||||
"started_at": "2026-05-10T00:00:00Z",
|
||||
"exo_sha": "deadbeef",
|
||||
"hostname": "test-host",
|
||||
"platform": "Linux 6.0 (x86_64)",
|
||||
"tags": {"operator": "tester"},
|
||||
},
|
||||
"params": {
|
||||
"pp_step": 256,
|
||||
"num_steps": 4,
|
||||
"tg": 32,
|
||||
"warmup": 1,
|
||||
"full_model_id": model,
|
||||
},
|
||||
"cluster": {},
|
||||
"runs": [
|
||||
{
|
||||
"step_index": 0,
|
||||
"phase": "cached_sweep",
|
||||
"pp_tokens": 256,
|
||||
"delta_tokens": 256,
|
||||
"prompt_tps": 1800.0,
|
||||
"generation_tps": 410.0,
|
||||
"prefix_cache_hit": "exact",
|
||||
"prompt_tokens": 256,
|
||||
"generation_tokens": 32,
|
||||
"elapsed_s": 0.14,
|
||||
"peak_memory_bytes": 1_000_000_000,
|
||||
"output_text_preview": "",
|
||||
},
|
||||
{
|
||||
"step_index": 1,
|
||||
"phase": "cached_sweep",
|
||||
"pp_tokens": 512,
|
||||
"delta_tokens": 256,
|
||||
"prompt_tps": 2000.0,
|
||||
"generation_tps": 395.0,
|
||||
"prefix_cache_hit": "partial",
|
||||
"prompt_tokens": 512,
|
||||
"generation_tokens": 32,
|
||||
"elapsed_s": 0.13,
|
||||
"peak_memory_bytes": 1_100_000_000,
|
||||
"output_text_preview": "",
|
||||
},
|
||||
{
|
||||
"step_index": 2,
|
||||
"phase": "cached_sweep",
|
||||
"pp_tokens": 768,
|
||||
"delta_tokens": 256,
|
||||
"prompt_tps": 2200.0,
|
||||
"generation_tps": 378.0,
|
||||
"prefix_cache_hit": "partial",
|
||||
"prompt_tokens": 768,
|
||||
"generation_tokens": 32,
|
||||
"elapsed_s": 0.12,
|
||||
"peak_memory_bytes": 1_200_000_000,
|
||||
"output_text_preview": "",
|
||||
},
|
||||
],
|
||||
"cold_controls": [
|
||||
{
|
||||
"phase": "cold_control",
|
||||
"pp_tokens": 512,
|
||||
"delta_tokens": 512,
|
||||
"prompt_tps": 3200.0,
|
||||
"generation_tps": 400.0,
|
||||
"prefix_cache_hit": "none",
|
||||
"prompt_tokens": 512,
|
||||
"generation_tokens": 32,
|
||||
"elapsed_s": 0.16,
|
||||
"peak_memory_bytes": 1_500_000_000,
|
||||
"output_text_preview": "",
|
||||
},
|
||||
],
|
||||
"derived": {
|
||||
"t_cum_seconds": [0.14, 0.27, 0.39],
|
||||
"control_gaps": [],
|
||||
},
|
||||
}
|
||||
_ = path.write_text(json.dumps(payload))
|
||||
|
||||
|
||||
def _png_is_valid(path: Path) -> bool:
|
||||
"""A PNG file starts with the 8-byte magic ``\\x89PNG\\r\\n\\x1a\\n``."""
|
||||
if not path.is_file():
|
||||
return False
|
||||
if path.stat().st_size < 100:
|
||||
return False
|
||||
head = path.read_bytes()[:8]
|
||||
return head == b"\x89PNG\r\n\x1a\n"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRenderContextScaling:
|
||||
def test_single_run(self, tmp_path: Path) -> None:
|
||||
json_path = tmp_path / "run.json"
|
||||
_write_synthetic_run(json_path, run_id="r1")
|
||||
out = tmp_path / "out.png"
|
||||
|
||||
returned = render_context_scaling(PlotInputs(results=[json_path], output=out))
|
||||
assert returned == out
|
||||
assert _png_is_valid(out)
|
||||
|
||||
def test_creates_output_parent_dir(self, tmp_path: Path) -> None:
|
||||
json_path = tmp_path / "run.json"
|
||||
_write_synthetic_run(json_path, run_id="r1")
|
||||
out = tmp_path / "nested" / "deep" / "out.png"
|
||||
|
||||
_ = render_context_scaling(PlotInputs(results=[json_path], output=out))
|
||||
assert _png_is_valid(out)
|
||||
|
||||
def test_comparison_two_runs(self, tmp_path: Path) -> None:
|
||||
a = tmp_path / "a.json"
|
||||
b = tmp_path / "b.json"
|
||||
_write_synthetic_run(a, run_id="run-a", model="test/model-a")
|
||||
_write_synthetic_run(b, run_id="run-b", model="test/model-b")
|
||||
out = tmp_path / "compare.png"
|
||||
|
||||
_ = render_context_scaling(PlotInputs(results=[a, b], output=out))
|
||||
assert _png_is_valid(out)
|
||||
|
||||
def test_label_tag_uses_metadata_tag(self, tmp_path: Path) -> None:
|
||||
# Smoke test: just confirm passing label_tag doesn't throw and the
|
||||
# PNG renders. Label content is too matplotlib-internal to inspect.
|
||||
json_path = tmp_path / "run.json"
|
||||
_write_synthetic_run(json_path, run_id="r1")
|
||||
out = tmp_path / "out.png"
|
||||
|
||||
_ = render_context_scaling(
|
||||
PlotInputs(results=[json_path], output=out, label_tag="operator")
|
||||
)
|
||||
assert _png_is_valid(out)
|
||||
|
||||
def test_explicit_title(self, tmp_path: Path) -> None:
|
||||
json_path = tmp_path / "run.json"
|
||||
_write_synthetic_run(json_path, run_id="r1")
|
||||
out = tmp_path / "out.png"
|
||||
|
||||
_ = render_context_scaling(
|
||||
PlotInputs(results=[json_path], output=out, title="Custom Title")
|
||||
)
|
||||
assert _png_is_valid(out)
|
||||
|
||||
def test_empty_results_raises(self, tmp_path: Path) -> None:
|
||||
with pytest.raises(ValueError, match="at least one"):
|
||||
_ = render_context_scaling(
|
||||
PlotInputs(results=[], output=tmp_path / "out.png")
|
||||
)
|
||||
|
||||
def test_wrong_benchmark_raises(self, tmp_path: Path) -> None:
|
||||
# Same shape but with the wrong metadata.benchmark
|
||||
json_path = tmp_path / "run.json"
|
||||
_write_synthetic_run(json_path, run_id="r1")
|
||||
raw_loaded: Any = json.loads(json_path.read_text()) # type: ignore[reportAny]
|
||||
assert isinstance(raw_loaded, dict)
|
||||
data = cast("dict[str, dict[str, str]]", raw_loaded)
|
||||
data["metadata"]["benchmark"] = "something_else"
|
||||
_ = json_path.write_text(json.dumps(data))
|
||||
|
||||
with pytest.raises(ValueError, match="context_scaling"):
|
||||
_ = render_context_scaling(
|
||||
PlotInputs(results=[json_path], output=tmp_path / "out.png")
|
||||
)
|
||||
@@ -16,6 +16,7 @@ dependencies = [
|
||||
"lm-eval[api,math]>=0.4.0",
|
||||
"human-eval>=1.0.3",
|
||||
"numpy>=1.24.0",
|
||||
"matplotlib>=3.8",
|
||||
]
|
||||
|
||||
[build-system]
|
||||
|
||||
+8
-3
@@ -149,7 +149,9 @@ root = "src"
|
||||
|
||||
[[tool.basedpyright.executionEnvironments]]
|
||||
root = "bench"
|
||||
extraPaths = ["tools/src"]
|
||||
# `.` keeps `from bench.lib.X import …` resolvable (pytest adds the project
|
||||
# root to sys.path; we want type-checking to agree with runtime).
|
||||
extraPaths = ["tools/src", "."]
|
||||
|
||||
[[tool.basedpyright.executionEnvironments]]
|
||||
root = "tools/src"
|
||||
@@ -224,9 +226,12 @@ extend-exclude = [
|
||||
extend-select = ["I", "N", "B", "A", "PIE", "SIM"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
pythonpath = "."
|
||||
pythonpath = ["."]
|
||||
asyncio_mode = "auto"
|
||||
markers = ["slow: marks tests as slow (deselected by default)"]
|
||||
env = ["EXO_TESTS=1"]
|
||||
addopts = "-m 'not slow' --ignore=tests"
|
||||
# `tests/` requires an eco cluster (opt-in). `tmp/` holds throwaway scripts
|
||||
# that run a top-level `sys.exit(...)` at import time, which otherwise blows
|
||||
# up the default pytest collection.
|
||||
addopts = "-m 'not slow' --ignore=tests --ignore=tmp"
|
||||
filterwarnings = ["ignore:builtin type Swig:DeprecationWarning"]
|
||||
@@ -12,6 +12,7 @@ import atexit
|
||||
import contextlib
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import signal
|
||||
import subprocess
|
||||
@@ -25,6 +26,7 @@ from .client import ExoClient
|
||||
class Thunderbolt(str, Enum):
|
||||
A2A = "a2a" # all-to-all (eco --tb-a2a)
|
||||
RING = "ring" # ring topology (eco --tb-ring)
|
||||
NONE = "none" # exclude Thunderbolt-connected hosts (eco --no-thunderbolt)
|
||||
|
||||
|
||||
class Chip(str, Enum):
|
||||
@@ -143,6 +145,9 @@ class EcoSession:
|
||||
thunderbolt: Thunderbolt | None = None,
|
||||
chip: Chip | None = None,
|
||||
min_memory_gb: float | None = None,
|
||||
max_memory_gb: float | None = None,
|
||||
min_disk_gb: float | None = None,
|
||||
max_disk_gb: float | None = None,
|
||||
wait: bool = True,
|
||||
ref: str | None = _EXO_REF,
|
||||
timeout: int = 600,
|
||||
@@ -151,18 +156,32 @@ class EcoSession:
|
||||
|
||||
By default, deploys from local source via rsync. Set EXO_REF
|
||||
or pass ref= to deploy from a GitHub branch/tag instead (for CI).
|
||||
|
||||
Selection constraints (memory/disk in GiB, chip substring,
|
||||
Thunderbolt topology) are forwarded as eco CLI flags. Pass
|
||||
``thunderbolt=Thunderbolt.NONE`` to exclude TB-connected hosts.
|
||||
"""
|
||||
cmd: list[str] = ["eco", "--json", "start", "--deploy"]
|
||||
if hosts:
|
||||
cmd.extend(hosts)
|
||||
if count is not None:
|
||||
cmd.extend(["--count", str(count)])
|
||||
if thunderbolt is not None:
|
||||
if thunderbolt is Thunderbolt.NONE:
|
||||
cmd.append("--no-thunderbolt")
|
||||
elif thunderbolt is not None:
|
||||
cmd.append(f"--tb-{thunderbolt.value}")
|
||||
if chip is not None:
|
||||
cmd.extend(["--chip", chip.value])
|
||||
# eco's GB args are integer-typed. Round mins up + maxes down so
|
||||
# we never relax the user's constraint.
|
||||
if min_memory_gb is not None:
|
||||
cmd.extend(["--min-memory", str(min_memory_gb)])
|
||||
cmd.extend(["--min-memory", str(math.ceil(min_memory_gb))])
|
||||
if max_memory_gb is not None:
|
||||
cmd.extend(["--max-memory", str(math.floor(max_memory_gb))])
|
||||
if min_disk_gb is not None:
|
||||
cmd.extend(["--min-disk", str(math.ceil(min_disk_gb))])
|
||||
if max_disk_gb is not None:
|
||||
cmd.extend(["--max-disk", str(math.floor(max_disk_gb))])
|
||||
if wait:
|
||||
cmd.append("--wait")
|
||||
if ref:
|
||||
|
||||
@@ -255,7 +255,7 @@ name = "contourpy"
|
||||
version = "1.3.3"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "numpy", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "numpy", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/58/01/1253e6698a07380cd31a736d248a3f2a50a7c88779a1813da27503cadc2a/contourpy-1.3.3.tar.gz", hash = "sha256:083e12155b210502d0bca491432bb04d56dc3432f95a979b429f2848c3dbe880", size = 13466174, upload-time = "2025-07-26T12:03:12.549Z" }
|
||||
wheels = [
|
||||
@@ -507,6 +507,7 @@ dependencies = [
|
||||
{ name = "lm-eval", extra = ["api", "math"], marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "loguru", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "math-verify", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "matplotlib", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "numpy", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "protobuf", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "tiktoken", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
@@ -523,6 +524,7 @@ requires-dist = [
|
||||
{ name = "lm-eval", extras = ["api", "math"], specifier = ">=0.4.0" },
|
||||
{ name = "loguru", specifier = ">=0.7.3" },
|
||||
{ name = "math-verify", specifier = ">=0.7.0" },
|
||||
{ name = "matplotlib", specifier = ">=3.8" },
|
||||
{ name = "numpy", specifier = ">=1.24.0" },
|
||||
{ name = "protobuf", specifier = ">=5.29.0" },
|
||||
{ name = "tiktoken", specifier = ">=0.12.0" },
|
||||
@@ -1144,15 +1146,15 @@ name = "matplotlib"
|
||||
version = "3.10.8"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "contourpy", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "cycler", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "fonttools", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "kiwisolver", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "numpy", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "packaging", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "pillow", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "pyparsing", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "python-dateutil", marker = "sys_platform == 'darwin' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "contourpy", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "cycler", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "fonttools", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "kiwisolver", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "numpy", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "packaging", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "pillow", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "pyparsing", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
{ name = "python-dateutil", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda12') or (extra == 'extra-3-exo-cpu' and extra == 'extra-3-exo-cuda13') or (extra == 'extra-3-exo-cuda12' and extra == 'extra-3-exo-cuda13')" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/8a/76/d3c6e3a13fe484ebe7718d14e269c9569c4eb0020a968a327acb3b9a8fe6/matplotlib-3.10.8.tar.gz", hash = "sha256:2299372c19d56bcd35cf05a2738308758d32b9eaed2371898d8f5bd33f084aa3", size = 34806269, upload-time = "2025-12-10T22:56:51.155Z" }
|
||||
wheels = [
|
||||
|
||||
Reference in new issue
Block a user