Validate checkpoint credentials once; log throughput

#3
Files changed (1) hide show
  1. app.py +49 -7
app.py CHANGED
@@ -21,6 +21,7 @@ ENV_NAME = os.environ.get("ENV_NAME", "G1JoystickRoughTerrain")
21
  NUM_TIMESTEPS = int(os.environ.get("NUM_TIMESTEPS", 200_000_000))
22
  SEED = int(os.environ.get("SEED", 0))
23
  HF_REPO = os.environ.get("HF_REPO", "").strip()
 
24
  AUTO_START = os.environ.get("AUTO_START", "1") == "1"
25
 
26
  # /data exists only when persistent storage is attached; fall back to /tmp.
@@ -39,18 +40,49 @@ class Stopped(Exception):
39
  pass
40
 
41
 
42
- def _upload(path: Path) -> None:
 
 
 
 
43
  if not HF_REPO:
 
 
 
 
44
  return
45
  try:
46
  from huggingface_hub import HfApi
47
- api = HfApi()
 
48
  api.create_repo(HF_REPO, repo_type="model", exist_ok=True)
49
- api.upload_file(path_or_fileobj=str(path), path_in_repo=path.name,
50
- repo_id=HF_REPO, repo_type="model")
51
- log(f"uploaded {path.name} -> {HF_REPO}")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52
  except Exception as e: # never let upload kill training
53
- log(f"upload failed ({type(e).__name__}): {e}")
 
 
 
 
 
54
 
55
 
56
  def install_jax_pmap_shims() -> None:
@@ -127,14 +159,24 @@ def train_worker() -> None:
127
  f"timesteps={NUM_TIMESTEPS:,}")
128
 
129
  OUT.mkdir(parents=True, exist_ok=True)
 
 
130
 
131
  def progress(step, metrics):
132
  if STATE["stop"]:
133
  raise Stopped()
 
134
  STATE["step"] = int(step)
135
  rew = metrics.get("eval/episode_reward", float("nan"))
136
  ln = metrics.get("eval/avg_episode_length", float("nan"))
137
- log(f"step {int(step):>12,} reward {rew:8.2f} ep_len {ln:7.1f}")
 
 
 
 
 
 
 
138
 
139
  def policy_params_fn(step, make_policy, params):
140
  p = OUT / f"ckpt_{int(step)}.pkl"
 
21
  NUM_TIMESTEPS = int(os.environ.get("NUM_TIMESTEPS", 200_000_000))
22
  SEED = int(os.environ.get("SEED", 0))
23
  HF_REPO = os.environ.get("HF_REPO", "").strip()
24
+ HF_TOKEN = os.environ.get("HF_TOKEN", "").strip() or None
25
  AUTO_START = os.environ.get("AUTO_START", "1") == "1"
26
 
27
  # /data exists only when persistent storage is attached; fall back to /tmp.
 
40
  pass
41
 
42
 
43
+ UPLOAD = {"enabled": False, "fails": 0}
44
+
45
+
46
+ def init_upload() -> None:
47
+ """Validate credentials once, up front, instead of failing on every eval."""
48
  if not HF_REPO:
49
+ log(f"HF_REPO unset -- checkpoints stay in {OUT} and are lost on restart")
50
+ return
51
+ if not HF_TOKEN:
52
+ log("HF_TOKEN secret unset -- cannot push checkpoints")
53
  return
54
  try:
55
  from huggingface_hub import HfApi
56
+ api = HfApi(token=HF_TOKEN)
57
+ who = api.whoami().get("name")
58
  api.create_repo(HF_REPO, repo_type="model", exist_ok=True)
59
+ UPLOAD["enabled"] = True
60
+ log(f"checkpoints -> https://huggingface.co/{HF_REPO} (as {who})")
61
+ except Exception as e:
62
+ log(f"cannot write to HF_REPO={HF_REPO}: {type(e).__name__}: "
63
+ f"{str(e).splitlines()[0][:200]}")
64
+ log("fix: HF_TOKEN needs write scope on that namespace. A token scoped to "
65
+ "your own user cannot write to an org repo -- point HF_REPO at a "
66
+ "namespace the token owns, e.g. <your-user>/g1-rough-terrain.")
67
+
68
+
69
+ def _upload(path: Path) -> None:
70
+ if not UPLOAD["enabled"]:
71
+ return
72
+ try:
73
+ from huggingface_hub import HfApi
74
+ HfApi(token=HF_TOKEN).upload_file(
75
+ path_or_fileobj=str(path), path_in_repo=path.name,
76
+ repo_id=HF_REPO, repo_type="model")
77
+ log(f"uploaded {path.name}")
78
+ UPLOAD["fails"] = 0
79
  except Exception as e: # never let upload kill training
80
+ UPLOAD["fails"] += 1
81
+ log(f"upload failed ({type(e).__name__}): {str(e).splitlines()[0][:160]}")
82
+ if UPLOAD["fails"] >= 3:
83
+ UPLOAD["enabled"] = False
84
+ log("disabling uploads after 3 failures; training continues, "
85
+ f"checkpoints remain in {OUT}")
86
 
87
 
88
  def install_jax_pmap_shims() -> None:
 
159
  f"timesteps={NUM_TIMESTEPS:,}")
160
 
161
  OUT.mkdir(parents=True, exist_ok=True)
162
+ init_upload()
163
+ t_start = [None]
164
 
165
  def progress(step, metrics):
166
  if STATE["stop"]:
167
  raise Stopped()
168
+ import time
169
  STATE["step"] = int(step)
170
  rew = metrics.get("eval/episode_reward", float("nan"))
171
  ln = metrics.get("eval/avg_episode_length", float("nan"))
172
+ if t_start[0] is None:
173
+ t_start[0] = (time.time(), int(step))
174
+ sps = float("nan")
175
+ else:
176
+ t0, s0 = t_start[0]
177
+ sps = (int(step) - s0) / max(time.time() - t0, 1e-9)
178
+ log(f"step {int(step):>12,} reward {rew:8.2f} ep_len {ln:7.1f} "
179
+ f"{sps:,.0f} steps/s")
180
 
181
  def policy_params_fn(step, make_policy, params):
182
  p = OUT / f"ckpt_{int(step)}.pkl"