robynsd Claude Opus 5.5 commited on
Commit
d1ec01f
·
1 Parent(s): e4f8ca8

Remove the PyTorch vs ONNX comparison

Browse files

On the Space, ONNX Runtime took 5.6 s per request against 1.2 s for
PyTorch (identical answers), so the demo stays on PyTorch. The export
also added about a minute to each start.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

Files changed (3) hide show
  1. app.py +0 -2
  2. onnx_backend.py +0 -110
  3. requirements.txt +0 -4
app.py CHANGED
@@ -28,8 +28,6 @@ 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
- # 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
 
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
onnx_backend.py DELETED
@@ -1,110 +0,0 @@
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 subprocess
9
- import sys
10
- import time
11
-
12
- ONNX_DIR = os.environ.get("LAYA_ONNX_DIR", os.path.join(os.path.expanduser("~"), ".cache", "laya-onnx"))
13
-
14
-
15
- def onnx_path_for(model_id):
16
- return os.path.join(ONNX_DIR, model_id.replace("/", "--"), "model.onnx")
17
-
18
-
19
- def export(model_id, path):
20
- """Trace the PyTorch model with variable batch and sequence sizes and save it as ONNX."""
21
- import torch
22
- from laya.agent import Agent
23
-
24
- agent = Agent(model_id, compile=False, device="cpu")
25
- inputs = (
26
- torch.randint(0, 100, (1, 16), dtype=torch.long), # input_ids
27
- torch.ones((1, 16), dtype=torch.long), # attention_mask
28
- torch.tensor([[1, 5]], dtype=torch.long), # marker_pos
29
- torch.tensor([[True, True]], dtype=torch.bool), # marker_mask
30
- torch.tensor([0], dtype=torch.long), # qtype
31
- )
32
- os.makedirs(os.path.dirname(path), exist_ok=True)
33
- torch.onnx.export(
34
- agent.model, inputs, path,
35
- export_params=True, opset_version=18, do_constant_folding=True,
36
- input_names=["input_ids", "attention_mask", "marker_pos", "marker_mask", "qtype"],
37
- output_names=["logits", "act_logits"],
38
- dynamic_axes={
39
- "input_ids": {0: "batch_size", 1: "seq_len"},
40
- "attention_mask": {0: "batch_size", 1: "seq_len"},
41
- "marker_pos": {0: "batch_size", 1: "num_markers"},
42
- "marker_mask": {0: "batch_size", 1: "num_markers"},
43
- "qtype": {0: "batch_size"},
44
- "logits": {0: "batch_size", 1: "num_markers"},
45
- "act_logits": {0: "batch_size"},
46
- },
47
- )
48
-
49
-
50
- def load(model_id):
51
- from laya.onnx_agent import ONNXAgent
52
-
53
- path = onnx_path_for(model_id)
54
- if not os.path.exists(path):
55
- start = time.perf_counter()
56
- # A separate process: on ZeroGPU hardware the `spaces` package, already imported here,
57
- # replaces torch's CUDA with an emulation that makes the export fail.
58
- subprocess.run([sys.executable, os.path.abspath(__file__), model_id, path], check=True)
59
- print(f"ONNX export of {model_id} took {time.perf_counter() - start:.0f} s", flush=True)
60
- return ONNXAgent(model_id, onnx_path=path)
61
-
62
-
63
- # Texts of increasing length for the comparison; answers must not depend on the engine.
64
- BENCH_TEXTS = [
65
- "Félicitations ! Vous avez gagné un iPhone 17. Cliquez ici pour récupérer votre cadeau.",
66
- "Bonjour, je suis CFO d'une scale-up de 120 personnes. Nous devons remplacer notre outil de "
67
- "reporting avant la clôture de décembre et nous avons prévu une enveloppe de 40k€.",
68
- "Bonjour, je suis directrice générale d'un groupe de distribution de 1 200 salariés. Notre "
69
- "plateforme e-commerce plante à chaque pic de ventes et nous perdons du chiffre d'affaires. "
70
- "Le conseil a approuvé un budget de 250k€ pour la refondre, je suis seule décisionnaire et la "
71
- "mise en ligne doit se faire avant le Black Friday. Pouvez-vous démarrer dans deux semaines ?",
72
- ]
73
-
74
-
75
- def _probabilities(answers):
76
- values = []
77
- for key in sorted(answers):
78
- answer = answers[key]
79
- values += [answer["noul"]] if answer["type"] == "noul" else list(answer["probabilities"].values())
80
- return values
81
-
82
-
83
- def compare(engines, cases, repeats=3):
84
- """Run every case on every text with each engine; report median times and the largest gap."""
85
- times = {name: [] for name in engines}
86
- worst_gap = 0.0
87
- for case in cases.values():
88
- questions = case["questions"]()
89
- for agent in engines.values():
90
- agent.predict("warmup", questions)
91
- for text in BENCH_TEXTS:
92
- results = {}
93
- for name, agent in engines.items():
94
- for _ in range(repeats):
95
- start = time.perf_counter()
96
- results[name] = _probabilities(agent.predict(text, questions)["answers"])
97
- times[name].append(time.perf_counter() - start)
98
- first, *others = results.values()
99
- for other in others:
100
- worst_gap = max(worst_gap, max(abs(a - b) for a, b in zip(first, other)))
101
- return {
102
- "median_ms": {name: round(statistics.median(t) * 1000) for name, t in times.items()},
103
- "max_probability_gap": round(worst_gap, 4),
104
- "cpu_count": os.cpu_count(),
105
- "requests_per_engine": len(times[next(iter(times))]),
106
- }
107
-
108
-
109
- if __name__ == "__main__":
110
- export(sys.argv[1], sys.argv[2])
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
requirements.txt CHANGED
@@ -5,7 +5,3 @@ laya==0.3.20
5
  fastapi
6
  uvicorn
7
  spaces
8
- # Temporary ONNX comparison (onnx_backend.py).
9
- onnx
10
- onnxruntime
11
- onnxscript
 
5
  fastapi
6
  uvicorn
7
  spaces