Spaces:
Sleeping
Sleeping
Make the timestep budget settable from the UI
#4
by arminfg - opened
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 =
|
| 157 |
log(f"num_envs={ppo_params.get('num_envs')} "
|
| 158 |
f"batch_size={ppo_params.get('batch_size')} "
|
| 159 |
-
f"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 |
-
|
|
|
|
| 240 |
return (f"**{ENV_NAME}** — status: `{STATE['status']}`"
|
| 241 |
f"{' (running)' if alive else ''} \n"
|
| 242 |
-
f"step {STATE['step']:,} / {
|
| 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.
|
|
|
|
|
|
|
| 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())
|