Temporarily compare PyTorch and ONNX on the Space's hardware
Browse filesserver.py gains an "onnx" engine (checkpoint exported at startup by
onnx_backend.py). With LAYA_BENCH=1, set by app.py for now, both engines
answer the same questions in a background thread after startup; results
are served on GET /bench and printed to the logs. The demo keeps running
on PyTorch meanwhile.
Co-Authored-By: Claude Opus 5.5 <[email protected]>
- app.py +2 -0
- onnx_backend.py +102 -0
- requirements.txt +4 -0
- server.py +44 -4
app.py
CHANGED
|
@@ -28,6 +28,8 @@ os.environ.setdefault("LAYA_MODELS", "typed")
|
|
| 28 |
os.environ.setdefault("LAYA_DEVICE", "cpu")
|
| 29 |
# About 1 s per request on the Space's CPU: wait for a pause in typing before calling the model.
|
| 30 |
os.environ.setdefault("LAYA_DEBOUNCE", "700")
|
|
|
|
|
|
|
| 31 |
|
| 32 |
import gradio as gr # noqa: E402
|
| 33 |
import uvicorn # noqa: E402
|
|
|
|
| 28 |
os.environ.setdefault("LAYA_DEVICE", "cpu")
|
| 29 |
# About 1 s per request on the Space's CPU: wait for a pause in typing before calling the model.
|
| 30 |
os.environ.setdefault("LAYA_DEBOUNCE", "700")
|
| 31 |
+
# Temporary: compare PyTorch and ONNX on the Space's hardware (see GET /bench).
|
| 32 |
+
os.environ.setdefault("LAYA_BENCH", "1")
|
| 33 |
|
| 34 |
import gradio as gr # noqa: E402
|
| 35 |
import uvicorn # noqa: E402
|
onnx_backend.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""ONNX Runtime engine for the demo: the original laya checkpoint exported to ONNX at startup.
|
| 2 |
+
|
| 3 |
+
The export follows laya's scripts/export_onnx.py (Apache 2.0). The resulting file is not on the
|
| 4 |
+
Hub, so it is rebuilt on each start (about a minute on CPU) and kept in LAYA_ONNX_DIR.
|
| 5 |
+
"""
|
| 6 |
+
import os
|
| 7 |
+
import statistics
|
| 8 |
+
import time
|
| 9 |
+
|
| 10 |
+
ONNX_DIR = os.environ.get("LAYA_ONNX_DIR", os.path.join(os.path.expanduser("~"), ".cache", "laya-onnx"))
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def onnx_path_for(model_id):
|
| 14 |
+
return os.path.join(ONNX_DIR, model_id.replace("/", "--"), "model.onnx")
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def export(model_id, path):
|
| 18 |
+
"""Trace the PyTorch model with variable batch and sequence sizes and save it as ONNX."""
|
| 19 |
+
import torch
|
| 20 |
+
from laya.agent import Agent
|
| 21 |
+
|
| 22 |
+
agent = Agent(model_id, compile=False, device="cpu")
|
| 23 |
+
inputs = (
|
| 24 |
+
torch.randint(0, 100, (1, 16), dtype=torch.long), # input_ids
|
| 25 |
+
torch.ones((1, 16), dtype=torch.long), # attention_mask
|
| 26 |
+
torch.tensor([[1, 5]], dtype=torch.long), # marker_pos
|
| 27 |
+
torch.tensor([[True, True]], dtype=torch.bool), # marker_mask
|
| 28 |
+
torch.tensor([0], dtype=torch.long), # qtype
|
| 29 |
+
)
|
| 30 |
+
os.makedirs(os.path.dirname(path), exist_ok=True)
|
| 31 |
+
torch.onnx.export(
|
| 32 |
+
agent.model, inputs, path,
|
| 33 |
+
export_params=True, opset_version=18, do_constant_folding=True,
|
| 34 |
+
input_names=["input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype"],
|
| 35 |
+
output_names=["logits", "act_logits"],
|
| 36 |
+
dynamic_axes={
|
| 37 |
+
"input_ids": {0: "batch_size", 1: "seq_len"},
|
| 38 |
+
"attention_mask": {0: "batch_size", 1: "seq_len"},
|
| 39 |
+
"marker_pos": {0: "batch_size", 1: "num_markers"},
|
| 40 |
+
"marker_mask": {0: "batch_size", 1: "num_markers"},
|
| 41 |
+
"qtype": {0: "batch_size"},
|
| 42 |
+
"logits": {0: "batch_size", 1: "num_markers"},
|
| 43 |
+
"act_logits": {0: "batch_size"},
|
| 44 |
+
},
|
| 45 |
+
)
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def load(model_id):
|
| 49 |
+
from laya.onnx_agent import ONNXAgent
|
| 50 |
+
|
| 51 |
+
path = onnx_path_for(model_id)
|
| 52 |
+
if not os.path.exists(path):
|
| 53 |
+
start = time.perf_counter()
|
| 54 |
+
export(model_id, path)
|
| 55 |
+
print(f"ONNX export of {model_id} took {time.perf_counter() - start:.0f} s", flush=True)
|
| 56 |
+
return ONNXAgent(model_id, onnx_path=path)
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
# Texts of increasing length for the comparison; answers must not depend on the engine.
|
| 60 |
+
BENCH_TEXTS = [
|
| 61 |
+
"Félicitations ! Vous avez gagné un iPhone 17. Cliquez ici pour récupérer votre cadeau.",
|
| 62 |
+
"Bonjour, je suis CFO d'une scale-up de 120 personnes. Nous devons remplacer notre outil de "
|
| 63 |
+
"reporting avant la clôture de décembre et nous avons prévu une enveloppe de 40k€.",
|
| 64 |
+
"Bonjour, je suis directrice générale d'un groupe de distribution de 1 200 salariés. Notre "
|
| 65 |
+
"plateforme e-commerce plante à chaque pic de ventes et nous perdons du chiffre d'affaires. "
|
| 66 |
+
"Le conseil a approuvé un budget de 250k€ pour la refondre, je suis seule décisionnaire et la "
|
| 67 |
+
"mise en ligne doit se faire avant le Black Friday. Pouvez-vous démarrer dans deux semaines ?",
|
| 68 |
+
]
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def _probabilities(answers):
|
| 72 |
+
values = []
|
| 73 |
+
for key in sorted(answers):
|
| 74 |
+
answer = answers[key]
|
| 75 |
+
values += [answer["noul"]] if answer["type"] == "noul" else list(answer["probabilities"].values())
|
| 76 |
+
return values
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def compare(engines, cases, repeats=3):
|
| 80 |
+
"""Run every case on every text with each engine; report median times and the largest gap."""
|
| 81 |
+
times = {name: [] for name in engines}
|
| 82 |
+
worst_gap = 0.0
|
| 83 |
+
for case in cases.values():
|
| 84 |
+
questions = case["questions"]()
|
| 85 |
+
for agent in engines.values():
|
| 86 |
+
agent.predict("warmup", questions)
|
| 87 |
+
for text in BENCH_TEXTS:
|
| 88 |
+
results = {}
|
| 89 |
+
for name, agent in engines.items():
|
| 90 |
+
for _ in range(repeats):
|
| 91 |
+
start = time.perf_counter()
|
| 92 |
+
results[name] = _probabilities(agent.predict(text, questions)["answers"])
|
| 93 |
+
times[name].append(time.perf_counter() - start)
|
| 94 |
+
first, *others = results.values()
|
| 95 |
+
for other in others:
|
| 96 |
+
worst_gap = max(worst_gap, max(abs(a - b) for a, b in zip(first, other)))
|
| 97 |
+
return {
|
| 98 |
+
"median_ms": {name: round(statistics.median(t) * 1000) for name, t in times.items()},
|
| 99 |
+
"max_probability_gap": round(worst_gap, 4),
|
| 100 |
+
"cpu_count": os.cpu_count(),
|
| 101 |
+
"requests_per_engine": len(times[next(iter(times))]),
|
| 102 |
+
}
|
requirements.txt
CHANGED
|
@@ -5,3 +5,7 @@ laya==0.3.20
|
|
| 5 |
fastapi
|
| 6 |
uvicorn
|
| 7 |
spaces
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
fastapi
|
| 6 |
uvicorn
|
| 7 |
spaces
|
| 8 |
+
# Temporary ONNX comparison (onnx_backend.py).
|
| 9 |
+
onnx
|
| 10 |
+
onnxruntime
|
| 11 |
+
onnxscript
|
server.py
CHANGED
|
@@ -13,12 +13,15 @@ from questions import CASES
|
|
| 13 |
|
| 14 |
WEB = Path(__file__).parent / "web"
|
| 15 |
|
| 16 |
-
#
|
| 17 |
# - "mlx": laya-mlx, Apple Silicon only, fast local runs;
|
| 18 |
-
# - "torch": the original laya package (PyTorch), for Linux servers
|
|
|
|
|
|
|
| 19 |
CHECKPOINTS = {
|
| 20 |
"mlx": {"typed": "aac6fef/laya-typed-decisions-mlx", "english": "aac6fef/laya-mlx"},
|
| 21 |
-
"torch":
|
|
|
|
| 22 |
}
|
| 23 |
APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64"
|
| 24 |
BACKEND = os.environ.get("LAYA_BACKEND", "mlx" if APPLE_SILICON else "torch")
|
|
@@ -41,7 +44,16 @@ if origins:
|
|
| 41 |
|
| 42 |
# Optional device override ("cpu" to mimic a GPU-less server when testing locally).
|
| 43 |
device = {"device": os.environ["LAYA_DEVICE"]} if os.environ.get("LAYA_DEVICE") else {}
|
| 44 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
for agent in agents.values():
|
| 46 |
agent.predict("warmup", {"q": {"type": "noul", "instructions": "Is this a test?"}})
|
| 47 |
|
|
@@ -93,6 +105,34 @@ def explain(name: str):
|
|
| 93 |
"routing": {key: MODELS[model_for(case, key)] for key in case["questions"]()}}
|
| 94 |
|
| 95 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
@app.get("/health")
|
| 97 |
def health():
|
| 98 |
return {"backend": BACKEND, "models": list(MODELS.values())}
|
|
|
|
| 13 |
|
| 14 |
WEB = Path(__file__).parent / "web"
|
| 15 |
|
| 16 |
+
# Three engines run the same Laya weights and return the same answers:
|
| 17 |
# - "mlx": laya-mlx, Apple Silicon only, fast local runs;
|
| 18 |
+
# - "torch": the original laya package (PyTorch), for Linux servers;
|
| 19 |
+
# - "onnx": the same checkpoint exported to ONNX at startup (onnx_backend.py), CPU only.
|
| 20 |
+
UPSTREAM = {"typed": "convaiinnovations/laya-typed-decisions", "english": "convaiinnovations/laya"}
|
| 21 |
CHECKPOINTS = {
|
| 22 |
"mlx": {"typed": "aac6fef/laya-typed-decisions-mlx", "english": "aac6fef/laya-mlx"},
|
| 23 |
+
"torch": UPSTREAM,
|
| 24 |
+
"onnx": UPSTREAM,
|
| 25 |
}
|
| 26 |
APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64"
|
| 27 |
BACKEND = os.environ.get("LAYA_BACKEND", "mlx" if APPLE_SILICON else "torch")
|
|
|
|
| 44 |
|
| 45 |
# Optional device override ("cpu" to mimic a GPU-less server when testing locally).
|
| 46 |
device = {"device": os.environ["LAYA_DEVICE"]} if os.environ.get("LAYA_DEVICE") else {}
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def load_agent(backend, model_id):
|
| 50 |
+
if backend == "onnx":
|
| 51 |
+
import onnx_backend
|
| 52 |
+
return onnx_backend.load(model_id)
|
| 53 |
+
return laya.load(model_id, **device)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
agents = {alias: load_agent(BACKEND, model_id) for alias, model_id in MODELS.items()}
|
| 57 |
for agent in agents.values():
|
| 58 |
agent.predict("warmup", {"q": {"type": "noul", "instructions": "Is this a test?"}})
|
| 59 |
|
|
|
|
| 105 |
"routing": {key: MODELS[model_for(case, key)] for key in case["questions"]()}}
|
| 106 |
|
| 107 |
|
| 108 |
+
# Temporary benchmark (LAYA_BENCH=1): times PyTorch and ONNX on this machine, in the
|
| 109 |
+
# background after startup since it takes about a minute. Results: GET /bench and the logs.
|
| 110 |
+
bench_state = {"status": "disabled"}
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
def run_bench():
|
| 114 |
+
import onnx_backend
|
| 115 |
+
bench_state["status"] = "running"
|
| 116 |
+
try:
|
| 117 |
+
model_id = UPSTREAM["typed"]
|
| 118 |
+
torch_agent = agents["typed"] if BACKEND == "torch" else load_agent("torch", model_id)
|
| 119 |
+
engines = {"torch": torch_agent, "onnx": load_agent("onnx", model_id)}
|
| 120 |
+
bench_state.update(status="done", result=onnx_backend.compare(engines, CASES))
|
| 121 |
+
except Exception as error:
|
| 122 |
+
bench_state.update(status="failed", error=f"{type(error).__name__}: {error}")
|
| 123 |
+
print(f"bench: {bench_state}", flush=True)
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
if os.environ.get("LAYA_BENCH") == "1":
|
| 127 |
+
import threading
|
| 128 |
+
threading.Thread(target=run_bench, daemon=True).start()
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
@app.get("/bench")
|
| 132 |
+
def bench():
|
| 133 |
+
return bench_state
|
| 134 |
+
|
| 135 |
+
|
| 136 |
@app.get("/health")
|
| 137 |
def health():
|
| 138 |
return {"backend": BACKEND, "models": list(MODELS.values())}
|