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

Export ONNX in a separate process

Browse files

On 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>

Files changed (1) hide show
  1. 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
- 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
 
@@ -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])