Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 51 additions & 0 deletions benchmarks/laya_mps/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,10 @@ JSONL to `results/` (kept out of the repository); `report.py` builds the tables
| `bench_inproc.py` | Laya in-process (no HTTP): load, warmup, first request, warm latency, memory |
| `bench_http.py` | a `/v1/systemone` server, optionally started by the script and optionally behind the frontend: time to ready, first request, warm latency, throughput |
| `paired.py` | two configurations compared request by request, both alive at once: two worker flag sets, or two running servers (e.g. a worker directly and through the frontend) |
| `lengths.py` | first request of new input lengths and the memory it adds, for one or two worker flag sets |
| `late_load.py` | a checkpoint loaded while serving: its first request directly and through the frontend |
| `fallback.py` | Laya's fallback to the CPU on a GPU out-of-memory error, triggered by a lowered MPS memory limit |
| `release.py` | in-process: what releasing PyTorch's MPS caches gives back after many lengths, and what it costs |
| `profile_mps.py` | where a request's time goes on MPS |
| `report.py` | tables from the JSONL, including the run-to-run gate and the answer comparison against a reference config |
| `env.py` | shared: versions, checkpoint, hardware, power and load recorded with each run; memory footprint |
Expand Down Expand Up @@ -43,6 +47,53 @@ python benchmarks/laya_mps/paired.py --run f1 --a-url http://127.0.0.1:8000 --b-
python benchmarks/laya_mps/paired.py --summarize benchmarks/laya_mps/results/paired_p1.jsonl
```

## Reproduce every number in the recipe

Each claim in the [recipe](../../recipe/laya/apple-silicon.md) comes from one of these commands. A
rerun on another Mac, or under different load, gives other numbers; the comparison each command makes
(A against B in the same run) is what carries over. Every script refuses a measured run on battery
power or above `--max-load`: a `--run` label other than `feasibility`, or for `late_load.py`,
`fallback.py` and `release.py` a run without `--feasibility`; paired and lengths runs hold up better
than separate ones under the load that remains, because both sides see it. Stop the recipe's worker and
frontend first: the scripts start their own on ports 8000, 8001 and 8080.

| recipe claim | command |
| --- | --- |
| checkpoint download time | `HF_HOME="$(mktemp -d)" python -c "import time, huggingface_hub as h; t = time.time(); h.snapshot_download('convaiinnovations/laya', allow_patterns=['rl_agent_config.json', 'model.safetensors', 'tokenizer/*', 'encoder/*']); print(f'{time.time() - t:.0f} s')"` |
| first request after ready, time to ready, memory: worker against laya-serve, with and without the options | `bench_http.py` C3, C3w and C3o above, then `report.py` (phases and memory tables) |
| first request after ready over many fresh starts | the loop below the table |
| warm latency and answers, with the options against without | `paired.py --run p1 --a "" --b "--compile --weights fp16"` |
| what each option contributes | `paired.py --run p2 --a "" --b=--compile`, `paired.py --run p3 --a "" --b "--weights fp16"`, and fp16 on top of compile: `paired.py --run p4 --a=--compile --b "--compile --weights fp16"` |
| frontend overhead | start a worker on 8000 and the frontend on 8080 as in the recipe, then `paired.py --run f1 --a-url http://127.0.0.1:8000 --b-url http://127.0.0.1:8080` |
| a request after an idle gap | `paired.py --run i1 --a "" --b "--compile --weights fp16" --gap 2 --only W1 -n 30 --discard 2` (also `--gap 0.5`) |
| first request of a new input length; memory growth with the lengths seen | `lengths.py --run l1 --a "" --b "--compile --weights fp16"` (add `--first-words 1 --step 1 --lengths 1000` for every length up to the window) |
| memory released by `torch.mps.empty_cache()` and the cost afterwards | `release.py --compile --weights fp16` |
| a checkpoint loaded while serving, directly and through the frontend | `late_load.py --flags "--compile --weights fp16" --frontend target/release/omni-jev`, and without `--flags` |
| fallback to the CPU on a GPU out-of-memory error | `fallback.py --limit-gb 3.5`, and `--limit-gb 2.5 --flags "--compile --weights fp16"` |
| where the time goes | `profile_mps.py --run p1` |
| two commits compared | start each worker from its own checkout on its own port, then `paired.py --a-url ... --b-url ...` |

