Make the timestep budget settable from the UI

#4
Files changed (1) hide show
  1. app.py +16 -7
app.py CHANGED
@@ -28,7 +28,8 @@ AUTO_START = os.environ.get("AUTO_START", "1") == "1"
28
  OUT = Path("/data/ckpt") if Path("/data").is_dir() else Path("/tmp/ckpt")
29
 
30
  LOG: deque[str] = deque(maxlen=800)
31
- STATE = {"thread": None, "status": "idle", "stop": False, "step": 0}
 
32
 
33
 
34
  def log(msg: str) -> None:
@@ -153,10 +154,10 @@ def train_worker() -> None:
153
  env = registry.load(ENV_NAME)
154
  env_cfg = registry.get_default_config(ENV_NAME)
155
  ppo_params = locomotion_params.brax_ppo_config(ENV_NAME)
156
- ppo_params.num_timesteps = NUM_TIMESTEPS
157
  log(f"num_envs={ppo_params.get('num_envs')} "
158
  f"batch_size={ppo_params.get('batch_size')} "
159
- f"timesteps={NUM_TIMESTEPS:,}")
160
 
161
  OUT.mkdir(parents=True, exist_ok=True)
162
  init_upload()
@@ -218,10 +219,15 @@ def train_worker() -> None:
218
  log(traceback.format_exc()[-2000:])
219
 
220
 
221
- def start():
 
 
222
  t = STATE["thread"]
223
  if t is not None and t.is_alive():
224
  return status_md()
 
 
 
225
  STATE["stop"] = False
226
  STATE["thread"] = threading.Thread(target=train_worker, daemon=True)
227
  STATE["thread"].start()
@@ -236,10 +242,11 @@ def stop():
236
 
237
  def status_md() -> str:
238
  alive = STATE["thread"] is not None and STATE["thread"].is_alive()
239
- pct = 100.0 * STATE["step"] / max(NUM_TIMESTEPS, 1)
 
240
  return (f"**{ENV_NAME}** — status: `{STATE['status']}`"
241
  f"{' (running)' if alive else ''} \n"
242
- f"step {STATE['step']:,} / {NUM_TIMESTEPS:,} ({pct:.1f}%) \n"
243
  f"checkpoints: `{OUT}`"
244
  + (f" → pushing to `{HF_REPO}`" if HF_REPO else
245
  " \n⚠️ `HF_REPO` unset — checkpoints are not pushed to the Hub"))
@@ -253,7 +260,9 @@ with gr.Blocks(title="G1 rough-terrain training") as demo:
253
  gr.Markdown("# Unitree G1 — rough-terrain locomotion training")
254
  st = gr.Markdown(status_md())
255
  with gr.Row():
256
- gr.Button("Start", variant="primary").click(start, outputs=st)
 
 
257
  gr.Button("Stop", variant="stop").click(stop, outputs=st)
258
  out = gr.Textbox(label="training log", lines=26, max_lines=26,
259
  autoscroll=True, value=logs())
 
28
  OUT = Path("/data/ckpt") if Path("/data").is_dir() else Path("/tmp/ckpt")
29
 
30
  LOG: deque[str] = deque(maxlen=800)
31
+ STATE = {"thread": None, "status": "idle", "stop": False, "step": 0,
32
+ "target": NUM_TIMESTEPS}
33
 
34
 
35
  def log(msg: str) -> None:
 
154
  env = registry.load(ENV_NAME)
155
  env_cfg = registry.get_default_config(ENV_NAME)
156
  ppo_params = locomotion_params.brax_ppo_config(ENV_NAME)
157
+ ppo_params.num_timesteps = int(STATE["target"])
158
  log(f"num_envs={ppo_params.get('num_envs')} "
159
  f"batch_size={ppo_params.get('batch_size')} "
160
+ f"timesteps={int(STATE['target']):,}")
161
 
162
  OUT.mkdir(parents=True, exist_ok=True)
163
  init_upload()
 
219
  log(traceback.format_exc()[-2000:])
220
 
221
 
222
+ def start(timesteps=None):
223
+ """Budget is settable from the UI so a rerun does not need a Space restart
224
+ (env vars are only injected at container start)."""
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)
233
  STATE["thread"].start()
 
242
 
243
  def status_md() -> str:
244
  alive = STATE["thread"] is not None and STATE["thread"].is_alive()
245
+ target = int(STATE["target"])
246
+ pct = 100.0 * STATE["step"] / max(target, 1)
247
  return (f"**{ENV_NAME}** — status: `{STATE['status']}`"
248
  f"{' (running)' if alive else ''} \n"
249
+ f"step {STATE['step']:,} / {target:,} ({pct:.1f}%) \n"
250
  f"checkpoints: `{OUT}`"
251
  + (f" → pushing to `{HF_REPO}`" if HF_REPO else
252
  " \n⚠️ `HF_REPO` unset — checkpoints are not pushed to the Hub"))
 
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())