1
Fork 0

refactor: use mjx eval in PPOTrainer

This commit is contained in:
Tibo De Peuter 2026-05-07 23:01:56 +02:00
parent a7db4c77bb
commit 1f14ccfd4a
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU
2 changed files with 49 additions and 203 deletions

View file

@ -58,11 +58,6 @@ def build_eval_rollout_fn(
) -> Callable: ) -> Callable:
"""Build and JIT-compile a single-episode MJX evaluation rollout. """Build and JIT-compile a single-episode MJX evaluation rollout.
The returned function has the signature::
eval_fn(params: dict, seed: int, max_steps: int)
-> (steps, reached_target, eval_return, final_xy_dist, initial_xy_dist)
All outputs are JAX arrays. Convert to Python scalars before logging. All outputs are JAX arrays. Convert to Python scalars before logging.
Args: Args:
@ -141,13 +136,13 @@ def evaluate_checkpoint_mjx(
"""Run one deterministic evaluation episode and return typed metrics. """Run one deterministic evaluation episode and return typed metrics.
Args: Args:
eval_fn: A JIT-compiled function as returned by :func:`build_eval_rollout_fn`. eval_fn: A JIT-compiled function as returned by `build_eval_rollout_fn`.
params: Agent parameter dict (e.g. ``agent_state.params``). params: Agent parameter dict (e.g. ``agent_state.params``).
seed: Random seed for environment reset (controls target placement). seed: Random seed for environment reset (controls target placement).
max_steps: Maximum number of control steps before the episode is cut off. max_steps: Maximum number of control steps before the episode is cut off.
Returns: Returns:
A :class:`CheckpointEvalResult` with all JAX arrays converted to A `CheckpointEvalResult` with all JAX arrays converted to
plain Python scalars. plain Python scalars.
""" """
steps, reached, eval_return, final_xy_dist, initial_xy_dist = eval_fn(params, seed, max_steps) steps, reached, eval_return, final_xy_dist, initial_xy_dist = eval_fn(params, seed, max_steps)
@ -223,7 +218,7 @@ def append_checkpoint_eval_row(
run_dir: Root directory of the training run (Hydra's output dir). run_dir: Root directory of the training run (Hydra's output dir).
iteration: Training iteration number, used as the checkpoint identifier. iteration: Training iteration number, used as the checkpoint identifier.
trained_timesteps: Total environment steps taken at this checkpoint. trained_timesteps: Total environment steps taken at this checkpoint.
result: Evaluation result as returned by :func:`evaluate_checkpoint_mjx`. result: Evaluation result as returned by `evaluate_checkpoint_mjx`.
Returns: Returns:
Absolute path to the CSV file (useful for W&B sync). Absolute path to the CSV file (useful for W&B sync).

View file

@ -1,17 +1,15 @@
import datetime import datetime
import random import random
import time import time
import csv
from dataclasses import asdict, dataclass from dataclasses import asdict, dataclass
from functools import partial from functools import partial
from pathlib import Path
from typing import Any
import jax import jax
import jax.numpy as jnp import jax.numpy as jnp
import numpy as np import numpy as np
import optax import optax
from flax.training.train_state import TrainState from flax.training.train_state import TrainState
from typing import Any
from experiment_logger import get_logger from experiment_logger import get_logger
@ -19,6 +17,11 @@ from brittle_star_project.configs.main_config import BrittleStarConfig
from brittle_star_project.dataclasses import EpisodeStatistics from brittle_star_project.dataclasses import EpisodeStatistics
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from brittle_star_project.environment.obs_processing import create_obs_processor from brittle_star_project.environment.obs_processing import create_obs_processor
from brittle_star_project.evaluation.evaluate_mjx import (
append_checkpoint_eval_row,
build_eval_rollout_fn,
evaluate_checkpoint_mjx,
)
from brittle_star_project.MLPs.mlps import ( from brittle_star_project.MLPs.mlps import (
Actor, Actor,
AgentParams, AgentParams,
@ -122,8 +125,13 @@ def _step_once(
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
def _reward_fn(env_state, next_env_state): def reward_fn(env_state, next_env_state):
# if delta distance positive ==> brittle star walking away from target """Shaped reward used during training and checkpoint evaluation.
Public so that ``evaluation.evaluate_mjx`` can import it and produce
metrics that are directly comparable to training-time returns.
"""
# Positive delta_distance means the brittle star is moving *away* from target.
delta_distance = ( delta_distance = (
next_env_state.observations["xy_distance_to_target"] next_env_state.observations["xy_distance_to_target"]
- env_state.observations["xy_distance_to_target"] - env_state.observations["xy_distance_to_target"]
@ -142,7 +150,7 @@ def _reward_fn(env_state, next_env_state):
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_processor): def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_processor):
next_env_state = env_step_fn(env_state, action) next_env_state = env_step_fn(env_state, action)
reward = _reward_fn(env_state, next_env_state) reward = reward_fn(env_state, next_env_state)
terminated = next_env_state.terminated terminated = next_env_state.terminated
truncated = next_env_state.truncated truncated = next_env_state.truncated
done = terminated | truncated done = terminated | truncated
@ -327,74 +335,8 @@ class PPOTrainer:
self.episode_stats = self._init_episode_stats() self.episode_stats = self._init_episode_stats()
self._init_random() self._init_random()
# Lazily created MJX/JAX evaluation rollout (compiled on first use) # Lazily-built JIT-compiled MJX eval rollout, created on first evaluation.
self._eval_rollout_mjx_fn = None self._eval_fn = None
def _get_or_create_eval_rollout_mjx_fn(self):
if self._eval_rollout_mjx_fn is not None:
return self._eval_rollout_mjx_fn
# Use the same backend as training (typically MJX).
if getattr(self.env, "backend", None) != Backend.MJX:
self.logger.warning(
f"[EVAL]: Training env backend is {self.env.backend}; "
"MJX evaluation may be unavailable/slow."
)
# We vmap over a single environment (batch size 1) for simplicity.
reset_1 = jax.vmap(self.env.raw.reset)
step_1 = jax.vmap(self.env.raw.step)
action_low = self._action_low
action_high = self._action_high
obs_processor = self.obs_processor
sensor_apply = self.sensor.apply
actor_apply = self.actor.apply
def _eval_rollout(params, seed: int, max_steps: int):
rng = jax.random.PRNGKey(seed)
rngs = jnp.asarray(jax.random.split(rng, 1))
state = reset_1(rng=rngs)
initial_xy_dist = jnp.squeeze(state.observations["xy_distance_to_target"])
t0 = jnp.asarray(0, dtype=jnp.int32)
done0 = jnp.squeeze(state.terminated | state.truncated)
return0 = jnp.asarray(0.0, dtype=jnp.float32)
def cond(carry):
t, _state, done, _return_ = carry
return jnp.logical_and(t < max_steps, jnp.logical_not(done))
def body(carry):
t, state, _done, return_ = carry
obs = obs_processor(state.observations)
hidden = sensor_apply(params["sensor_params"], obs)
mean, _log_std = actor_apply(params["actor_params"], hidden)
action = jnp.clip(mean, action_low, action_high)
next_state = step_1(state=state, action=action)
# Match training's shaped reward as closely as possible.
shaped_reward = _reward_fn(state, next_state)
return_ = return_ + jnp.squeeze(shaped_reward)
done_next = jnp.squeeze(next_state.terminated | next_state.truncated)
return (t + 1, next_state, done_next, return_)
t, final_state, _done, return_ = jax.lax.while_loop(
cond, body, (t0, state, done0, return0)
)
reached_target = jnp.squeeze(final_state.terminated)
final_xy_dist_raw = jnp.squeeze(final_state.observations["xy_distance_to_target"])
final_xy_dist = jnp.where(reached_target, 0.0, final_xy_dist_raw)
return t, reached_target, return_, final_xy_dist, initial_xy_dist
self._eval_rollout_mjx_fn = jax.jit(_eval_rollout)
return self._eval_rollout_mjx_fn
def _init_random(self): def _init_random(self):
self.logger.info(f"[RANDOM]: Setting random seed to {self.experiment.seed}") self.logger.info(f"[RANDOM]: Setting random seed to {self.experiment.seed}")
@ -613,104 +555,12 @@ class PPOTrainer:
params=self.agent_state.params, step=iteration, metadata=asdict(self.cfg) 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,
trained_timesteps: int,
eval_steps: int,
eval_return: float,
final_xy_dist: float,
initial_xy_dist: float,
reached_target: bool,
) -> Path:
metrics_dir = Path(self.run_dir) / "metrics"
metrics_dir.mkdir(parents=True, exist_ok=True)
csv_path = metrics_dir / "checkpoint_evaluation.csv"
fieldnames = [
"checkpoint",
"trained_timesteps",
"eval_steps",
"eval_return",
"final_xy_dist",
"initial_xy_dist",
"reached_target",
]
# If a previous version created this CSV with a different header, migrate it.
if csv_path.exists():
try:
with open(csv_path, "r", newline="") as f:
reader = csv.reader(f)
header = next(reader, None)
if header is not None and list(header) != fieldnames:
migrated_rows: list[dict[str, Any]] = []
with open(csv_path, "r", newline="") as f:
dict_reader = csv.DictReader(f)
for row in dict_reader:
# Support older schemas best-effort.
checkpoint = row.get("checkpoint", row.get("iteration"))
steps = row.get("eval_steps", row.get("steps_to_target"))
migrated_rows.append(
{
"checkpoint": checkpoint,
"trained_timesteps": row.get("trained_timesteps"),
"eval_steps": steps,
"eval_return": row.get("eval_return"),
"final_xy_dist": row.get("final_xy_dist"),
"initial_xy_dist": row.get("initial_xy_dist"),
"reached_target": row.get("reached_target"),
}
)
with open(csv_path, "w", newline="") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
for row in migrated_rows:
writer.writerow(row)
except Exception:
# Best-effort only; do not fail training on migration issues.
pass
file_exists = csv_path.exists()
with open(csv_path, "a", newline="") as f:
writer = csv.DictWriter(
f,
fieldnames=fieldnames,
)
if not file_exists:
writer.writeheader()
writer.writerow(
{
"checkpoint": int(iteration),
"trained_timesteps": int(trained_timesteps),
"eval_steps": int(eval_steps),
"eval_return": float(eval_return),
"final_xy_dist": float(final_xy_dist),
"initial_xy_dist": float(initial_xy_dist),
"reached_target": bool(reached_target),
}
)
return csv_path
def _evaluate_checkpoint(self, iteration: int, *, trained_timesteps: int) -> None: def _evaluate_checkpoint(self, iteration: int, *, trained_timesteps: int) -> None:
"""Evaluate the current checkpoint and persist metrics to CSV.
Delegates all evaluation logic to `evaluation.evaluate_mjx`.
Best-effort: a failure here must never abort training.
"""
if not self.evaluation_cfg.evaluate_checkpoints: if not self.evaluation_cfg.evaluate_checkpoints:
return return
@ -728,35 +578,36 @@ class PPOTrainer:
) )
return return
# Run evaluation best-effort; never fail training because evaluation failed.
try: try:
eval_fn = self._get_or_create_eval_rollout_mjx_fn() if self._eval_fn is None:
( if getattr(self.env, "backend", None) != Backend.MJX:
steps, self.logger.warning(
reached, f"[EVAL]: Training env backend is {self.env.backend}; "
eval_return, "MJX evaluation may be unavailable/slow."
final_xy_dist, )
initial_xy_dist, self._eval_fn = build_eval_rollout_fn(
) = eval_fn(self.agent_state.params, seed, max_steps) env=self.env,
obs_processor=self.obs_processor,
sensor_apply=self.sensor.apply,
actor_apply=self.actor.apply,
action_low=self._action_low,
action_high=self._action_high,
reward_fn=reward_fn,
)
eval_steps = int(steps) result = evaluate_checkpoint_mjx(
reached_target = bool(reached) self._eval_fn,
self.agent_state.params,
# Keep numeric conversions explicit (JAX scalars -> Python scalars). seed=seed,
eval_return_f = float(eval_return) max_steps=max_steps,
final_xy_dist_f = float(final_xy_dist) )
initial_xy_dist_f = float(initial_xy_dist) csv_path = append_checkpoint_eval_row(
self.run_dir,
csv_path = self._append_checkpoint_eval_row(
iteration=iteration, iteration=iteration,
trained_timesteps=int(trained_timesteps), trained_timesteps=int(trained_timesteps),
eval_steps=eval_steps, result=result,
eval_return=eval_return_f,
final_xy_dist=final_xy_dist_f,
initial_xy_dist=initial_xy_dist_f,
reached_target=reached_target,
) )
self._maybe_sync_csv_to_wandb(csv_path) self.logger.sync_file(csv_path)
except Exception as e: except Exception as e:
self.logger.warning(f"[EVAL]: Checkpoint evaluation failed: {e}") self.logger.warning(f"[EVAL]: Checkpoint evaluation failed: {e}")