From 3de090fece0f800738bb793a0741f40e333e869d Mon Sep 17 00:00:00 2001 From: cacheline999 <326908201+cacheline999@users.noreply.github.com> Date: Sat, 3 Oct 2026 15:51:18 +0800 Subject: [PATCH 1/4] Laya: format the worker, benchmark scripts and tests with ruff 0.16.10 defaults Formatting only (88 columns, as the rest of the repository); no behaviour change. --- benchmarks/laya_mps/bench_http.py | 137 +++++++--- benchmarks/laya_mps/bench_inproc.py | 42 ++- benchmarks/laya_mps/env.py | 32 ++- benchmarks/laya_mps/paired.py | 100 ++++++-- benchmarks/laya_mps/profile_mps.py | 49 +++- benchmarks/laya_mps/report.py | 135 ++++++++-- src/frontend/laya_mps.py | 100 ++++++-- src/models/laya/engine.py | 47 +++- src/models/laya/optimize.py | 8 +- tests/laya/test_bench.py | 115 +++++++-- tests/laya/test_contract.py | 84 +++++- tests/laya/test_worker.py | 382 ++++++++++++++++++++++------ 12 files changed, 992 insertions(+), 239 deletions(-) diff --git a/benchmarks/laya_mps/bench_http.py b/benchmarks/laya_mps/bench_http.py index da8eb4a..b8e5c7b 100644 --- a/benchmarks/laya_mps/bench_http.py +++ b/benchmarks/laya_mps/bench_http.py @@ -44,7 +44,9 @@ class Client: def __init__(self, url, token=None): 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=120 + ) self.headers = {"Content-Type": "application/json"} if token: self.headers["Authorization"] = f"Bearer {token}" @@ -107,18 +109,24 @@ 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 + barrier = threading.Barrier( + concurrency + 1, timeout=120 + ) # a thread that fails to connect breaks it def worker(index, count): client = Client(url, token) - client.request("POST", "/v1/systemone", body, retry=True) # connect outside the timed window + client.request( + "POST", "/v1/systemone", body, retry=True + ) # connect outside the timed window barrier.wait() mine = [] for _ in range(count): @@ -131,7 +139,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() @@ -146,24 +156,58 @@ 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") @@ -175,7 +219,11 @@ def main(): 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)] + 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) @@ -184,7 +232,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" @@ -201,12 +251,20 @@ 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 + spawn_log = open( + Path(args.out) / f"http_{args.config}_{args.run}.worker.log", "w" + ) # noqa: SIM115 processes["worker"] = subprocess.Popen( command, env=env, stdout=spawn_log, stderr=subprocess.STDOUT, cwd=REPO ) @@ -217,7 +275,9 @@ 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 + frontend_log = open( + Path(args.out) / f"http_{args.config}_{args.run}.frontend.log", "w" + ) # noqa: SIM115 processes["frontend"] = subprocess.Popen( [args.frontend], env=env, stdout=frontend_log, stderr=subprocess.STDOUT ) @@ -240,11 +300,20 @@ 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", ) ) @@ -252,18 +321,24 @@ def emit(record): 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, @@ -277,12 +352,16 @@ 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 + if ( + error + ): # the parity section of report.py reports the workload as missing 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( diff --git a/benchmarks/laya_mps/bench_inproc.py b/benchmarks/laya_mps/bench_inproc.py index ea048fd..14ecb55 100644 --- a/benchmarks/laya_mps/bench_inproc.py +++ b/benchmarks/laya_mps/bench_inproc.py @@ -57,12 +57,21 @@ def main(): parser.add_argument("--run", required=True, help="feasibility, m1, m2, ...") parser.add_argument("--checkpoint", default="convaiinnovations/laya") 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( + "--only", nargs="*", help="bench workload ids to run (default: all)" + ) parser.add_argument("-n", type=int, default=300, help="timed requests per workload") - parser.add_argument("--discard", type=int, default=20, help="warmup requests per workload") + parser.add_argument( + "--discard", type=int, default=20, help="warmup requests per workload" + ) 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() problems = noise_problems(args.max_load) @@ -72,7 +81,11 @@ def main(): print(f"warning: {problem}", file=sys.stderr) workloads = load_workloads(args.workloads) - bench = [w for w in workloads if w["kind"] == "bench" and (not args.only or w["id"] in args.only)] + 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"] out = Path(args.out) / f"inproc_{args.config}_{args.device}_{args.run}.jsonl" @@ -101,7 +114,10 @@ def emit(record): ) ) if agent.device.type != args.device: - print(f"warning: asked for {args.device}, laya is on {agent.device}", file=sys.stderr) + print( + f"warning: asked for {args.device}, laya is on {agent.device}", + file=sys.stderr, + ) # Warmup: the first call per workload is kept apart, it is the first-shape cost. started = time.perf_counter() @@ -140,13 +156,25 @@ def emit(record): } ) if i == 0: - emit({"type": "answers", "workload": w["id"], "answers": result["answers"]}) + emit( + { + "type": "answers", + "workload": w["id"], + "answers": result["answers"], + } + ) for w in parity: _, result = timed_call(agent, w) emit({"type": "answers", "workload": w["id"], "answers": result["answers"]}) - emit({"type": "end", "total_s": round(time.perf_counter() - T_START, 1), **memory(agent.device)}) + emit( + { + "type": "end", + "total_s": round(time.perf_counter() - T_START, 1), + **memory(agent.device), + } + ) print(out) diff --git a/benchmarks/laya_mps/env.py b/benchmarks/laya_mps/env.py index 8ddcdd3..8d33245 100644 --- a/benchmarks/laya_mps/env.py +++ b/benchmarks/laya_mps/env.py @@ -17,7 +17,9 @@ def _run(*cmd): try: - return subprocess.run(cmd, capture_output=True, text=True, timeout=10, check=True).stdout.strip() + return subprocess.run( + cmd, capture_output=True, text=True, timeout=10, check=True + ).stdout.strip() except (OSError, subprocess.SubprocessError): return None @@ -35,7 +37,9 @@ def _checkpoint_revision(repo_id, ref="main"): try: from huggingface_hub.constants import HF_HUB_CACHE - ref_file = Path(HF_HUB_CACHE) / f"models--{repo_id.replace('/', '--')}" / "refs" / ref + ref_file = ( + Path(HF_HUB_CACHE) / f"models--{repo_id.replace('/', '--')}" / "refs" / ref + ) return ref_file.read_text().strip() except (ImportError, OSError): return None @@ -66,13 +70,25 @@ def noise_problems(max_load): load = os.getloadavg()[0] if load > max_load: top = _run("ps", "-Ao", "pcpu=,comm=", "-r") or "" - busiest = "; ".join(" ".join(line.split()[:1] + [line.split("/")[-1]]) for line in top.splitlines()[:3]) + busiest = "; ".join( + " ".join(line.split()[:1] + [line.split("/")[-1]]) + for line in top.splitlines()[:3] + ) problems.append(f"1-min load {load:.1f} > {max_load} (busiest: {busiest})") return problems def header(checkpoint, **extra): - status = _run("git", "-C", str(REPO), "status", "--porcelain", "--", ".", ":!benchmarks/laya_mps/results") + status = _run( + "git", + "-C", + str(REPO), + "status", + "--porcelain", + "--", + ".", + ":!benchmarks/laya_mps/results", + ) return { "type": "env", "utc": datetime.now(timezone.utc).isoformat(timespec="seconds"), @@ -84,7 +100,9 @@ def header(checkpoint, **extra): "torch": _version("torch"), "transformers": _version("transformers"), "python": platform.python_version(), - "os": f"macOS {platform.mac_ver()[0]}" if sys.platform == "darwin" else platform.platform(), + "os": f"macOS {platform.mac_ver()[0]}" + if sys.platform == "darwin" + else platform.platform(), "chip": _run("sysctl", "-n", "machdep.cpu.brand_string"), "cpu_perf_cores": _run("sysctl", "-n", "hw.perflevel0.physicalcpu"), "cpu_eff_cores": _run("sysctl", "-n", "hw.perflevel1.physicalcpu"), @@ -151,7 +169,9 @@ class RusageInfoV4(ctypes.Structure): # , rusage_info_v4 return {} info = RusageInfoV4() libc = ctypes.CDLL("/usr/lib/libSystem.B.dylib", use_errno=True) - if libc.proc_pid_rusage(pid or os.getpid(), 4, ctypes.byref(info)) != 0: # RUSAGE_INFO_V4 + if ( + libc.proc_pid_rusage(pid or os.getpid(), 4, ctypes.byref(info)) != 0 + ): # RUSAGE_INFO_V4 return {} return { "footprint_mb": round(info.phys_footprint / 2**20), diff --git a/benchmarks/laya_mps/paired.py b/benchmarks/laya_mps/paired.py index dfed172..38ec721 100644 --- a/benchmarks/laya_mps/paired.py +++ b/benchmarks/laya_mps/paired.py @@ -48,7 +48,9 @@ def spawn(flags, port, python, model, log_path): *shlex.split(flags), ] log = open(log_path, "w") # noqa: SIM115 - return subprocess.Popen(command, env=env, stdout=log, stderr=subprocess.STDOUT, cwd=REPO) + return subprocess.Popen( + command, env=env, stdout=log, stderr=subprocess.STDOUT, cwd=REPO + ) def run(args): @@ -74,8 +76,19 @@ def run(args): try: if not (args.a_url or args.b_url): for s, (flags, port) in sides.items(): - procs[s] = spawn(flags, port, args.python, args.model, out_dir / f"paired_{args.run}_{s}.log") - health = {s: wait_ready(urls[s], {s: procs[s]} if s in procs else {}, args.ready_timeout)[1] for s in sides} + procs[s] = spawn( + flags, + port, + args.python, + args.model, + out_dir / f"paired_{args.run}_{s}.log", + ) + health = { + s: wait_ready( + urls[s], {s: procs[s]} if s in procs else {}, args.ready_timeout + )[1] + for s in sides + } with open(out_dir / f"paired_{args.run}.jsonl", "w") as f: def emit(record): @@ -98,7 +111,15 @@ def emit(record): body = body_for(w, args.model) for s in sides: answers, error = fetch_answers(clients[s], body) - emit({"type": "answers", "side": s, "workload": w["id"], "answers": answers, "error": error}) + emit( + { + "type": "answers", + "side": s, + "workload": w["id"], + "answers": answers, + "error": error, + } + ) order = bench[:] random.Random(args.seed).shuffle(order) for w in order: @@ -108,7 +129,9 @@ def emit(record): ms = {} for s in (first, second): try: - t, status, _ = clients[s].request("POST", "/v1/systemone", body) + t, status, _ = clients[s].request( + "POST", "/v1/systemone", body + ) except (OSError, http.client.HTTPException): t, status = None, 0 ms[s] = t if status == 200 else None @@ -124,12 +147,17 @@ def emit(record): "rows": len(w["questions"]), } ) - end_health = {s: json.loads(clients[s].request("GET", "/health", retry=True)[2]) for s in sides} + end_health = { + s: json.loads(clients[s].request("GET", "/health", retry=True)[2]) + for s in sides + } emit( { "type": "end", "health": end_health, - "footprint_mb": {s: footprint_mb(procs[s].pid).get("footprint_mb") for s in procs}, + "footprint_mb": { + s: footprint_mb(procs[s].pid).get("footprint_mb") for s in procs + }, } ) finally: @@ -145,8 +173,14 @@ def emit(record): def median_interval(ratios, seed=0, resamples=2000): rng = random.Random(seed) - meds = sorted(statistics.median(rng.choices(ratios, k=len(ratios))) for _ in range(resamples)) - return statistics.median(ratios), meds[int(0.025 * resamples)], meds[int(0.975 * resamples) - 1] + meds = sorted( + statistics.median(rng.choices(ratios, k=len(ratios))) for _ in range(resamples) + ) + return ( + statistics.median(ratios), + meds[int(0.025 * resamples)], + meds[int(0.975 * resamples) - 1], + ) def flat(answer): @@ -171,8 +205,12 @@ def summarize(paths): with open(path) as f: records = [json.loads(line) for line in f if line.strip()] env = next(r for r in records if r["type"] == "env") - print(f"## {env['run']}: A = `{env['a']}`, B = `{env['b']}`, load at start {env['loadavg_1m']}\n") - print("| input | pairs | A p50 ms | B p50 ms | median B/A | 95% interval |\n|---|---|---|---|---|---|") + print( + f"## {env['run']}: A = `{env['a']}`, B = `{env['b']}`, load at start {env['loadavg_1m']}\n" + ) + print( + "| input | pairs | A p50 ms | B p50 ms | median B/A | 95% interval |\n|---|---|---|---|---|---|" + ) pairs, failed = {}, 0 for r in records: if r["type"] == "pair" and r["a_ms"] and r["b_ms"]: @@ -201,23 +239,37 @@ def summarize(paths): if a is None or b is None or a.get("type") != b.get("type"): errors.append(f"{wid}/{q}") continue - worst = max(worst, max(abs(flat(a)[k] - flat(b).get(k, 0.0)) for k in flat(a))) + worst = max( + worst, max(abs(flat(a)[k] - flat(b).get(k, 0.0)) for k in flat(a)) + ) if decision(a) != decision(b): flips.append((wid, q, round(margin(a), 4))) - print(f"\nB vs A answers: max |Δp| {worst:.4f}, flips {flips}, errors {errors}; failed pairs: {failed}") + print( + f"\nB vs A answers: max |Δp| {worst:.4f}, flips {flips}, errors {errors}; failed pairs: {failed}" + ) end = next((r for r in records if r["type"] == "end"), None) if end is None: - print("**The run did not finish: no end record, so no device, recompile or memory check.**\n") + print( + "**The run did not finish: no end record, so no device, recompile or memory check.**\n" + ) continue - compile_state = {s: h.get("compile", {}).get("recompiled_after_ready") for s, h in end["health"].items()} + compile_state = { + s: h.get("compile", {}).get("recompiled_after_ready") + for s, h in end["health"].items() + } devices = {s: h.get("device") for s, h in end["health"].items()} off_gpu = [ s for s, h in end["health"].items() if h.get("device_mismatch") - or (h.get("compile", {}).get("enabled") and not h["compile"].get("active", True)) + or ( + h.get("compile", {}).get("enabled") + and not h["compile"].get("active", True) + ) ] - print(f"recompiled after ready: {compile_state}; device at end: {devices}; footprint MB: {end['footprint_mb']}") + print( + f"recompiled after ready: {compile_state}; device at end: {devices}; footprint MB: {end['footprint_mb']}" + ) if off_gpu: print( f"**Side {', '.join(off_gpu)} left its device or compiled path during the run; the ratios above mix both.**" @@ -229,9 +281,17 @@ def main(): parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) parser.add_argument("--summarize", nargs="+", metavar="JSONL") parser.add_argument("--run") - parser.add_argument("--a", help="extra frontend.laya_mps flags for side A (default: none)") - parser.add_argument("--b", help="extra frontend.laya_mps flags for side B (default: --compile --weights fp16)") - parser.add_argument("--a-url", help="instead of starting workers: an already running server for side A") + parser.add_argument( + "--a", help="extra frontend.laya_mps flags for side A (default: none)" + ) + parser.add_argument( + "--b", + help="extra frontend.laya_mps flags for side B (default: --compile --weights fp16)", + ) + parser.add_argument( + "--a-url", + help="instead of starting workers: an already running server for side A", + ) parser.add_argument("--b-url", help="... and for side B") parser.add_argument("--port-a", type=int, default=8000) parser.add_argument("--port-b", type=int, default=8001) diff --git a/benchmarks/laya_mps/profile_mps.py b/benchmarks/laya_mps/profile_mps.py index c7a8ff5..93f44c6 100644 --- a/benchmarks/laya_mps/profile_mps.py +++ b/benchmarks/laya_mps/profile_mps.py @@ -81,7 +81,10 @@ def forward(b): if agent.device.type == "mps": torch.mps.synchronize() t2 = time.perf_counter() - out = logits.float().cpu().numpy(), torch.softmax(act.float(), -1).cpu().numpy() + out = ( + logits.float().cpu().numpy(), + torch.softmax(act.float(), -1).cpu().numpy(), + ) t3 = time.perf_counter() self._add("dispatch", (t1 - t0) * 1000) self._add("gpu_wait", (t2 - t1) * 1000) @@ -100,7 +103,9 @@ def request(self, state, questions): torch.mps.synchronize() stages, self.current = self.current, None stages["wall"] = (time.perf_counter() - started) * 1000 - stages["other"] = stages["wall"] - sum(stages.get(s, 0.0) for s in STAGES if s != "other") + stages["other"] = stages["wall"] - sum( + stages.get(s, 0.0) for s in STAGES if s != "other" + ) return stages, result @@ -124,7 +129,9 @@ def main(): parser.add_argument("--stage-workloads", nargs="+", default=["W1", "W3", "W5"]) parser.add_argument("-n", type=int, default=100, help="timed requests per point") parser.add_argument("--discard", type=int, default=10) - parser.add_argument("--sweep-words", type=int, nargs="+", default=[1, 8, 24, 56, 120, 250, 380]) + parser.add_argument( + "--sweep-words", type=int, nargs="+", default=[1, 8, 24, 56, 120, 250, 380] + ) parser.add_argument("--ops-requests", type=int, default=20) parser.add_argument("--out", default=str(HERE / "results")) parser.add_argument("--max-load", type=float, default=2.0) @@ -137,7 +144,9 @@ def main(): print(f"warning: {problem}", file=sys.stderr) with open(args.workloads) as f: - workloads = {w["id"]: w for w in (json.loads(line) for line in f if line.strip())} + workloads = { + w["id"]: w for w in (json.loads(line) for line in f if line.strip()) + } route = workloads["W1"]["questions"] agent = laya.load(args.checkpoint, device=args.device) @@ -187,12 +196,24 @@ def emit(record): ) print(f"| {words} | {tokens} | {p50:.1f} |") a, b, r2 = fit([t for t, _ in points], [p for _, p in points]) - emit({"type": "fit", "a_ms": round(a, 3), "b_ms_per_token": round(b, 5), "r2": round(r2, 4)}) + emit( + { + "type": "fit", + "a_ms": round(a, 3), + "b_ms_per_token": round(b, 5), + "r2": round(r2, 4), + } + ) print(f"\nwall ≈ {a:.1f} ms + {b:.3f} ms/token × tokens (R² {r2:.3f})") print("\n## Stages (median ms per request)\n") columns = ["wall", *STAGES, "gpu_exec"] - print("| workload | tokens | rows | " + " | ".join(columns) + " |\n|" + "---|" * (len(columns) + 3)) + print( + "| workload | tokens | rows | " + + " | ".join(columns) + + " |\n|" + + "---|" * (len(columns) + 3) + ) for wid in args.stage_workloads: w = workloads[wid] for _ in range(args.discard): @@ -211,15 +232,21 @@ def emit(record): "tokens": tokens, "rows": len(w["questions"]), "median_ms": {k: round(v, 3) for k, v in medians.items()}, - "samples": {k: [round(x, 3) for x in v] for k, v in per_stage.items()}, + "samples": { + k: [round(x, 3) for x in v] for k, v in per_stage.items() + }, } ) - cells = " | ".join(f"{medians[c]:.1f}" if c in medians else "" for c in columns) + cells = " | ".join( + f"{medians[c]:.1f}" if c in medians else "" for c in columns + ) print(f"| {wid} | {tokens} | {len(w['questions'])} | {cells} |") print("\n## Host operators for W1 (torch.profiler, CPU)\n") w = workloads["W1"] - with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU]) as prof: + with torch.profiler.profile( + activities=[torch.profiler.ProfilerActivity.CPU] + ) as prof: for _ in range(args.ops_requests): timer.request(w["state"], w["questions"]) events = [e for e in prof.key_averages() if e.key.startswith("aten::")] @@ -235,7 +262,9 @@ def emit(record): { "op": e.key, "calls_per_request": e.count / args.ops_requests, - "self_cpu_ms_per_request": e.self_cpu_time_total / 1000 / args.ops_requests, + "self_cpu_ms_per_request": e.self_cpu_time_total + / 1000 + / args.ops_requests, } for e in top ], diff --git a/benchmarks/laya_mps/report.py b/benchmarks/laya_mps/report.py index 4ec1aa8..355ae9d 100644 --- a/benchmarks/laya_mps/report.py +++ b/benchmarks/laya_mps/report.py @@ -18,7 +18,12 @@ from collections import defaultdict GATE = 0.10 -PHASES = ["import_s", "load_s", "process_to_ready_s", "warmup_s"] # whichever a result file has +PHASES = [ + "import_s", + "load_s", + "process_to_ready_s", + "warmup_s", +] # whichever a result file has def percentile(sorted_values, p): @@ -30,8 +35,13 @@ def read(paths): for path in paths: with open(path) as f: rows = [json.loads(line) for line in f if line.strip()] - if any("config" not in r for r in rows): # e.g. paired.py results: summarize those with paired.py - print(f"skipping {path}: not a bench_inproc/bench_http result", file=sys.stderr) + if any( + "config" not in r for r in rows + ): # e.g. paired.py results: summarize those with paired.py + print( + f"skipping {path}: not a bench_inproc/bench_http result", + file=sys.stderr, + ) continue records.extend(rows) return records @@ -39,7 +49,10 @@ def read(paths): def table(headers, rows): lines = ["| " + " | ".join(headers) + " |", "|" + "---|" * len(headers)] - lines += ["| " + " | ".join("" if c is None else str(c) for c in row) + " |" for row in rows] + lines += [ + "| " + " | ".join("" if c is None else str(c) for c in row) + " |" + for row in rows + ] return "\n".join(lines) @@ -84,13 +97,24 @@ def environment(envs): def phase_table(phases): - present = [p for p in PHASES if any(ph.get(p) is not None for ph in phases.values())] + present = [ + p for p in PHASES if any(ph.get(p) is not None for ph in phases.values()) + ] workload_ids = sorted({w for ph in phases.values() for w in ph["first_ms"]}) rows = [ - [*k, *(ph.get(p) for p in present), *(ph["first_ms"].get(w) for w in workload_ids)] + [ + *k, + *(ph.get(p) for p in present), + *(ph["first_ms"].get(w) for w in workload_ids), + ] for k, ph in sorted(phases.items()) ] - headers = ["config", "run", *(p.removesuffix("_s") for p in present), *(f"first {w}" for w in workload_ids)] + headers = [ + "config", + "run", + *(p.removesuffix("_s") for p in present), + *(f"first {w}" for w in workload_ids), + ] return table(headers, rows) @@ -98,7 +122,9 @@ def latency(records): samples = defaultdict(list) for r in records: if r["type"] == "req" and r.get("status", 200) == 200: - samples[(r["config"], r["workload"], r.get("concurrency", 1), r["run"])].append(r["wall_ms"]) + samples[ + (r["config"], r["workload"], r.get("concurrency", 1), r["run"]) + ].append(r["wall_ms"]) rows, p50s = [], defaultdict(dict) for (config, workload, concurrency, run), values in sorted(samples.items()): values.sort() @@ -120,27 +146,68 @@ def latency(records): f"{cv:.1%}", ] ) - return table(["config", "workload", "conc", "run", "n", "p50", "p95", "mean", "CV"], rows), p50s + return table( + ["config", "workload", "conc", "run", "n", "p50", "p95", "mean", "CV"], rows + ), p50s def gate(p50s): rows = [] for (config, workload, concurrency), by_run in sorted(p50s.items()): if len(by_run) < 2: - rows.append([config, workload, concurrency, len(by_run), "", "needs 2 measured runs"]) + rows.append( + [ + config, + workload, + concurrency, + len(by_run), + "", + "needs 2 measured runs", + ] + ) continue spread = (max(by_run.values()) - min(by_run.values())) / min(by_run.values()) - rows.append([config, workload, concurrency, len(by_run), f"{spread:.1%}", "PASS" if spread <= GATE else "FAIL"]) + rows.append( + [ + config, + workload, + concurrency, + len(by_run), + f"{spread:.1%}", + "PASS" if spread <= GATE else "FAIL", + ] + ) return table(["config", "workload", "conc", "runs", "spread", "gate"], rows) def throughput(records): rows = [ - [r["config"], r["workload"], r["concurrency"], r["run"], r["n"], r["errors"], r["elapsed_s"], r["rps"]] + [ + r["config"], + r["workload"], + r["concurrency"], + r["run"], + r["n"], + r["errors"], + r["elapsed_s"], + r["rps"], + ] for r in records if r["type"] == "throughput" ] - return table(["config", "workload", "conc", "run", "n", "errors", "elapsed s", "successful req/s"], sorted(rows)) + return table( + [ + "config", + "workload", + "conc", + "run", + "n", + "errors", + "elapsed s", + "successful req/s", + ], + sorted(rows), + ) def memory(phases, ends): @@ -189,7 +256,9 @@ def read_answers(records): elif r["type"] == "answers": answers.setdefault(key, {})[r["workload"]] = r["answers"] elif r["type"] == "answers_error": - errors.setdefault(key, {})[r["workload"]] = f"status {r.get('status')}: {r.get('detail')}" + errors.setdefault(key, {})[r["workload"]] = ( + f"status {r.get('status')}: {r.get('detail')}" + ) return envs, answers, errors, benchmarks @@ -236,7 +305,9 @@ def parity(records, ref): "|---|---|---|---|---|---|---|---|---|---|", ] failed = total = 0 - for key in sorted(benchmarks | answers.keys() | errors.keys()): # a benchmark run without answers still counts + for key in sorted( + benchmarks | answers.keys() | errors.keys() + ): # a benchmark run without answers still counts if key == ref_key: continue env = envs.get(key, {}) @@ -244,7 +315,9 @@ def parity(records, ref): got = answers.get(key, {}).get(workload) if got is None: why = errors.get(key, {}).get(workload, "missing") - lines.append(f"| {key[0]} | {key[1]} | {workload} | | | | {why} | | | FAIL |") + lines.append( + f"| {key[0]} | {key[1]} | {workload} | | | | {why} | | | FAIL |" + ) failed += 1 total += 1 continue @@ -252,15 +325,24 @@ def parity(records, ref): for qid, ref_answer in sorted(questions.items()): total += 1 if qid not in got: - lines.append(f"| {key[0]} | {key[1]} | {workload} | {qid} | | | missing | | | FAIL |") + lines.append( + f"| {key[0]} | {key[1]} | {workload} | {qid} | | | missing | | | FAIL |" + ) failed += 1 continue ref_decision, ref_probs = outcome(ref_answer) decision, probs = outcome(got[qid]) - delta = max(abs(ref_probs.get(o, 0.0) - probs.get(o, 0.0)) for o in ref_probs.keys() | probs.keys()) + delta = max( + abs(ref_probs.get(o, 0.0) - probs.get(o, 0.0)) + for o in ref_probs.keys() | probs.keys() + ) ok = decision == ref_decision and delta <= TOLERANCE[precision] failed += not ok - flip = f"{ref_decision} → {decision}" if decision != ref_decision else f"{decision}" + flip = ( + f"{ref_decision} → {decision}" + if decision != ref_decision + else f"{decision}" + ) lines.append( f"| {key[0]} | {key[1]} | {workload} | {qid} | {ref_answer['type']} | {precision} | {flip} " f"| {delta:.4f} | {margin(ref_probs):.4f} | {'PASS' if ok else 'FAIL'} |" @@ -272,7 +354,11 @@ def parity(records, ref): def main(): parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) parser.add_argument("files", nargs="+") - parser.add_argument("--ref", default="C1", help="reference config for the parity section (default C1)") + parser.add_argument( + "--ref", + default="C1", + help="reference config for the parity section (default C1)", + ) args = parser.parse_args() records = read(args.files) @@ -282,9 +368,14 @@ def by_run(kind): envs, phases, ends = by_run("env"), by_run("phase"), by_run("end") latency_table, p50s = latency(records) print("## Environment\n\n" + environment(envs)) - print("\n## Phases (s) and first request per workload (ms)\n\n" + phase_table(phases)) + print( + "\n## Phases (s) and first request per workload (ms)\n\n" + phase_table(phases) + ) print("\n## Warm latency (ms)\n\n" + latency_table) - print(f"\n## Run-to-run gate (p50 spread across measured runs <= {GATE:.0%})\n\n" + gate(p50s)) + print( + f"\n## Run-to-run gate (p50 spread across measured runs <= {GATE:.0%})\n\n" + + gate(p50s) + ) if any(r["type"] == "throughput" for r in records): print("\n## Throughput\n\n" + throughput(records)) print("\n## Memory (MB)\n\n" + memory(phases, ends)) diff --git a/src/frontend/laya_mps.py b/src/frontend/laya_mps.py index c29eda4..48cc6e6 100644 --- a/src/frontend/laya_mps.py +++ b/src/frontend/laya_mps.py @@ -72,22 +72,33 @@ def build_app( nothing is preloaded. Raises if preparing fails, so the caller never binds a worker that cannot answer.""" from laya.serve import create_app - resident: dict[str, tuple[Any, dict[str, Any]]] = {} # prepared checkpoints: name -> (agent, warmup result) + resident: dict[ + str, tuple[Any, dict[str, Any]] + ] = {} # prepared checkpoints: name -> (agent, warmup result) preparing: set[str] = set() graphs_at_ready = None def apply_options(name: str, agent: Any) -> None: if (fp16 or compile) and not optimize.apply(agent, fp16=fp16, compile=compile): - log.warning("%s is on the CPU: --compile and --weights fp16 apply on the GPU only", name) + log.warning( + "%s is on the CPU: --compile and --weights fp16 apply on the GPU only", + name, + ) def make_ready(name: str, agent: Any) -> str | None: """Warm the checkpoint up and start describing it. Returns where it is if not on the requested device.""" nonlocal graphs_at_ready warmed = engine.warmup(router, name) - warmed["revision"] = engine.loaded_revision(revisions, warmed["routing"]) # fixed here: see engine + warmed["revision"] = engine.loaded_revision( + revisions, warmed["routing"] + ) # fixed here: see engine resident[name] = (agent, warmed) autocast_rows = getattr(agent, "mps_amp_min_rows", None) - if str(agent.device).startswith("mps") and autocast_rows and autocast_rows > engine.WARMUP_MAX_ROWS: + if ( + str(agent.device).startswith("mps") + and autocast_rows + and autocast_rows > engine.WARMUP_MAX_ROWS + ): log.warning( "%s: laya autocasts from %d questions but the warmup stops at %d; the first request that " "large is not warm", @@ -97,8 +108,14 @@ def make_ready(name: str, agent: Any) -> str | None: ) if compile: graphs_at_ready = graph_counter() - described = engine.describe(agent, requested, warmed["routing"], warmed["revision"]) - return f"{name} is on {described['device']}" if described["device_mismatch"] else None + described = engine.describe( + agent, requested, warmed["routing"], warmed["revision"] + ) + return ( + f"{name} is on {described['device']}" + if described["device_mismatch"] + else None + ) def check_device(misplaced: list[str]) -> None: if misplaced: @@ -107,7 +124,9 @@ def check_device(misplaced: list[str]) -> None: raise RuntimeError(message) log.warning(message) - evicted: list[str] = [] # what laya dropped to make room for the checkpoint it is loading + evicted: list[ + str + ] = [] # what laya dropped to make room for the checkpoint it is loading def on_evict(ctx: Any) -> None: resident.pop(ctx.model, None) @@ -131,18 +150,28 @@ def on_load(ctx: Any) -> None: try: router.load(name) except Exception: # noqa: BLE001 -- the request fails for the first reason either way - log.exception("%s was evicted for %s and could not be loaded again", name, ctx.model) + log.exception( + "%s was evicted for %s and could not be loaded again", + name, + ctx.model, + ) if compile: - graphs_at_ready = graph_counter() # graphs the unloaded checkpoint compiled are not recompiles + graphs_at_ready = ( + graph_counter() + ) # graphs the unloaded checkpoint compiled are not recompiles raise finally: preparing.discard(ctx.model) - names = list(router.loaded) or [model] # never load a model the worker was not asked to serve + names = list(router.loaded) or [ + model + ] # never load a model the worker was not asked to serve startup = {name: router.load(name) for name in names} for name, agent in startup.items(): apply_options(name, agent) - check_device(list(filter(None, [make_ready(name, agent) for name, agent in startup.items()]))) + check_device( + list(filter(None, [make_ready(name, agent) for name, agent in startup.items()])) + ) for hook in [h for h in getattr(router, "hooks", ()) if isinstance(h, Lifecycle)]: router.remove_hook(hook) # an app built earlier on this router router.add_hook(Lifecycle(on_load, on_evict)) @@ -151,7 +180,9 @@ def current() -> dict[str, Any]: """The agents as they are now, not as they were at startup (see the module docstring).""" models = { name: { - **engine.describe(agent, requested, warmed["routing"], warmed["revision"]), + **engine.describe( + agent, requested, warmed["routing"], warmed["revision"] + ), "warmup_ms": warmed["warmup_ms"], } for name, (agent, warmed) in list(resident.items()) @@ -164,7 +195,9 @@ def current() -> dict[str, Any]: } app = create_app(router) - app.router.routes[:] = [r for r in app.router.routes if getattr(r, "path", None) != "/health"] + app.router.routes[:] = [ + r for r in app.router.routes if getattr(r, "path", None) != "/health" + ] @app.get("/health") def health() -> dict[str, Any]: @@ -172,7 +205,10 @@ def health() -> dict[str, Any]: if compile: now = graph_counter() compiled.update( - active=any(optimize.compile_active(agent) for agent, _ in list(resident.values())), + active=any( + optimize.compile_active(agent) + for agent, _ in list(resident.values()) + ), graphs_at_ready=graphs_at_ready, graphs_now=now, recompiled_after_ready=now > graphs_at_ready and not preparing, @@ -199,11 +235,32 @@ def make_router(device: str | None, model: str) -> Any: def main() -> None: parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) - parser.add_argument("--device", default=None, help="torch device for laya: mps, cpu (default: laya's choice)") - parser.add_argument("--model", default="english", help="laya checkpoint to serve: english, multilingual, ...") - parser.add_argument("--compile", action="store_true", help="torch.compile the GPU path during warmup") - parser.add_argument("--weights", default="fp32", choices=["fp32", "fp16"], help="weight precision on the GPU") - parser.add_argument("--require-device", action="store_true", help="exit if a model is not on --device") + parser.add_argument( + "--device", + default=None, + help="torch device for laya: mps, cpu (default: laya's choice)", + ) + parser.add_argument( + "--model", + default="english", + help="laya checkpoint to serve: english, multilingual, ...", + ) + parser.add_argument( + "--compile", + action="store_true", + help="torch.compile the GPU path during warmup", + ) + parser.add_argument( + "--weights", + default="fp32", + choices=["fp32", "fp16"], + help="weight precision on the GPU", + ) + parser.add_argument( + "--require-device", + action="store_true", + help="exit if a model is not on --device", + ) parser.add_argument("--host", default="127.0.0.1") parser.add_argument("--port", type=int, default=8000) parser.add_argument( @@ -219,7 +276,10 @@ def main() -> None: logging.basicConfig(level=args.log_level.upper(), format="%(name)s: %(message)s") unread = [name for name in UNREAD_LAYA_SERVE_VARIABLES if os.environ.get(name)] if unread: - log.warning("%s: read by laya-serve's launcher, not by this worker; use the flags", ", ".join(unread)) + log.warning( + "%s: read by laya-serve's launcher, not by this worker; use the flags", + ", ".join(unread), + ) revisions = engine.record_snapshot_revisions() try: app = build_app( diff --git a/src/models/laya/engine.py b/src/models/laya/engine.py index 2ec8cd5..586b7a4 100644 --- a/src/models/laya/engine.py +++ b/src/models/laya/engine.py @@ -15,22 +15,38 @@ _CHOICE = { "type": "choice", "instructions": "Which team should handle this?", - "criteria": {"billing": "Charges and refunds", "technical": "Software problems", "other": "Anything else"}, + "criteria": { + "billing": "Charges and refunds", + "technical": "Software problems", + "other": "Anything else", + }, +} +_SCORE = { + "type": "score", + "instructions": "How urgent is it?", + "criteria": ["Low", "Medium", "High"], } -_SCORE = {"type": "score", "instructions": "How urgent is it?", "criteria": ["Low", "Medium", "High"]} _NOUL = {"type": "noul", "instructions": "Does the customer ask for a refund?"} WARMUP_SHAPES = [ (10, {"q": _CHOICE}), (150, {"q": _CHOICE}), (400, {"q": _CHOICE}), (10, {"a": _CHOICE, "b": _SCORE, "c": _NOUL}), - (10, {f"q{i}": q for i, q in enumerate([_CHOICE, _SCORE, _NOUL, _CHOICE, _SCORE, _NOUL])}), + ( + 10, + { + f"q{i}": q + for i, q in enumerate([_CHOICE, _SCORE, _NOUL, _CHOICE, _SCORE, _NOUL]) + }, + ), ] WARMUP_REPEATS = 2 WARMUP_MAX_ROWS = max(len(questions) for _, questions in WARMUP_SHAPES) -def warmup(router: Any, model: str, shapes=WARMUP_SHAPES, repeats: int = WARMUP_REPEATS) -> dict[str, Any]: +def warmup( + router: Any, model: str, shapes=WARMUP_SHAPES, repeats: int = WARMUP_REPEATS +) -> dict[str, Any]: """Any failure propagates: a worker that cannot answer must not bind.""" started = time.perf_counter() routing = None @@ -39,13 +55,20 @@ def warmup(router: Any, model: str, shapes=WARMUP_SHAPES, repeats: int = WARMUP_ for _ in range(repeats): result = router.predict(state, questions, model=model) routing = result.get("routing") or routing - return {"warmup_ms": round((time.perf_counter() - started) * 1000, 1), "routing": routing} + return { + "warmup_ms": round((time.perf_counter() - started) * 1000, 1), + "routing": routing, + } def _checkpoint_name(repo_id: str, allow_patterns: Any) -> str: """laya's name for what one download fetched: "", or "/" for a bundled checkpoint. laya restricts each download to one checkpoint's files, which all sit under its subfolder if it has one.""" - patterns = [allow_patterns] if isinstance(allow_patterns, str) else list(allow_patterns or [""]) + patterns = ( + [allow_patterns] + if isinstance(allow_patterns, str) + else list(allow_patterns or [""]) + ) folders = {pattern.split("/")[0] if "/" in pattern else "" for pattern in patterns} subfolder = folders.pop() if len(folders) == 1 else "" return f"{repo_id}/{subfolder}" if subfolder else repo_id @@ -78,13 +101,18 @@ def recording(repo_id, *args, **kwargs): return revisions -def loaded_revision(revisions: dict[str, str] | None, routing: dict[str, Any] | None) -> str | None: +def loaded_revision( + revisions: dict[str, str] | None, routing: dict[str, Any] | None +) -> str | None: """The commit of the checkpoint that has just loaded, if it was downloaded.""" return (revisions or {}).get((routing or {}).get("repo")) def describe( - agent: Any, requested: str | None, routing: dict[str, Any] | None, revision: str | None = None + agent: Any, + requested: str | None, + routing: dict[str, Any] | None, + revision: str | None = None, ) -> dict[str, Any]: """What /health reports about one loaded agent, read from the agent as it is now.""" device = str(getattr(agent, "device", "unknown")) @@ -98,7 +126,8 @@ def describe( return { "device": device, "requested_device": requested or "auto", - "device_mismatch": bool(requested_type) and device.split(":")[0] != requested_type, + "device_mismatch": bool(requested_type) + and device.split(":")[0] != requested_type, "weights_dtype": weights, "autocast_dtype": str(getattr(agent, "dtype", None)), "mps_amp_min_rows": getattr(agent, "mps_amp_min_rows", None), diff --git a/src/models/laya/optimize.py b/src/models/laya/optimize.py index 2b6d369..d02d4c0 100644 --- a/src/models/laya/optimize.py +++ b/src/models/laya/optimize.py @@ -32,7 +32,9 @@ def forward(self, input_ids, *args, **kwargs): if self.paths is None: return self.eager(input_ids, *args, **kwargs) whole, encoder_only = self.paths - return (whole if input_ids.shape[0] == 1 else encoder_only)(input_ids, *args, **kwargs) + return (whole if input_ids.shape[0] == 1 else encoder_only)( + input_ids, *args, **kwargs + ) return Served @@ -102,7 +104,9 @@ def apply(agent: Any, *, fp16: bool, compile: bool) -> bool: def compile_active(agent: Any) -> bool: """Whether requests to this agent run the compiled paths now. False on the CPU, also after a fallback.""" - return getattr(getattr(agent, "model", None), "paths", None) is not None and not _on_cpu(agent) + return getattr( + getattr(agent, "model", None), "paths", None + ) is not None and not _on_cpu(agent) def compiled_graphs() -> int: diff --git a/tests/laya/test_bench.py b/tests/laya/test_bench.py index 9baa7a1..5f276d1 100644 --- a/tests/laya/test_bench.py +++ b/tests/laya/test_bench.py @@ -18,8 +18,14 @@ def rows(path): def test_header_records_what_makes_two_runs_comparable(monkeypatch): record = bench_env.header("some/repo", extra=1) - assert record["type"] == "env" and record["extra"] == 1 and record["checkpoint"] == "some/repo" - head = subprocess.run(["git", "-C", str(REPO), "rev-parse", "HEAD"], capture_output=True, text=True).stdout.strip() + assert ( + record["type"] == "env" + and record["extra"] == 1 + and record["checkpoint"] == "some/repo" + ) + head = subprocess.run( + ["git", "-C", str(REPO), "rev-parse", "HEAD"], capture_output=True, text=True + ).stdout.strip() assert record["omni_sha"] == head and isinstance(record["omni_dirty"], bool) import laya # noqa: F401 from importlib.metadata import version @@ -27,9 +33,13 @@ def test_header_records_what_makes_two_runs_comparable(monkeypatch): assert (record["laya"], record["torch"], record["transformers"]) == tuple( version(package) for package in ("laya", "torch", "transformers") ) - assert record["argv"] == sys.argv and record["python"] == ".".join(map(str, sys.version_info[:3])) + assert record["argv"] == sys.argv and record["python"] == ".".join( + map(str, sys.version_info[:3]) + ) assert record["loadavg_1m"] >= 0 - if sys.platform == "darwin": # the machine probes use macOS tools; elsewhere they are None (test below) + if ( + sys.platform == "darwin" + ): # the machine probes use macOS tools; elsewhere they are None (test below) assert record["power"] and record["chip"] and record["mem_gb"] > 0 assert record["utc"].endswith("+00:00") @@ -39,18 +49,34 @@ def test_the_fixed_inputs_span_question_types_lengths_and_option_counts(): assert len({w["id"] for w in workloads}) == len(workloads) questions = [q for w in workloads for q in w["questions"].values()] assert {q["type"] for q in questions} == {"choice", "score", "noul"} - assert {len(q["criteria"]) for q in questions if q["type"] == "choice"} >= {2, 5, 10} + assert {len(q["criteria"]) for q in questions if q["type"] == "choice"} >= { + 2, + 5, + 10, + } bench = [w for w in workloads if w["kind"] == "bench"] words = sorted(len(w["state"].split()) for w in bench) assert ( words[0] < 20 and 100 < max(w for w in words if w < 200) and words[-1] > 300 ) # short, medium, near the window - assert {len(w["questions"]) for w in bench} >= {1, 3, 6} # below and above laya's autocast threshold of 5 rows + assert {len(w["questions"]) for w in bench} >= { + 1, + 3, + 6, + } # below and above laya's autocast threshold of 5 rows parity = [w for w in workloads if w["kind"] == "parity"] - assert {q["type"] for w in parity for q in w["questions"].values()} == {"choice", "score", "noul"} - - -DOCUMENTS = ["recipe/laya/apple-silicon.md", "benchmarks/laya_mps/README.md", "src/models/laya/README.md"] + assert {q["type"] for w in parity for q in w["questions"].values()} == { + "choice", + "score", + "noul", + } + + +DOCUMENTS = [ + "recipe/laya/apple-silicon.md", + "benchmarks/laya_mps/README.md", + "src/models/laya/README.md", +] SCRIPTS = { "frontend.laya_mps": "src/frontend/laya_mps.py", **{f"benchmarks/laya_mps/{name}.py": f"benchmarks/laya_mps/{name}.py" @@ -77,24 +103,48 @@ def test_documented_commands_use_flags_and_files_that_exist(): assert len(commands) >= 15 for document, command, name, source in commands: assert (REPO / source).exists(), f"{document}: {source}" - options = set(re.findall(r'add_argument\(\s*"(--?[a-z][a-z-]*)"', (REPO / source).read_text())) - ours = command.split("--spawn")[0] if name.endswith(".py") else command.split(name, 1)[1] + options = set( + re.findall( + r'add_argument\(\s*"(--?[a-z][a-z-]*)"', (REPO / source).read_text() + ) + ) + ours = ( + command.split("--spawn")[0] + if name.endswith(".py") + else command.split(name, 1)[1] + ) if name == "benchmarks/laya_mps/paired.py": - ours = re.sub(r'"[^"]*"', "", ours) # --a/--b carry flags of the worker, checked below + ours = re.sub( + r'"[^"]*"', "", ours + ) # --a/--b carry flags of the worker, checked below for worker_flags in re.findall(r'--[ab] "([^"]*)"', command): assert set(re.findall(r"--[a-z-]+", worker_flags)) <= set( - re.findall(r'add_argument\(\s*"(--[a-z-]+)"', (REPO / SCRIPTS["frontend.laya_mps"]).read_text()) + re.findall( + r'add_argument\(\s*"(--[a-z-]+)"', + (REPO / SCRIPTS["frontend.laya_mps"]).read_text(), + ) ), f"{document}: {command}" used = set(re.findall(r"(? 0 + assert ( + startup["first_health"]["ready"] is True + and startup["first_health"]["warmup_ms"] > 0 + ) def test_the_checkpoint_stays_loaded_across_requests(worker): @@ -142,7 +155,10 @@ def test_the_checkpoint_stays_loaded_across_requests(worker): for _ in range(5): assert decide(port, {"q": NOUL})[0] == 200 health = json.loads(call(port, "GET", "/health")[1]) - assert health["loaded"] == ["english"] and health["warmup_ms"] == startup["first_health"]["warmup_ms"] + assert ( + health["loaded"] == ["english"] + and health["warmup_ms"] == startup["first_health"]["warmup_ms"] + ) def test_first_request_after_ready_is_not_a_cold_start(worker): @@ -160,7 +176,10 @@ def test_first_request_after_ready_is_not_a_cold_start(worker): ({"q": CHOICE}, {"q": "choice"}), ({"q": SCORE}, {"q": "score"}), ({"q": NOUL}, {"q": "noul"}), - ({"a": CHOICE, "b": SCORE, "c": NOUL}, {"a": "choice", "b": "score", "c": "noul"}), + ( + {"a": CHOICE, "b": SCORE, "c": NOUL}, + {"a": "choice", "b": "score", "c": "noul"}, + ), ], ids=["choice", "score", "noul", "combined"], ) @@ -187,7 +206,13 @@ def test_decisions(worker, questions, kinds): @pytest.mark.parametrize( "questions", - [{"q": CHOICE}, {"q": SCORE}, {"q": NOUL}, {"a": CHOICE, "b": SCORE, "c": NOUL}, SIX], + [ + {"q": CHOICE}, + {"q": SCORE}, + {"q": NOUL}, + {"a": CHOICE, "b": SCORE, "c": NOUL}, + SIX, + ], ids=["choice", "score", "noul", "combined", "six-questions"], ) def test_answers_match_laya_itself(worker, reference, questions): @@ -198,7 +223,10 @@ def test_answers_match_laya_itself(worker, reference, questions): expected = reference.system_one(STATE, questions) assert set(expected) <= set(served) # the worker adds `routing`, it drops nothing assert served["usage"] == expected["usage"] - reduced = "fp16" in FLAGS or (DEVICE == "mps" and len(questions) >= startup["first_health"]["mps_amp_min_rows"]) + reduced = "fp16" in FLAGS or ( + DEVICE == "mps" + and len(questions) >= startup["first_health"]["mps_amp_min_rows"] + ) tolerance = 1e-2 if reduced else 1e-3 for qid, want in expected["answers"].items(): got = served["answers"][qid] @@ -210,12 +238,19 @@ def test_answers_match_laya_itself(worker, reference, questions): for option, p in want["probabilities"].items(): assert got["probabilities"][option] == pytest.approx(p, abs=tolerance) if want["type"] == "choice": - assert got["choice"] == want["choice"] == max(got["probabilities"], key=got["probabilities"].get) + assert ( + got["choice"] + == want["choice"] + == max(got["probabilities"], key=got["probabilities"].get) + ) def test_same_request_same_answer(worker): port, _ = worker - first, second = (json.loads(decide(port, {"a": CHOICE, "b": NOUL})[1])["answers"] for _ in range(2)) + first, second = ( + json.loads(decide(port, {"a": CHOICE, "b": NOUL})[1])["answers"] + for _ in range(2) + ) assert first == second @@ -225,16 +260,37 @@ def test_same_request_same_answer(worker): (b"{not json", True, TOKEN, 400), ({"model": "english", "state": STATE}, False, TOKEN, 400), ( - {"model": "english", "state": STATE, "questions": {"q": {"type": "bogus", "instructions": "?"}}}, + { + "model": "english", + "state": STATE, + "questions": {"q": {"type": "bogus", "instructions": "?"}}, + }, False, TOKEN, 422, ), (b"x" * (2 * 1024 * 1024 + 1), True, TOKEN, 413), - ({"model": "english", "state": STATE, "questions": {"q": NOUL}}, False, "wrong", 401), - ({"model": "english", "state": STATE, "questions": {"q": NOUL}}, False, None, 401), + ( + {"model": "english", "state": STATE, "questions": {"q": NOUL}}, + False, + "wrong", + 401, + ), + ( + {"model": "english", "state": STATE, "questions": {"q": NOUL}}, + False, + None, + 401, + ), + ], + ids=[ + "malformed-json", + "no-questions", + "bad-question", + "too-large", + "wrong-token", + "no-token", ], - ids=["malformed-json", "no-questions", "bad-question", "too-large", "wrong-token", "no-token"], ) def test_errors(worker, body, raw, token, expected): port, _ = worker diff --git a/tests/laya/test_worker.py b/tests/laya/test_worker.py index b8f81f7..43749ba 100644 --- a/tests/laya/test_worker.py +++ b/tests/laya/test_worker.py @@ -83,7 +83,9 @@ def test_warmup_covers_short_long_and_fp16_multi_question_shapes(): rows = {len(questions) for _, questions, _ in router.calls} assert min(words) <= 20 and max(words) >= 400 assert max(rows) >= router.agent.mps_amp_min_rows - assert {q["type"] for _, questions, _ in router.calls for q in questions.values()} == {"choice", "score", "noul"} + assert { + q["type"] for _, questions, _ in router.calls for q in questions.values() + } == {"choice", "score", "noul"} assert len(router.calls) == len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS assert all(model == "english" for _, _, model in router.calls) @@ -101,7 +103,9 @@ def test_warmup_failure_raises_and_no_app_is_built(): def test_health_reports_the_agent_device_not_the_requested_one(): router = FakeRouter(FakeAgent(device="cpu", dtype="torch.float32")) - health = TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + health = ( + TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + ) assert health["device"] == "cpu" assert health["requested_device"] == "mps" assert health["device_mismatch"] is True @@ -109,7 +113,11 @@ def test_health_reports_the_agent_device_not_the_requested_one(): def test_health_on_the_requested_device(): - health = TestClient(worker.build_app(FakeRouter(), "english", "mps")).get("/health").json() + health = ( + TestClient(worker.build_app(FakeRouter(), "english", "mps")) + .get("/health") + .json() + ) assert health["device"] == "mps" assert health["device_mismatch"] is False assert health["autocast_dtype"] == "torch.float16" @@ -119,7 +127,12 @@ def test_health_on_the_requested_device(): def test_device_index_is_not_a_mismatch(): router = FakeRouter(FakeAgent(device="cuda:0")) - assert TestClient(worker.build_app(router, "english", "cuda")).get("/health").json()["device_mismatch"] is False + assert ( + TestClient(worker.build_app(router, "english", "cuda")) + .get("/health") + .json()["device_mismatch"] + is False + ) def test_auto_device_is_never_a_mismatch(): @@ -137,7 +150,9 @@ def test_require_device_refuses_to_serve_on_another_device(): def test_only_one_health_route_remains(): app = worker.build_app(FakeRouter(), "english", "mps") - assert [r.path for r in app.router.routes if getattr(r, "path", None) == "/health"] == ["/health"] + assert [ + r.path for r in app.router.routes if getattr(r, "path", None) == "/health" + ] == ["/health"] def test_decisions_still_go_through_laya_serve(): @@ -146,7 +161,11 @@ def test_decisions_still_go_through_laya_serve(): before = len(router.calls) response = client.post( "/v1/systemone", - json={"model": "english", "state": "refund me", "questions": {"r": {"type": "noul", "instructions": "?"}}}, + json={ + "model": "english", + "state": "refund me", + "questions": {"r": {"type": "noul", "instructions": "?"}}, + }, ) assert response.status_code == 200 assert response.json()["answers"]["r"]["noul"] == 0.9 @@ -154,32 +173,57 @@ def test_decisions_still_go_through_laya_serve(): def test_main_exits_non_zero_when_warmup_fails(monkeypatch, caplog): - monkeypatch.setattr(worker, "make_router", lambda device, model: FakeRouter(fail_on_call=1)) + monkeypatch.setattr( + worker, "make_router", lambda device, model: FakeRouter(fail_on_call=1) + ) monkeypatch.setattr(sys, "argv", ["laya_mps", "--device", "mps"]) monkeypatch.setattr("uvicorn.run", lambda *a, **k: pytest.fail("must not bind")) - with caplog.at_level("ERROR", logger="laya-worker"), pytest.raises(SystemExit, match="not starting"): + with ( + caplog.at_level("ERROR", logger="laya-worker"), + pytest.raises(SystemExit, match="not starting"), + ): worker.main() - assert "startup failed" in caplog.text and "Traceback" in caplog.text and "out of memory" in caplog.text + assert ( + "startup failed" in caplog.text + and "Traceback" in caplog.text + and "out of memory" in caplog.text + ) def test_compile_wraps_the_model_before_warmup(monkeypatch): router = FakeRouter() order = [] - monkeypatch.setattr(optimize, "compile_agent", lambda agent: order.append((agent, len(router.calls)))) + monkeypatch.setattr( + optimize, + "compile_agent", + lambda agent: order.append((agent, len(router.calls))), + ) worker.build_app(router, "english", "mps", compile=True, graph_counter=lambda: 3) assert order == [(router.agent, 0)] def test_health_reports_compile_off_by_default(): - health = TestClient(worker.build_app(FakeRouter(), "english", "mps")).get("/health").json() + health = ( + TestClient(worker.build_app(FakeRouter(), "english", "mps")) + .get("/health") + .json() + ) assert health["compile"] == {"enabled": False} def test_health_flags_graphs_compiled_after_ready(monkeypatch): monkeypatch.setattr(optimize, "compile_agent", lambda agent: None) - graphs = iter([4, 4, 5]) # at readiness, first /health, second /health after a new shape compiled + graphs = iter( + [4, 4, 5] + ) # at readiness, first /health, second /health after a new shape compiled client = TestClient( - worker.build_app(FakeRouter(), "english", "mps", compile=True, graph_counter=lambda: next(graphs)) + worker.build_app( + FakeRouter(), + "english", + "mps", + compile=True, + graph_counter=lambda: next(graphs), + ) ) first = client.get("/health").json()["compile"] assert first == { @@ -227,7 +271,11 @@ def forward(self, *args): def fake_compile(module, dynamic): assert dynamic is True - return Stub("compiled encoder" if isinstance(module, Encoder) else "whole model compiled") + return Stub( + "compiled encoder" + if isinstance(module, Encoder) + else "whole model compiled" + ) monkeypatch.setattr(torch, "compile", fake_compile) agent = FakeAgent() @@ -237,14 +285,23 @@ def fake_compile(module, dynamic): assert agent.model(on_gpu(1)) == "whole model compiled" assert agent.model(on_gpu(3)) == ("head", "compiled encoder") assert agent.model(torch.zeros(1, 7)) == ("head", "eager encoder") - assert model.encoder(torch.zeros(1, 7)) == "eager encoder" # the original model is left as it was - assert len(list(agent.model.parameters())) == len(list(model.parameters())) # one set of weights + assert ( + model.encoder(torch.zeros(1, 7)) == "eager encoder" + ) # the original model is left as it was + assert len(list(agent.model.parameters())) == len( + list(model.parameters()) + ) # one set of weights def test_every_loaded_model_is_warmed_and_described(): - agents = {"english": FakeAgent(), "multilingual": FakeAgent(device="cpu", dtype="torch.float32")} + agents = { + "english": FakeAgent(), + "multilingual": FakeAgent(device="cpu", dtype="torch.float32"), + } router = FakeRouter(agents=agents) - health = TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + health = ( + TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + ) per_model = len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS assert [m for _, _, m in router.calls].count("multilingual") == per_model assert [m for _, _, m in router.calls].count("english") == per_model @@ -256,7 +313,9 @@ def test_every_loaded_model_is_warmed_and_described(): def test_a_model_that_is_not_preloaded_is_not_loaded_for_warmup(): router = FakeRouter(agents={"multilingual": FakeAgent()}) - health = TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + health = ( + TestClient(worker.build_app(router, "english", "mps")).get("/health").json() + ) assert "english" not in router.loads assert {m for _, _, m in router.calls} == {"multilingual"} assert set(health["models"]) == {"multilingual"} @@ -270,9 +329,15 @@ def test_nothing_preloaded_warms_the_worker_model(): def test_revision_comes_from_the_loaded_snapshot_not_a_guess(): revisions = {"convaiinnovations/laya": "55cf4c4"} - health = TestClient(worker.build_app(FakeRouter(), "english", "mps", revisions=revisions)).get("/health") + health = TestClient( + worker.build_app(FakeRouter(), "english", "mps", revisions=revisions) + ).get("/health") assert health.json()["revision"] == "55cf4c4" - unknown = TestClient(worker.build_app(FakeRouter(), "english", "mps")).get("/health").json() + unknown = ( + TestClient(worker.build_app(FakeRouter(), "english", "mps")) + .get("/health") + .json() + ) assert unknown["revision"] is None @@ -283,19 +348,29 @@ def test_record_snapshot_revisions_reads_the_downloaded_path(monkeypatch): "convaiinnovations/laya": "/cache/models--convaiinnovations--laya/snapshots/55cf4c4abc/multilingual", "/local/checkpoint": "/local/checkpoint", } - monkeypatch.setattr(huggingface_hub, "snapshot_download", lambda repo_id, **kwargs: paths[repo_id]) + monkeypatch.setattr( + huggingface_hub, "snapshot_download", lambda repo_id, **kwargs: paths[repo_id] + ) revisions = engine.record_snapshot_revisions() assert ( - huggingface_hub.snapshot_download("convaiinnovations/laya", allow_patterns=["*"]) + huggingface_hub.snapshot_download( + "convaiinnovations/laya", allow_patterns=["*"] + ) == paths["convaiinnovations/laya"] ) huggingface_hub.snapshot_download("/local/checkpoint") assert revisions == {"convaiinnovations/laya": "55cf4c4abc"} huggingface_hub.snapshot_download( - "convaiinnovations/laya", allow_patterns=["multilingual/model.safetensors", "multilingual/tokenizer/*"] + "convaiinnovations/laya", + allow_patterns=["multilingual/model.safetensors", "multilingual/tokenizer/*"], ) - huggingface_hub.snapshot_download("convaiinnovations/laya", allow_patterns=["model.safetensors", "tokenizer/*"]) - assert set(revisions) == {"convaiinnovations/laya", "convaiinnovations/laya/multilingual"} + huggingface_hub.snapshot_download( + "convaiinnovations/laya", allow_patterns=["model.safetensors", "tokenizer/*"] + ) + assert set(revisions) == { + "convaiinnovations/laya", + "convaiinnovations/laya/multilingual", + } def test_fp16_weights_keep_act_head_in_fp32(): @@ -318,24 +393,44 @@ def test_fp16_weights_are_applied_to_every_loaded_model_before_warmup(monkeypatc order = [] agents = {"english": FakeAgent(), "multilingual": FakeAgent()} router = FakeRouter(agents=agents) - monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: order.append((agent, len(router.calls)))) + monkeypatch.setattr( + optimize, + "use_fp16_weights", + lambda agent: order.append((agent, len(router.calls))), + ) worker.build_app(router, "english", "mps", fp16=True) assert order == [(agents["english"], 0), (agents["multilingual"], 0)] def test_options_are_not_applied_to_a_model_on_the_cpu(monkeypatch, caplog): applied = [] - monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: applied.append("fp16")) - monkeypatch.setattr(optimize, "compile_agent", lambda agent: applied.append("compile")) + monkeypatch.setattr( + optimize, "use_fp16_weights", lambda agent: applied.append("fp16") + ) + monkeypatch.setattr( + optimize, "compile_agent", lambda agent: applied.append("compile") + ) with caplog.at_level("WARNING", logger="laya-worker"): worker.build_app( - FakeRouter(FakeAgent(device="cpu")), "english", "cpu", fp16=True, compile=True, graph_counter=lambda: 0 + FakeRouter(FakeAgent(device="cpu")), + "english", + "cpu", + fp16=True, + compile=True, + graph_counter=lambda: 0, ) assert applied == [] assert "apply on the GPU only" in caplog.text caplog.clear() with caplog.at_level("WARNING", logger="laya-worker"): - worker.build_app(FakeRouter(), "english", "mps", fp16=True, compile=True, graph_counter=lambda: 0) + worker.build_app( + FakeRouter(), + "english", + "mps", + fp16=True, + compile=True, + graph_counter=lambda: 0, + ) assert applied == ["fp16", "compile"] assert "GPU only" not in caplog.text @@ -352,7 +447,9 @@ def __init__(self): def forward(self, input_ids): return ("eager", self.encoder.weight.dtype) - monkeypatch.setattr(torch, "compile", lambda module, dynamic: lambda *a, **k: "compiled") + monkeypatch.setattr( + torch, "compile", lambda module, dynamic: lambda *a, **k: "compiled" + ) agent = FakeAgent() agent.model = Model() optimize.use_fp16_weights(agent) @@ -366,7 +463,9 @@ def forward(self, input_ids): def test_health_follows_a_fallback_to_cpu_after_startup(): agent = FakeAgent(device="mps") - client = TestClient(worker.build_app(FakeRouter(agent), "english", "mps", require_device=True)) + client = TestClient( + worker.build_app(FakeRouter(agent), "english", "mps", require_device=True) + ) assert client.get("/health").json()["device_mismatch"] is False agent.device = "cpu" agent.dtype = "torch.float32" @@ -381,13 +480,21 @@ def test_health_compile_active_follows_a_fallback_to_cpu(monkeypatch): import torch compiles = [] - monkeypatch.setattr(torch, "compile", lambda module, dynamic: compiles.append(module) or module) + monkeypatch.setattr( + torch, "compile", lambda module, dynamic: compiles.append(module) or module + ) agent = FakeAgent() agent.model = torch.nn.Sequential() agent.model.encoder = torch.nn.Identity() - client = TestClient(worker.build_app(FakeRouter(agent), "english", "mps", compile=True, graph_counter=lambda: 2)) + client = TestClient( + worker.build_app( + FakeRouter(agent), "english", "mps", compile=True, graph_counter=lambda: 2 + ) + ) assert client.get("/health").json()["compile"]["active"] is True - optimize.compile_agent(agent) # a second name for the same agent must not compile again + optimize.compile_agent( + agent + ) # a second name for the same agent must not compile again assert len(compiles) == 2 agent.device = "cpu" assert client.get("/health").json()["compile"]["active"] is False @@ -410,21 +517,30 @@ def test_log_level_applies_to_the_workers_own_log(monkeypatch): monkeypatch.setattr(worker, "make_router", lambda device, model: FakeRouter()) monkeypatch.setattr(worker.logging, "basicConfig", lambda **kw: seen.update(kw)) monkeypatch.setattr("uvicorn.run", lambda app, **kw: seen.update(uvicorn=kw)) - monkeypatch.setattr(sys, "argv", ["laya_mps", "--device", "mps", "--log-level", "warning"]) + monkeypatch.setattr( + sys, "argv", ["laya_mps", "--device", "mps", "--log-level", "warning"] + ) worker.main() assert (seen["level"], seen["uvicorn"]["log_level"]) == ("WARNING", "warning") - assert (seen["uvicorn"]["host"], seen["uvicorn"]["port"]) == ("127.0.0.1", 8000) # local only by default + assert (seen["uvicorn"]["host"], seen["uvicorn"]["port"]) == ( + "127.0.0.1", + 8000, + ) # local only by default def test_a_checkpoint_loaded_while_serving_is_prepared_and_described(monkeypatch): applied = [] - monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: applied.append(agent)) + monkeypatch.setattr( + optimize, "use_fp16_weights", lambda agent: applied.append(agent) + ) router = FakeRouter() client = TestClient(worker.build_app(router, "english", "mps", fp16=True)) before = len(router.calls) late = FakeAgent(device="cpu", dtype="torch.float32") router.load_while_serving("multilingual", late) - assert applied == [router.agent] # the late one is on the CPU, where the options do not apply + assert applied == [ + router.agent + ] # the late one is on the CPU, where the options do not apply assert [m for _, _, m in router.calls[before:]] == ["multilingual"] * ( len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS ) @@ -456,16 +572,30 @@ def test_compile_baseline_moves_when_a_late_checkpoint_compiles(monkeypatch): monkeypatch.setattr(optimize, "compile_agent", lambda agent: None) graphs = iter([3, 7, 7]) # startup, after the late checkpoint compiled, /health router = FakeRouter() - client = TestClient(worker.build_app(router, "english", "mps", compile=True, graph_counter=lambda: next(graphs))) + client = TestClient( + worker.build_app( + router, "english", "mps", compile=True, graph_counter=lambda: next(graphs) + ) + ) router.load_while_serving("multilingual", FakeAgent()) compiled = client.get("/health").json()["compile"] - assert (compiled["graphs_at_ready"], compiled["recompiled_after_ready"]) == (7, False) + assert (compiled["graphs_at_ready"], compiled["recompiled_after_ready"]) == ( + 7, + False, + ) def test_startup_error_names_every_checkpoint_off_the_requested_device(): - agents = {"english": FakeAgent(device="cpu"), "multilingual": FakeAgent(device="cpu")} - with pytest.raises(RuntimeError, match="asked for mps, english is on cpu, multilingual is on cpu"): - worker.build_app(FakeRouter(agents=agents), "english", "mps", require_device=True) + agents = { + "english": FakeAgent(device="cpu"), + "multilingual": FakeAgent(device="cpu"), + } + with pytest.raises( + RuntimeError, match="asked for mps, english is on cpu, multilingual is on cpu" + ): + worker.build_app( + FakeRouter(agents=agents), "english", "mps", require_device=True + ) def test_health_names_a_checkpoint_while_it_is_being_prepared(): @@ -495,7 +625,11 @@ def __init__(self, repo, device=None, token=None, subfolder=None): self.mps_amp_min_rows = 5 def system_one(self, state, questions, lang=None, **_): - return {"model": "stub", "answers": {qid: ANSWER for qid in questions}, "usage": {}} + return { + "model": "stub", + "answers": {qid: ANSWER for qid in questions}, + "usage": {}, + } def test_a_failed_late_load_gives_back_the_checkpoint_laya_evicted_for_it(monkeypatch): @@ -508,7 +642,11 @@ def test_a_failed_late_load_gives_back_the_checkpoint_laya_evicted_for_it(monkey router.preload(["english", "multilingual"]) # laya keeps two checkpoints by default client = TestClient(worker.build_app(router, "english", "mps", require_device=True)) with pytest.raises(RuntimeError, match="typed-decisions is on cpu"): - router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") + router.predict( + "refund me", + {"r": {"type": "noul", "instructions": "?"}}, + model="typed-decisions", + ) assert sorted(router.loaded) == ["english", "multilingual"] health = client.get("/health").json() assert set(health["models"]) == {"english", "multilingual"} @@ -522,7 +660,9 @@ def test_building_a_second_app_on_a_router_replaces_the_first_apps_hooks(): assert len(router.hooks) == 1 before = len(router.calls) router.load_while_serving("multilingual", FakeAgent()) - assert len(router.calls) - before == len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS + assert ( + len(router.calls) - before == len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS + ) def test_main_warns_about_laya_serve_variables_it_does_not_read(monkeypatch, caplog): @@ -538,7 +678,9 @@ def test_main_warns_about_laya_serve_variables_it_does_not_read(monkeypatch, cap assert "LAYA_API_KEY" not in caplog.text -def test_graphs_compiled_by_a_late_checkpoint_that_fails_are_not_reported_as_recompiles(monkeypatch): +def test_graphs_compiled_by_a_late_checkpoint_that_fails_are_not_reported_as_recompiles( + monkeypatch, +): monkeypatch.setattr(optimize, "compile_agent", lambda agent: None) router = FakeRouter() graphs = {"n": 0} @@ -549,7 +691,11 @@ def predict_and_compile(state, questions, model=None): return predict(state, questions, model=model) router.predict = predict_and_compile - client = TestClient(worker.build_app(router, "english", "mps", compile=True, graph_counter=lambda: graphs["n"])) + client = TestClient( + worker.build_app( + router, "english", "mps", compile=True, graph_counter=lambda: graphs["n"] + ) + ) router.fail_on_call = len(router.calls) + 3 with pytest.raises(RuntimeError, match="out of memory"): router.load_while_serving("multilingual", FakeAgent()) @@ -572,7 +718,9 @@ def __init__(self, repo, device=None, token=None, subfolder=None): prefix = f"{subfolder}/" if subfolder else "" path = huggingface_hub.snapshot_download( - repo, token=token, allow_patterns=[prefix + "model.safetensors", prefix + "tokenizer/*"] + repo, + token=token, + allow_patterns=[prefix + "model.safetensors", prefix + "tokenizer/*"], ) self.loaded_commit = Path(path).name self.subfolder = subfolder @@ -583,7 +731,11 @@ def __init__(self, repo, device=None, token=None, subfolder=None): def system_one(self, state, questions, lang=None, **_): if self.subfolder in Repository.broken: raise RuntimeError("MPS backend out of memory") - return {"model": "stub", "answers": {qid: ANSWER for qid in questions}, "usage": {}} + return { + "model": "stub", + "answers": {qid: ANSWER for qid in questions}, + "usage": {}, + } @pytest.fixture @@ -597,7 +749,9 @@ def laya_router(monkeypatch): monkeypatch.setattr(Repository, "device", {}) monkeypatch.setattr(Repository, "broken", set()) monkeypatch.setattr( - huggingface_hub, "snapshot_download", lambda repo, **kw: f"/hf/snapshots/commit-{Repository.commit}" + huggingface_hub, + "snapshot_download", + lambda repo, **kw: f"/hf/snapshots/commit-{Repository.commit}", ) monkeypatch.setattr(laya.agent, "Agent", DownloadingAgent) return Router(device="mps"), engine.record_snapshot_revisions() @@ -611,7 +765,9 @@ def assert_health_matches(client, router): for name, agent in agents.items(): assert health["models"][name]["device"] == str(agent.device) assert health["models"][name]["revision"] == agent.loaded_commit - assert health["device_mismatch"] == any(str(agent.device) != "mps" for agent in agents.values()) + assert health["device_mismatch"] == any( + str(agent.device) != "mps" for agent in agents.values() + ) def test_checkpoints_of_one_repository_keep_the_revision_of_their_own_load(laya_router): @@ -622,22 +778,38 @@ def test_checkpoints_of_one_repository_keep_the_revision_of_their_own_load(laya_ client = TestClient(worker.build_app(router, "english", "mps", revisions=revisions)) assert_health_matches(client, router) Repository.commit = 2 - router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") + router.predict( + "refund me", + {"r": {"type": "noul", "instructions": "?"}}, + model="typed-decisions", + ) assert_health_matches(client, router) @pytest.mark.parametrize("require_device", [False, True]) @pytest.mark.parametrize("seed", range(8)) -def test_health_matches_the_router_after_any_sequence_of_loads(laya_router, seed, require_device): +def test_health_matches_the_router_after_any_sequence_of_loads( + laya_router, seed, require_device +): import random router, revisions = laya_router names = ["english", "multilingual", "typed-decisions"] - subfolder = {"english": None, "multilingual": "multilingual", "typed-decisions": "typed-decisions"} + subfolder = { + "english": None, + "multilingual": "multilingual", + "typed-decisions": "typed-decisions", + } rng = random.Random(seed) router.preload(rng.sample(names, rng.choice([1, 2]))) client = TestClient( - worker.build_app(router, router.loaded[0], "mps", require_device=require_device, revisions=revisions) + worker.build_app( + router, + router.loaded[0], + "mps", + require_device=require_device, + revisions=revisions, + ) ) assert_health_matches(client, router) for _ in range(12): @@ -645,7 +817,11 @@ def test_health_matches_the_router_after_any_sequence_of_loads(laya_router, seed Repository.device = {subfolder[n]: "cpu" for n in names if rng.random() < 0.25} Repository.broken = {subfolder[n] for n in names if rng.random() < 0.15} try: - router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model=rng.choice(names)) + router.predict( + "refund me", + {"r": {"type": "noul", "instructions": "?"}}, + model=rng.choice(names), + ) except RuntimeError: pass # a failed late load or forward: the request fails, /health must still match the router Repository.broken = set() @@ -660,12 +836,24 @@ def test_health_answers_with_no_resident_checkpoint(laya_router): assert client.get("/health").json()["models"] == {} -def test_a_checkpoint_loaded_while_serving_gets_the_gpu_options_before_its_warmup(monkeypatch): +def test_a_checkpoint_loaded_while_serving_gets_the_gpu_options_before_its_warmup( + monkeypatch, +): applied = [] router = FakeRouter() - monkeypatch.setattr(optimize, "use_fp16_weights", lambda agent: applied.append(("fp16", len(router.calls)))) - monkeypatch.setattr(optimize, "compile_agent", lambda agent: applied.append(("compile", len(router.calls)))) - worker.build_app(router, "english", "mps", fp16=True, compile=True, graph_counter=lambda: 0) + monkeypatch.setattr( + optimize, + "use_fp16_weights", + lambda agent: applied.append(("fp16", len(router.calls))), + ) + monkeypatch.setattr( + optimize, + "compile_agent", + lambda agent: applied.append(("compile", len(router.calls))), + ) + worker.build_app( + router, "english", "mps", fp16=True, compile=True, graph_counter=lambda: 0 + ) warmed_at_startup = len(router.calls) router.load_while_serving("multilingual", FakeAgent()) assert applied[2:] == [("fp16", warmed_at_startup), ("compile", warmed_at_startup)] @@ -689,39 +877,65 @@ def test_failed_late_loads_are_logged_with_their_cause(laya_router, caplog): router, revisions = laya_router router.preload(["english", "multilingual"]) worker.build_app(router, "english", "mps", require_device=True, revisions=revisions) - Repository.device = {"typed-decisions": "cpu", None: "cpu"} # english cannot come back either + Repository.device = { + "typed-decisions": "cpu", + None: "cpu", + } # english cannot come back either with caplog.at_level("ERROR", logger="laya-worker"), pytest.raises(RuntimeError): - router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") + router.predict( + "refund me", + {"r": {"type": "noul", "instructions": "?"}}, + model="typed-decisions", + ) assert "typed-decisions could not be prepared and is unloaded" in caplog.text - assert "english was evicted for typed-decisions and could not be loaded again" in caplog.text - assert "asked for mps, typed-decisions is on cpu" in caplog.text # the traceback of the cause + assert ( + "english was evicted for typed-decisions and could not be loaded again" + in caplog.text + ) + assert ( + "asked for mps, typed-decisions is on cpu" in caplog.text + ) # the traceback of the cause assert router.loaded == ["multilingual"] -def test_a_failed_late_load_reloads_only_what_was_evicted_for_it(laya_router, monkeypatch): +def test_a_failed_late_load_reloads_only_what_was_evicted_for_it( + laya_router, monkeypatch +): router, revisions = laya_router router.max_loaded = 1 router.load("english") worker.build_app(router, "english", "mps", require_device=True, revisions=revisions) - router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="multilingual") # evicts english + router.predict( + "refund me", {"r": {"type": "noul", "instructions": "?"}}, model="multilingual" + ) # evicts english built = [] original = DownloadingAgent.__init__ monkeypatch.setattr( DownloadingAgent, "__init__", - lambda self, repo, **kw: built.append(kw.get("subfolder")) or original(self, repo, **kw), + lambda self, repo, **kw: ( + built.append(kw.get("subfolder")) or original(self, repo, **kw) + ), ) Repository.device = {"typed-decisions": "cpu"} with pytest.raises(RuntimeError): - router.predict("refund me", {"r": {"type": "noul", "instructions": "?"}}, model="typed-decisions") - assert built == ["typed-decisions", "multilingual"] and router.loaded == ["multilingual"] + router.predict( + "refund me", + {"r": {"type": "noul", "instructions": "?"}}, + model="typed-decisions", + ) + assert built == ["typed-decisions", "multilingual"] and router.loaded == [ + "multilingual" + ] def test_describe_reads_the_weight_dtype_from_the_model_as_it_is(): import torch agent = FakeAgent() - assert engine.describe(agent, "mps", None)["weights_dtype"] is None # no model to read + assert ( + engine.describe(agent, "mps", None)["weights_dtype"] is None + ) # no model to read agent.model = torch.nn.Linear(2, 2) assert engine.describe(agent, "mps", None)["weights_dtype"] == "torch.float32" agent.model.half() @@ -740,16 +954,24 @@ def test_apply_says_whether_the_options_were_applied(monkeypatch): assert optimize.apply(FakeAgent(device="mps"), fp16=True, compile=False) is True -def test_the_checkpoint_is_built_once_and_serves_every_request(laya_router, monkeypatch): +def test_the_checkpoint_is_built_once_and_serves_every_request( + laya_router, monkeypatch +): router, revisions = laya_router router.preload(["english"]) built = [] original = DownloadingAgent.__init__ monkeypatch.setattr( - DownloadingAgent, "__init__", lambda self, repo, **kw: built.append(repo) or original(self, repo, **kw) + DownloadingAgent, + "__init__", + lambda self, repo, **kw: built.append(repo) or original(self, repo, **kw), ) client = TestClient(worker.build_app(router, "english", "mps", revisions=revisions)) - body = {"model": "english", "state": "refund me", "questions": {"r": {"type": "noul", "instructions": "?"}}} + body = { + "model": "english", + "state": "refund me", + "questions": {"r": {"type": "noul", "instructions": "?"}}, + } for _ in range(3): response = client.post("/v1/systemone", json=body) assert response.status_code == 200 and response.json()["answers"]["r"] == ANSWER @@ -761,12 +983,16 @@ def test_main_binds_only_after_every_warmup_request(monkeypatch): bound_after = [] monkeypatch.setattr(worker, "make_router", lambda device, model: router) monkeypatch.setattr(worker.logging, "basicConfig", lambda **kw: None) - monkeypatch.setattr("uvicorn.run", lambda app, **kw: bound_after.append(len(router.calls))) + monkeypatch.setattr( + "uvicorn.run", lambda app, **kw: bound_after.append(len(router.calls)) + ) monkeypatch.setattr(sys, "argv", ["laya_mps", "--device", "mps"]) worker.main() assert bound_after == [len(engine.WARMUP_SHAPES) * engine.WARMUP_REPEATS] def test_the_warmup_asks_every_question_type(): - kinds = {q["type"] for _, questions in engine.WARMUP_SHAPES for q in questions.values()} + kinds = { + q["type"] for _, questions in engine.WARMUP_SHAPES for q in questions.values() + } assert kinds == {"choice", "score", "noul"} From c80d7ebeff5235323390d1acde80e3582d713435 Mon Sep 17 00:00:00 2001 From: cacheline999 <326908201+cacheline999@users.noreply.github.com> Date: Sat, 3 Oct 2026 15:52:41 +0800 Subject: [PATCH 2/4] Laya: readiness test checks success, not latency; pin the test tools; M4 in the recipe - The contract test of the first request after ready asserts that it succeeds with a valid answer. Its latency depends on how long the GPU has been idle and on the Mac (an M5 answered a first request without warmup in 132-143 ms), so it is left to the benchmarks; that the warmup ran before ready is tested separately. - requirements-mps.txt pins pytest, httpx2 and ruff; the recipe's Test section runs the ruff checks. - The recipe lists the reviewer's M4 among the Macs the tests ran on. - Test docstring points at tests/laya. --- recipe/laya/apple-silicon.md | 11 ++++++++--- recipe/laya/requirements-mps.txt | 4 ++++ tests/laya/test_contract.py | 16 +++++++++------- tests/laya/test_worker.py | 2 +- 4 files changed, 22 insertions(+), 11 deletions(-) diff --git a/recipe/laya/apple-silicon.md b/recipe/laya/apple-silicon.md index 1344ea1..e453655 100644 --- a/recipe/laya/apple-silicon.md +++ b/recipe/laya/apple-silicon.md @@ -7,7 +7,9 @@ runs the benchmarks. The model-side code is in [`src/models/laya/`](../../src/mo Validated on an M1 Pro (16 GB, 16-core GPU), macOS 26.1, Python 3.12, `laya[serve]==0.3.20`, torch 2.14.0 and the `english` checkpoint (`convaiinnovations/laya` at `55cf4c4`), and by another -contributor on an M5 (10-core GPU, 32 GB, macOS 26.5.2). Other M-series Macs have not been tested. +contributor on an M5 (10-core GPU, 32 GB, macOS 26.5.2). A reviewer ran the tests on an M4 (10-core GPU, +16 GB, macOS 26, Python 3.13), including the contract tests on the GPU. Other M-series Macs have not been +tested. Run all commands from the repository root. @@ -142,12 +144,15 @@ The frontend forwards the worker's response unchanged; `compare_with_backend.py` ## Test -The tests need `pytest` and `httpx2` (Starlette's `TestClient`; `httpx` works with a deprecation warning): +The tests need `pytest` and `httpx2` (Starlette's `TestClient`; `httpx` works with a deprecation warning); +`requirements-mps.txt` pins them and ruff: ```sh -.venv/bin/python -m pip install pytest httpx2 +.venv/bin/python -m pip install -r recipe/laya/requirements-mps.txt PYTHONPATH=src .venv/bin/python -m pytest tests/laya # unit tests, no model LAYA_CONTRACT=1 PYTHONPATH=src .venv/bin/python -m pytest tests/laya # plus contract tests against a CPU worker +.venv/bin/ruff format --check src/frontend/laya_mps.py src/models/laya benchmarks/laya_mps tests/laya +.venv/bin/ruff check --select E4,E7,E9,F src/frontend/laya_mps.py src/models/laya benchmarks/laya_mps tests/laya ``` The contract tests start a real worker and check readiness, the three decision types, error responses, and diff --git a/recipe/laya/requirements-mps.txt b/recipe/laya/requirements-mps.txt index dca8e7a..0fe46fd 100644 --- a/recipe/laya/requirements-mps.txt +++ b/recipe/laya/requirements-mps.txt @@ -8,3 +8,7 @@ safetensors==0.8.0 huggingface-hub==1.33.0 fastapi==0.141.1 uvicorn==0.54.0 +# For the tests and checks in the recipe's Test section. +pytest==9.1.1 +httpx2==2.13.1 +ruff==0.16.10 diff --git a/tests/laya/test_contract.py b/tests/laya/test_contract.py index 4ea9ee2..9ef6017 100644 --- a/tests/laya/test_contract.py +++ b/tests/laya/test_contract.py @@ -13,7 +13,6 @@ import os import shlex import socket -import statistics import subprocess import sys import time @@ -161,13 +160,16 @@ def test_the_checkpoint_stays_loaded_across_requests(worker): ) -def test_first_request_after_ready_is_not_a_cold_start(worker): - port, (status, _, first_ms) = worker +def test_the_first_request_after_ready_succeeds(worker): + # Its latency is measured by the benchmarks, not asserted here: on MPS it depends on + # how long the GPU has been idle since the warmup (tens to over a hundred ms on an + # M1 Pro), and on an M5 even a worker without warmup answered its first request in + # 132-143 ms, so no bound separates warm from cold on every Mac. That the warmup ran + # before ready is checked above. + _, (status, body, _) = worker assert status == 200 - warm = statistics.median(decide(port, {"q": CHOICE})[2] for _ in range(20)) - # Without the warmup the first request costs several hundred ms more than a warm one. A few tens of ms - # remain on MPS: the GPU has been idle since the warmup, and any request after a pause pays that. - assert first_ms <= warm + 100, f"first {first_ms:.0f} ms, warm p50 {warm:.0f} ms" + answer = json.loads(body)["answers"]["q"] + assert answer["type"] == "choice" and answer["choice"] in CHOICE["criteria"] @pytest.mark.parametrize( diff --git a/tests/laya/test_worker.py b/tests/laya/test_worker.py index 43749ba..3c86061 100644 --- a/tests/laya/test_worker.py +++ b/tests/laya/test_worker.py @@ -1,6 +1,6 @@ """Unit tests for the Laya worker. A fake Router stands in for laya's; no model is loaded. -python -m pytest src/models/laya/tests +PYTHONPATH=src python -m pytest tests/laya """ import sys From 744b3b638a11ea89c6e98e7ef93211d0aedf4c98 Mon Sep 17 00:00:00 2001 From: cacheline999 <326908201+cacheline999@users.noreply.github.com> Date: Sun, 4 Oct 2026 17:09:54 +0800 Subject: [PATCH 3/4] Laya: a reproduction command for every number in the Apple Silicon recipe The benchmark README now has a table from each recipe claim to the command behind it. New scripts cover what #30 measured with one-off probes: - paired.py --gap SECONDS (each request after that much idle, on a fresh connection) and --only; - lengths.py: the first request of new input lengths and the footprint they add, for one or two flag sets; it skips the lengths the warmup ran and stops at the window; - late_load.py: a checkpoint loaded while serving, directly and through the frontend; - fallback.py: Laya's fallback to the CPU under a lowered MPS memory limit; - release.py: in one process prepared by the worker's build_app, what torch.mps.empty_cache() gives back after new lengths. It releases once before the walk and reads every footprint --settle seconds after a release, because the footprint shows a release up to ~2 s late and the starting footprint varies by a few hundred MB between runs; where a release lands does not; - a loop of fresh starts alternating the worker with plain laya-serve. Every measuring script refuses a run on battery or above --max-load (late_load, fallback and release unless --feasibility); env.py holds the shared check and workload reader. Recipe changes from the reruns: a late load with --compile took 19-22 s (the ~70 s in #30 was not reproduced, so the recipe gives no number for a load past the frontend's 60 s); the CPU fallback request took 30-73 s; every length up to the window is 454 lengths, 5.2 GB with the options and 3.7 GB without; a release brings the footprint back to about 2.95 GB. The heartbeat numbers had no script and are gone. Comments that ruff format had pushed onto closing brackets are back above their statements (syntax trees unchanged). No change to the worker's runtime code. --- benchmarks/laya_mps/README.md | 51 +++ benchmarks/laya_mps/bench_http.py | 40 +-- benchmarks/laya_mps/bench_inproc.py | 15 +- benchmarks/laya_mps/env.py | 32 +- benchmarks/laya_mps/fallback.py | 103 ++++++ benchmarks/laya_mps/late_load.py | 167 ++++++++++ benchmarks/laya_mps/lengths.py | 239 +++++++++++++ benchmarks/laya_mps/paired.py | 58 +++- benchmarks/laya_mps/profile_mps.py | 13 +- benchmarks/laya_mps/release.py | 165 +++++++++ benchmarks/laya_mps/report.py | 13 +- recipe/laya/apple-silicon.md | 28 +- src/frontend/laya_mps.py | 25 +- tests/laya/test_bench.py | 501 +++++++++++++++++++++++++--- tests/laya/test_worker.py | 44 ++- 15 files changed, 1326 insertions(+), 168 deletions(-) create mode 100644 benchmarks/laya_mps/fallback.py create mode 100644 benchmarks/laya_mps/late_load.py create mode 100644 benchmarks/laya_mps/lengths.py create mode 100644 benchmarks/laya_mps/release.py diff --git a/benchmarks/laya_mps/README.md b/benchmarks/laya_mps/README.md index c296329..5f4c864 100644 --- a/benchmarks/laya_mps/README.md +++ b/benchmarks/laya_mps/README.md @@ -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 | @@ -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`, and the same with `--gap 0.2`, `--gap 0.5`, `--gap 1` and `--gap 5` (each with its own `--run`) | +| 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/