Each fresh start is one `bench_http.py` run, alternating the worker with the options and plain
laya-serve; the phases table of `report.py` then lists the first request of every start on both
sides (one measured run per start, so it needs an idle machine):

```sh
for i in $(seq 23); do
python benchmarks/laya_mps/bench_http.py --config C3o --run s$i --only W1 -n 1 --discard 1 \
--concurrency 1 --spawn .venv/bin/python -m frontend.laya_mps --device {device} --model {model} \
--compile --weights fp16 --port {port}
python benchmarks/laya_mps/bench_http.py --config C3 --run s$i --only W1 -n 1 --discard 1 \
--concurrency 1 --spawn .venv/bin/laya-serve
done
python benchmarks/laya_mps/report.py benchmarks/laya_mps/results/http_C3o_s*.jsonl \
benchmarks/laya_mps/results/http_C3_s*.jsonl --ref C3
```

All scripts are run as `python benchmarks/laya_mps/<script>` from the repository root; `--summarize`
rebuilds a paired or lengths table from its JSONL. The published paired runs were made with the earlier
environment-variable form of the options (`LAYA_WORKER_COMPILE=on LAYA_WORKER_WEIGHTS=fp16`), which the
flags replace one for one.

## Results

The measured runs on an M1 Pro are published as assets of one release on the fork,
Expand Down
145 changes: 107 additions & 38 deletions benchmarks/laya_mps/bench_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,17 +34,19 @@
HERE = Path(__file__).resolve().parent
REPO = HERE.parents[1]
sys.path.insert(0, str(HERE))
from env import footprint_mb, header, noise_problems # noqa: E402
from env import footprint_mb, header, read_workloads, refuse_if_noisy # noqa: E402

CHECKPOINT = "convaiinnovations/laya" # what laya-serve's "english" model resolves to (laya/router.py)


class Client:
"""One keep-alive connection. Not thread-safe: one per thread."""

def __init__(self, url, token=None):
def __init__(self, url, token=None, timeout=120):
parts = urlsplit(url)
self.conn = http.client.HTTPConnection(parts.hostname, parts.port or 80, timeout=120)
self.conn = http.client.HTTPConnection(
parts.hostname, parts.port or 80, timeout=timeout
)
self.headers = {"Content-Type": "application/json"}
if token:
self.headers["Authorization"] = f"Bearer {token}"
Expand Down Expand Up @@ -107,18 +109,22 @@ def fetch_answers(client, body):


def body_for(workload, model):
return json.dumps({"model": model, "state": workload["state"], "questions": workload["questions"]}).encode()
return json.dumps(
{"model": model, "state": workload["state"], "questions": workload["questions"]}
).encode()


