dexifried commited on
Commit
f57d914
·
1 Parent(s): 4553741

fix: skip ONNX export on HF, download checkpoint zip, export locally

Browse files
Files changed (1) hide show
  1. app.py +9 -29
app.py CHANGED
@@ -102,38 +102,18 @@ def train_and_export(encoder: str, epochs: int, batch_size: int, lr: float, max_
102
  eval_results = json.load(f)
103
  logs.append(f"\n📈 Results: {json.dumps(eval_results, indent=2)}")
104
 
105
- # ── Export ONNX ─────────────────────────────────────────────────────
106
- logs.append("\n📦 Exporting to ONNX...")
107
- onnx_dir = os.path.join(tmpdir, "onnx_export")
108
- export_cmd = _build_export_cmd(model_dir, onnx_dir)
109
- proc = subprocess.run(export_cmd, capture_output=True, text=True, cwd=str(Path(__file__).parent))
110
- logs.append(proc.stdout[-1000:] if proc.stdout else "")
111
- if proc.returncode != 0:
112
- logs.append(f"❌ Export failed:\n{proc.stderr[-1000:]}")
113
- return "\n".join(logs), None, None
114
-
115
- onnx_path = os.path.join(onnx_dir, "tiny_router.onnx")
116
- if not os.path.exists(onnx_path):
117
- logs.append("❌ ONNX file not found after export")
118
- return "\n".join(logs), None, None
119
-
120
- size_mb = os.path.getsize(onnx_path) / (1024 * 1024)
121
- logs.append(f"✅ ONNX exported! Size: {size_mb:.1f} MB")
122
-
123
- # Also grab temperature scaling config
124
- temp_path = os.path.join(model_dir, "temperature_scaling.json")
125
- temp_dest = os.path.join(tmpdir, "temperature_scaling.json")
126
- if os.path.exists(temp_path):
127
- shutil.copy2(temp_path, temp_dest)
128
 
129
- # Copy config
130
- config_path = os.path.join(model_dir, "config.json")
131
- config_dest = os.path.join(tmpdir, "config.json")
132
- if os.path.exists(config_path):
133
- shutil.copy2(config_path, config_dest)
 
134
 
135
  summary = json.dumps(eval_results, indent=2) if eval_results else "No eval results"
136
- return "\n".join(logs), onnx_path, summary
137
 
138
 
139
  def quick_predict(text: str):
 
102
  eval_results = json.load(f)
103
  logs.append(f"\n📈 Results: {json.dumps(eval_results, indent=2)}")
104
 
105
+ # ── Export ONNX (skip on HF — export locally after download) ─────
106
+ logs.append("\n📦 ONNX export skipped on HF (do locally after download).")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
107
 
108
+ # Zip the checkpoint for download
109
+ zip_path = os.path.join(tmpdir, "tiny-router-checkpoint")
110
+ shutil.make_archive(zip_path, "zip", model_dir)
111
+ zip_file = zip_path + ".zip"
112
+ size_mb = os.path.getsize(zip_file) / (1024 * 1024)
113
+ logs.append(f"✅ Checkpoint ready! ({size_mb:.1f} MB)")
114
 
115
  summary = json.dumps(eval_results, indent=2) if eval_results else "No eval results"
116
+ return "\n".join(logs), zip_file, summary
117
 
118
 
119
  def quick_predict(text: str):