1
Fork 0

feat: implemented checkpoint evaluation

This commit is contained in:
Jona Reynaert 2026-05-03 15:53:47 +02:00
parent 90c2c3d1f0
commit a3d6b32eaa

View file

@ -1,8 +1,10 @@
import datetime
import random
import time
import csv
from dataclasses import asdict, dataclass
from functools import partial
from pathlib import Path
from typing import Any
import jax
@ -26,6 +28,10 @@ from brittle_star_project.MLPs.mlps import (
)
from brittle_star_project.ppo import PPO
from brittle_star_project import Backend, BrittleStarEnv, BrittleStarEnvFactory
from brittle_star_project.evaluation.policy import PolicyAgent
from brittle_star_project.evaluation.rollout import rollout_headless
# TODO: clip scaled reward?
@ -538,6 +544,144 @@ class PPOTrainer:
params=self.agent_state.params, step=iteration, metadata=asdict(self.cfg)
)
def _maybe_sync_csv_to_wandb(self, csv_path: Path) -> None:
if not self.logging_cfg.track:
return
try:
import wandb
if wandb.run is None:
return
# "Simple sync" behavior: wandb will copy this file into the run.
wandb.save(str(csv_path), base_path=str(csv_path.parent))
except Exception as e:
self.logger.warning(f"[EVAL]: Failed to sync CSV to wandb: {e}")
def _append_checkpoint_eval_row(
self,
*,
iteration: int,
reached_target: bool,
steps_to_target: int,
max_steps: int,
seed: int,
) -> Path:
metrics_dir = Path(self.run_dir) / "metrics"
metrics_dir.mkdir(parents=True, exist_ok=True)
csv_path = metrics_dir / "checkpoint_evaluation.csv"
file_exists = csv_path.exists()
with open(csv_path, "a", newline="") as f:
writer = csv.DictWriter(
f,
fieldnames=[
"iteration",
"steps_to_target",
"reached_target",
"max_steps",
"seed",
],
)
if not file_exists:
writer.writeheader()
writer.writerow(
{
"iteration": int(iteration),
"steps_to_target": int(steps_to_target),
"reached_target": bool(reached_target),
"max_steps": int(max_steps),
"seed": int(seed),
}
)
return csv_path
def _evaluate_checkpoint(self, iteration: int) -> None:
if not self.logging_cfg.evaluate_checkpoints:
return
checkpoint_path = Path(self.run_dir) / "checkpoints" / f"checkpoint_step_{iteration}.flax"
if not checkpoint_path.exists():
self.logger.warning(
f"[EVAL]: Checkpoint not found at {checkpoint_path} (skipping evaluation)"
)
return
max_steps = int(self.logging_cfg.eval_max_steps)
seed = int(self.logging_cfg.eval_seed)
# Run evaluation best-effort; never fail training because evaluation failed.
env = None
try:
backend = Backend.MJC
factory = BrittleStarEnvFactory()
raw_env = factory.create_environment(
backend,
self.cfg.morphology,
self.cfg.arena,
self.cfg.environment,
)
env = BrittleStarEnv(
raw_env,
backend=backend,
config=self.cfg.environment,
morphology_config=self.cfg.morphology,
)
action_space = getattr(raw_env, "action_space", None)
action_dim = (
int(np.asarray(action_space.shape).reshape(-1)[0])
if action_space is not None and hasattr(action_space, "shape")
else sum(self.cfg.morphology.segments_per_arm) * 2
)
action_low = (
None
if action_space is None
else np.asarray(action_space.low, dtype=np.float32).ravel()
)
action_high = (
None
if action_space is None
else np.asarray(action_space.high, dtype=np.float32).ravel()
)
policy = PolicyAgent.from_checkpoint(
checkpoint_path,
action_dim=action_dim,
obs_processor=self.obs_processor,
)
result = rollout_headless(
env=env,
policy=policy,
seed=seed,
max_steps=max_steps,
action_low=action_low,
action_high=action_high,
action_mask=None,
)
steps_to_target = int(result.length) if result.reached_target else int(max_steps)
csv_path = self._append_checkpoint_eval_row(
iteration=iteration,
reached_target=bool(result.reached_target),
steps_to_target=steps_to_target,
max_steps=max_steps,
seed=seed,
)
self._maybe_sync_csv_to_wandb(csv_path)
except Exception as e:
self.logger.warning(f"[EVAL]: Checkpoint evaluation failed: {e}")
finally:
try:
if env is not None:
env.close()
except Exception:
pass
def train(self):
"""
Train the PPO agent for a specified number of iterations.
@ -591,6 +735,7 @@ class PPOTrainer:
if self.logging_cfg.save_checkpoints and self.logging_cfg.checkpoint_frequency > 0:
if iteration % self.logging_cfg.checkpoint_frequency == 0:
self._save_checkpoint(iteration)
self._evaluate_checkpoint(iteration)
if getattr(self.cfg.experiment, "debug_sanity", False):
self.logger.info("\n[SANITY CHECK] Successfully completed 1 epoch")