Spaces:
Sleeping
Sleeping
Stop external callers from killing runs via the public API
#5
by arminfg - opened
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.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 229 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 265 |
-
|
| 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)
|