Export ONNX in a separate process
Browse filesOn ZeroGPU hardware the already imported spaces package replaces torch's
CUDA with an emulation, and torch.onnx.export failed on an operation it
did not intercept. A fresh process never imports spaces.
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
- onnx_backend.py +9 -1
onnx_backend.py
CHANGED
|
@@ -5,6 +5,8 @@ Hub, so it is rebuilt on each start (about a minute on CPU) and kept in LAYA_ONN
|
|
| 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"))
|
|
@@ -51,7 +53,9 @@ def load(model_id):
|
|
| 51 |
path = onnx_path_for(model_id)
|
| 52 |
if not os.path.exists(path):
|
| 53 |
start = time.perf_counter()
|
| 54 |
-
|
|
|
|
|
|
|
| 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 |
|
|
@@ -100,3 +104,7 @@ def compare(engines, cases, repeats=3):
|
|
| 100 |
"cpu_count": os.cpu_count(),
|
| 101 |
"requests_per_engine": len(times[next(iter(times))]),
|
| 102 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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"))
|
|
|
|
| 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 |
|
|
|
|
| 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])
|