robynsd Claude Opus 5.5 commited on
Commit
c966bc6
·
1 Parent(s): 26f8059

Temporarily compare PyTorch and ONNX on the Space's hardware

Browse files

server.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]>

Files changed (4) hide show
  1. app.py +2 -0
  2. onnx_backend.py +102 -0
  3. requirements.txt +4 -0
  4. 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
- # Two 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
  CHECKPOINTS = {
20
  "mlx": {"typed": "aac6fef/laya-typed-decisions-mlx", "english": "aac6fef/laya-mlx"},
21
- "torch": {"typed": "convaiinnovations/laya-typed-decisions", "english": "convaiinnovations/laya"},
 
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
- agents = {alias: laya.load(model_id, **device) for alias, model_id in MODELS.items()}
 
 
 
 
 
 
 
 
 
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())}