refactor: use mjx eval in PPOTrainer
This commit is contained in:
parent
a7db4c77bb
commit
1f14ccfd4a
2 changed files with 49 additions and 203 deletions
|
|
@ -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).
|
||||||
|
|
|
||||||
|
|
@ -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}")
|
||||||
|
|
||||||
|
|
|
||||||
Reference in a new issue