def run_level(url, token, body, n, concurrency):
"""n requests split over `concurrency` threads. Returns [(thread, ms, status)], elapsed seconds."""
per_thread = [n // concurrency + (i < n % concurrency) for i in range(concurrency)]
results, lock = [], threading.Lock()
barrier = threading.Barrier(concurrency + 1, timeout=120) # a thread that fails to connect breaks it
# a thread that fails to connect breaks it
barrier = threading.Barrier(concurrency + 1, timeout=120)

def worker(index, count):
client = Client(url, token)
client.request("POST", "/v1/systemone", body, retry=True) # connect outside the timed window
# connect outside the timed window
client.request("POST", "/v1/systemone", body, retry=True)
barrier.wait()
mine = []
for _ in range(count):
Expand All @@ -131,7 +137,9 @@ def worker(index, count):
with lock:
results.extend(mine)

threads = [threading.Thread(target=worker, args=(i, c)) for i, c in enumerate(per_thread)]
threads = [
threading.Thread(target=worker, args=(i, c)) for i, c in enumerate(per_thread)
]
for t in threads:
t.start()
barrier.wait()
Expand All @@ -146,36 +154,69 @@ def main():
parser.add_argument("--config", required=True, help="label, e.g. C3 or C4")
parser.add_argument("--run", required=True, help="feasibility, m1, m2, ...")
parser.add_argument("--url", default="http://127.0.0.1:8000")
parser.add_argument("--model", default="english", help="model name the worker serves")
parser.add_argument("--frontend", help="Rust frontend binary to start on --url in front of the spawned worker")
parser.add_argument("--backend-port", type=int, default=8000, help="worker port when --frontend is used")
parser.add_argument("--spawn", nargs=argparse.REMAINDER, help="start this worker command, then benchmark it")
parser.add_argument("--device", default="mps", help="device for a spawned worker: LAYA_DEVICE and {device}")
parser.add_argument(
"--model", default="english", help="model name the worker serves"
)
parser.add_argument(
"--frontend",
help="Rust frontend binary to start on --url in front of the spawned worker",
)
parser.add_argument(
"--backend-port",
type=int,
default=8000,
help="worker port when --frontend is used",
)
parser.add_argument(
"--spawn",
nargs=argparse.REMAINDER,
help="start this worker command, then benchmark it",
)
parser.add_argument(
"--device",
default="mps",
help="device for a spawned worker: LAYA_DEVICE and {device}",
)
parser.add_argument("--ready-timeout", type=float, default=600)
parser.add_argument("--workloads", default=str(HERE / "workloads.jsonl"))
parser.add_argument("--only", nargs="*", help="bench workload ids to run (default: all)")
parser.add_argument("-n", type=int, default=300, help="timed requests per workload and concurrency level")
parser.add_argument("--discard", type=int, default=20, help="warmup requests per workload")
parser.add_argument(
"--only", nargs="*", help="bench workload ids to run (default: all)"
)
parser.add_argument(
"-n",
type=int,
default=300,
help="timed requests per workload and concurrency level",
)
parser.add_argument(
"--discard", type=int, default=20, help="warmup requests per workload"
)
parser.add_argument("--concurrency", type=int, nargs="+", default=[1, 4])
parser.add_argument("--seed", type=int, default=0, help="workload order seed")
parser.add_argument("--out", default=str(HERE / "results"))
parser.add_argument("--max-load", type=float, default=2.0, help="1-min load average allowed for measured runs")
parser.add_argument(
"--max-load",
type=float,
default=2.0,
help="1-min load average allowed for measured runs",
)
args = parser.parse_args()
token = os.environ.get("OMNI_JEV_TEST_TOKEN")
if args.frontend and not args.spawn:
parser.error("--frontend needs --spawn: the frontend is started in front of a spawned worker")
parser.error(
"--frontend needs --spawn: the frontend is started in front of a spawned worker"
)
if args.frontend and urlsplit(args.url).port == args.backend_port:
parser.error("--url and --backend-port must differ when --frontend is used")

problems = noise_problems(args.max_load)
if problems and args.run != "feasibility":
sys.exit("refusing a measured run: " + "; ".join(problems))
for problem in problems:
print(f"warning: {problem}", file=sys.stderr)
problems = refuse_if_noisy(args.max_load, args.run != "feasibility")

with open(args.workloads) as f:
workloads = [json.loads(line) for line in f if line.strip()]
bench = [w for w in workloads if w["kind"] == "bench" and (not args.only or w["id"] in args.only)]
workloads = read_workloads(args.workloads).values()
bench = [
w
for w in workloads
if w["kind"] == "bench" and (not args.only or w["id"] in args.only)
]
parity = [w for w in workloads if w["kind"] == "parity"]

Path(args.out).mkdir(parents=True, exist_ok=True)
Expand All @@ -184,7 +225,9 @@ def main():
def memory():
mem = footprint_mb(processes["worker"].pid) if "worker" in processes else {}
if "frontend" in processes:
mem["frontend_footprint_mb"] = footprint_mb(processes["frontend"].pid).get("footprint_mb")
mem["frontend_footprint_mb"] = footprint_mb(processes["frontend"].pid).get(
"footprint_mb"
)
return mem

out = Path(args.out) / f"http_{args.config}_{args.run}.jsonl"
Expand All @@ -201,12 +244,19 @@ def memory():
"LAYA_PRELOAD": "1",
"LAYA_LOG_LEVEL": "warning",
}
env["PYTHONPATH"] = str(REPO / "src") + (os.pathsep + env["PYTHONPATH"] if env.get("PYTHONPATH") else "")
placeholders = {"{port}": str(port), "{device}": args.device, "{model}": args.model}
env["PYTHONPATH"] = str(REPO / "src") + (
os.pathsep + env["PYTHONPATH"] if env.get("PYTHONPATH") else ""
)
placeholders = {
"{port}": str(port),
"{device}": args.device,
"{model}": args.model,
}
command = list(args.spawn)
for placeholder, value in placeholders.items():
command = [arg.replace(placeholder, value) for arg in command]
spawn_log = open(Path(args.out) / f"http_{args.config}_{args.run}.worker.log", "w") # noqa: SIM115
log_path = Path(args.out) / f"http_{args.config}_{args.run}.worker.log"
spawn_log = open(log_path, "w") # noqa: SIM115
processes["worker"] = subprocess.Popen(
command, env=env, stdout=spawn_log, stderr=subprocess.STDOUT, cwd=REPO
)
Expand All @@ -217,7 +267,8 @@ def memory():
"OMNI_JEV_BIND": f"{parts.hostname}:{parts.port}",
"OMNI_JEV_BACKEND_URL": f"http://127.0.0.1:{args.backend_port}",
}
frontend_log = open(Path(args.out) / f"http_{args.config}_{args.run}.frontend.log", "w") # noqa: SIM115
log_path = Path(args.out) / f"http_{args.config}_{args.run}.frontend.log"
frontend_log = open(log_path, "w") # noqa: SIM115
processes["frontend"] = subprocess.Popen(
[args.frontend], env=env, stdout=frontend_log, stderr=subprocess.STDOUT
)
Expand All @@ -240,30 +291,45 @@ def emit(record):
health=health,
device_actual=health.get("device"),
# The worker reports these in /health; laya-serve does not, so its values are assumed.
mps_amp_min_rows=health.get("mps_amp_min_rows", int(os.environ.get("LAYA_MPS_AMP_MIN_ROWS", "5"))),
mps_amp_min_rows=health.get(
"mps_amp_min_rows",
int(os.environ.get("LAYA_MPS_AMP_MIN_ROWS", "5")),
),
amp_dtype=health.get("autocast_dtype")
or ("torch.float16" if health.get("device") == "mps" else "torch.float32"),
or (
"torch.float16"
if health.get("device") == "mps"
else "torch.float32"
),
weights_dtype=health.get("weights_dtype") or "torch.float32",
dtype_source=None if "weights_dtype" in health else "assumed: laya 0.3.20 defaults",
dtype_source=None
if "weights_dtype" in health
else "assumed: laya 0.3.20 defaults",
)
)

