Spaces:
Runtime error
Runtime error
Commit ·
a798c01
1
Parent(s): 812dbcc
fix: push reward curve + logs every 10k steps, not just at training end
Browse files
app.py
CHANGED
|
@@ -231,17 +231,58 @@ def _training_thread():
|
|
| 231 |
return True
|
| 232 |
|
| 233 |
class PeriodicHubPush(BaseCallback):
|
| 234 |
-
"""Pushes
|
| 235 |
-
Ensures no work is lost if the Space is interrupted."""
|
| 236 |
|
| 237 |
-
def __init__(self, api, hf_repo, hf_token, vec_env,
|
|
|
|
| 238 |
super().__init__()
|
| 239 |
-
self._api
|
| 240 |
-
self._repo
|
| 241 |
-
self._token
|
| 242 |
-
self._vec_env
|
| 243 |
-
self.
|
| 244 |
-
self.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 245 |
|
| 246 |
def _on_step(self):
|
| 247 |
if self.num_timesteps - self._last_push < self._push_every:
|
|
@@ -251,10 +292,13 @@ def _training_thread():
|
|
| 251 |
_log(f"Periodic save at step {self.num_timesteps:,} ...")
|
| 252 |
self.model.save("/home/user/app/spindleflow_model_latest")
|
| 253 |
self._vec_env.save("/home/user/app/vec_normalize_latest.pkl")
|
|
|
|
| 254 |
candidates = [
|
| 255 |
("/home/user/app/spindleflow_model_latest.zip", "spindleflow_model_latest.zip"),
|
| 256 |
("/home/user/app/vec_normalize_latest.pkl", "vec_normalize_latest.pkl"),
|
| 257 |
("/home/user/app/assets/training_log.txt", "training_log.txt"),
|
|
|
|
|
|
|
| 258 |
]
|
| 259 |
ops = [
|
| 260 |
CommitOperationAdd(path_in_repo=dst, path_or_fileobj=src)
|
|
@@ -328,7 +372,7 @@ def _training_thread():
|
|
| 328 |
)
|
| 329 |
periodic_push = PeriodicHubPush(
|
| 330 |
api=api, hf_repo=HF_REPO, hf_token=HF_TOKEN,
|
| 331 |
-
vec_env=vec_env, push_every=10_000,
|
| 332 |
)
|
| 333 |
|
| 334 |
model.learn(
|
|
|
|
| 231 |
return True
|
| 232 |
|
| 233 |
class PeriodicHubPush(BaseCallback):
|
| 234 |
+
"""Pushes checkpoint + log + reward curve to HF Hub every N steps."""
|
|
|
|
| 235 |
|
| 236 |
+
def __init__(self, api, hf_repo, hf_token, vec_env,
|
| 237 |
+
reward_logger_ref, push_every=10_000):
|
| 238 |
super().__init__()
|
| 239 |
+
self._api = api
|
| 240 |
+
self._repo = hf_repo
|
| 241 |
+
self._token = hf_token
|
| 242 |
+
self._vec_env = vec_env
|
| 243 |
+
self._rl_ref = reward_logger_ref
|
| 244 |
+
self._push_every = push_every
|
| 245 |
+
self._last_push = 0
|
| 246 |
+
|
| 247 |
+
def _save_curve(self):
|
| 248 |
+
ep = self._rl_ref.episode_rewards
|
| 249 |
+
if len(ep) < 2:
|
| 250 |
+
return
|
| 251 |
+
window = max(10, len(ep) // 20)
|
| 252 |
+
smoothed = [
|
| 253 |
+
float(np.mean(ep[max(0, i - window):i + 1]))
|
| 254 |
+
for i in range(len(ep))
|
| 255 |
+
]
|
| 256 |
+
step = max(1, len(ep) // 200)
|
| 257 |
+
with open("/home/user/app/assets/reward_curve.json", "w") as f:
|
| 258 |
+
json.dump({
|
| 259 |
+
"episodes": list(range(len(ep)))[::step],
|
| 260 |
+
"mean_rewards": smoothed[::step],
|
| 261 |
+
"raw_rewards": ep[::step],
|
| 262 |
+
"step": self.num_timesteps,
|
| 263 |
+
}, f)
|
| 264 |
+
import matplotlib, matplotlib.pyplot as plt
|
| 265 |
+
matplotlib.use("Agg")
|
| 266 |
+
plt.figure(figsize=(10, 4))
|
| 267 |
+
every = max(1, len(ep) // 500)
|
| 268 |
+
plt.plot(range(0, len(ep), every), ep[::every],
|
| 269 |
+
"o", markersize=2, alpha=0.2, color="#00d4ff",
|
| 270 |
+
label="Episode reward")
|
| 271 |
+
plt.plot(range(0, len(ep), every), smoothed[::every],
|
| 272 |
+
linewidth=2.5, color="#ff6b35",
|
| 273 |
+
label=f"Smoothed ({window}-ep mean)")
|
| 274 |
+
if len(ep) >= 5:
|
| 275 |
+
plt.axhline(float(np.mean(ep[:5])),
|
| 276 |
+
color="#94a3b8", linestyle="--", alpha=0.8,
|
| 277 |
+
label="Early baseline")
|
| 278 |
+
plt.axhline(float(np.mean(ep[-min(200, len(ep)):])),
|
| 279 |
+
color="#34d399", linestyle="--", alpha=0.8,
|
| 280 |
+
label="Current mean")
|
| 281 |
+
plt.xlabel("Episode"); plt.ylabel("Reward")
|
| 282 |
+
plt.title(f"SpindleFlow RL — Learning Curve (step {self.num_timesteps:,})")
|
| 283 |
+
plt.legend(); plt.grid(alpha=0.2); plt.tight_layout()
|
| 284 |
+
plt.savefig("/home/user/app/assets/reward_curve.png", dpi=150)
|
| 285 |
+
plt.close()
|
| 286 |
|
| 287 |
def _on_step(self):
|
| 288 |
if self.num_timesteps - self._last_push < self._push_every:
|
|
|
|
| 292 |
_log(f"Periodic save at step {self.num_timesteps:,} ...")
|
| 293 |
self.model.save("/home/user/app/spindleflow_model_latest")
|
| 294 |
self._vec_env.save("/home/user/app/vec_normalize_latest.pkl")
|
| 295 |
+
self._save_curve()
|
| 296 |
candidates = [
|
| 297 |
("/home/user/app/spindleflow_model_latest.zip", "spindleflow_model_latest.zip"),
|
| 298 |
("/home/user/app/vec_normalize_latest.pkl", "vec_normalize_latest.pkl"),
|
| 299 |
("/home/user/app/assets/training_log.txt", "training_log.txt"),
|
| 300 |
+
("/home/user/app/assets/reward_curve.json", "reward_curve.json"),
|
| 301 |
+
("/home/user/app/assets/reward_curve.png", "reward_curve.png"),
|
| 302 |
]
|
| 303 |
ops = [
|
| 304 |
CommitOperationAdd(path_in_repo=dst, path_or_fileobj=src)
|
|
|
|
| 372 |
)
|
| 373 |
periodic_push = PeriodicHubPush(
|
| 374 |
api=api, hf_repo=HF_REPO, hf_token=HF_TOKEN,
|
| 375 |
+
vec_env=vec_env, reward_logger_ref=reward_logger, push_every=10_000,
|
| 376 |
)
|
| 377 |
|
| 378 |
model.learn(
|