Stop external callers from killing runs via the public API

#5
Files changed (1) hide show
  1. app.py +15 -7
app.py CHANGED
@@ -1,7 +1,12 @@
1
  """Unitree G1 rough-terrain locomotion training, on the Space's GPU.
2
 
3
  Training runs in a background thread so the Gradio server stays responsive; the
4
- UI is a log tail plus start/stop. Brax PPO cannot be interrupted from outside,
 
 
 
 
 
5
  so the stop button sets a flag that the progress callback checks and raises on.
6
  """
7
 
@@ -225,8 +230,11 @@ def start(timesteps=None):
225
  t = STATE["thread"]
226
  if t is not None and t.is_alive():
227
  return status_md()
228
- if timesteps:
229
- STATE["target"] = int(timesteps)
 
 
 
230
  STATE["step"] = 0
231
  STATE["stop"] = False
232
  STATE["thread"] = threading.Thread(target=train_worker, daemon=True)
@@ -260,10 +268,10 @@ with gr.Blocks(title="G1 rough-terrain training") as demo:
260
  gr.Markdown("# Unitree G1 — rough-terrain locomotion training")
261
  st = gr.Markdown(status_md())
262
  with gr.Row():
263
- steps_in = gr.Number(value=NUM_TIMESTEPS, precision=0, label="timesteps",
264
- minimum=100_000, maximum=2_000_000_000)
265
- gr.Button("Start", variant="primary").click(start, inputs=steps_in, outputs=st)
266
- gr.Button("Stop", variant="stop").click(stop, outputs=st)
267
  out = gr.Textbox(label="training log", lines=26, max_lines=26,
268
  autoscroll=True, value=logs())
269
  timer = gr.Timer(3.0)
 
1
  """Unitree G1 rough-terrain locomotion training, on the Space's GPU.
2
 
3
  Training runs in a background thread so the Gradio server stays responsive; the
4
+ UI is a log tail plus start/stop.
5
+
6
+ Start/stop are deliberately NOT exposed as named API endpoints (api_name=False).
7
+ On a public Space, auto-named endpoints get probed -- a caller sent the string
8
+ "ping" to /start and then hit /stop, which killed a run mid-flight. The read-only
9
+ log and status endpoints stay exposed so the run can be monitored remotely. Brax PPO cannot be interrupted from outside,
10
  so the stop button sets a flag that the progress callback checks and raises on.
11
  """
12
 
 
230
  t = STATE["thread"]
231
  if t is not None and t.is_alive():
232
  return status_md()
233
+ try:
234
+ if timesteps is not None:
235
+ STATE["target"] = max(100_000, int(float(timesteps)))
236
+ except (TypeError, ValueError):
237
+ log(f"ignoring non-numeric timesteps input: {timesteps!r}")
238
  STATE["step"] = 0
239
  STATE["stop"] = False
240
  STATE["thread"] = threading.Thread(target=train_worker, daemon=True)
 
268
  gr.Markdown("# Unitree G1 — rough-terrain locomotion training")
269
  st = gr.Markdown(status_md())
270
  with gr.Row():
271
+ steps_in = gr.Number(value=NUM_TIMESTEPS, precision=0, label="timesteps")
272
+ gr.Button("Start", variant="primary").click(
273
+ start, inputs=steps_in, outputs=st, api_name=False)
274
+ gr.Button("Stop", variant="stop").click(stop, outputs=st, api_name=False)
275
  out = gr.Textbox(label="training log", lines=26, max_lines=26,
276
  autoscroll=True, value=logs())
277
  timer = gr.Timer(3.0)