client = Client(args.url, token)
started = time.perf_counter()
first_ms, routing = {}, None
for w in bench:
ms, status, data = client.request("POST", "/v1/systemone", body_for(w, args.model), retry=True)
ms, status, data = client.request(
"POST", "/v1/systemone", body_for(w, args.model), retry=True
)
if status != 200:
sys.exit(f"{w['id']}: status {status}: {data[:200]!r}")
first_ms[w["id"]] = ms
routing = json.loads(data).get("routing")
for _ in range(args.discard - 1):
client.request("POST", "/v1/systemone", body_for(w, args.model), retry=True)
client.request(
"POST", "/v1/systemone", body_for(w, args.model), retry=True
)
warmup_s = time.perf_counter() - started
emit(
{
"type": "phase",
"process_to_ready_s": round(ready_s, 3) if "worker" in processes else None,
"process_to_ready_s": round(ready_s, 3)
if "worker" in processes
else None,
"warmup_s": round(warmup_s, 3),
"first_ms": {k: round(v, 2) for k, v in first_ms.items()},
"routing": routing,
Expand All @@ -277,12 +343,15 @@ def emit(record):
for w in order:
body = body_for(w, args.model)
answers, error = fetch_answers(client, body)
if error: # the parity section of report.py reports the workload as missing
# the parity section of report.py reports the workload as missing
if error:
emit({"type": "answers_error", "workload": w["id"], **error})
else:
emit({"type": "answers", "workload": w["id"], "answers": answers})
for concurrency in args.concurrency:
results, elapsed = run_level(args.url, token, body, args.n, concurrency)
results, elapsed = run_level(
args.url, token, body, args.n, concurrency
)
ok = sum(s == 200 for _, _, s in results)
for i, (thread, ms, status) in enumerate(results):
emit(
Expand Down
Loading
Loading