1
Fork 0

refactor(log): PPOTrainer

This commit is contained in:
Tibo De Peuter 2026-04-08 19:32:48 +02:00
parent 795ae48520
commit faf31567e8
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU

View file

@ -13,7 +13,8 @@ import numpy as np
import optax import optax
import tqdm import tqdm
from flax.training.train_state import TrainState from flax.training.train_state import TrainState
from torch.utils.tensorboard import SummaryWriter
from experiment_logger import get_logger
from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
@ -207,7 +208,7 @@ class PPOTrainer:
self.env = env self.env = env
self.run_dir = run_dir self.run_dir = run_dir
self.run_name = run_name self.run_name = run_name
self.writer = SummaryWriter(self.run_dir) self.logger = get_logger()
self.key = jax.random.PRNGKey(args.seed) self.key = jax.random.PRNGKey(args.seed)
@ -247,16 +248,14 @@ class PPOTrainer:
self._init_random() self._init_random()
def _init_random(self, log: bool = True): def _init_random(self):
if log: self.logger.info(f"[RANDOM]: Setting random seed to {self.args.seed}")
print(f"[RANDOM]: Setting random seed to {self.args.seed}")
random.seed(self.args.seed) random.seed(self.args.seed)
np.random.seed(self.args.seed) np.random.seed(self.args.seed)
def _init_agent(self, log: bool = True): def _init_agent(self):
if log: self.logger.info("[AGENT]: Initializing agent...")
print("[AGENT]: Initializing agent...")
sensor = GenericDenseLayersWithActivation() sensor = GenericDenseLayersWithActivation()
feature_extractor = GenericDenseLayersWithActivation() feature_extractor = GenericDenseLayersWithActivation()
@ -267,9 +266,8 @@ class PPOTrainer:
# messenger = OneDenseLayerMLP() # messenger = OneDenseLayerMLP()
return sensor, feature_extractor, actor, critic return sensor, feature_extractor, actor, critic
def _init_agent_state(self, log: bool = True) -> TrainState: def _init_agent_state(self) -> TrainState:
if log: self.logger.info("[AGENT STATE]: Initializing agent state...")
print("[AGENT STATE]: Initializing agent state...")
self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split( self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(
self.key, 5 self.key, 5
@ -313,9 +311,8 @@ class PPOTrainer:
), ),
) )
def _init_episode_stats(self, log: bool = True) -> EpisodeStatistics: def _init_episode_stats(self) -> EpisodeStatistics:
if log: self.logger.info("[EPISODE STATS]: Initializing episode stats...")
print("[EPISODE STATS]: Initializing episode stats...")
return EpisodeStatistics( return EpisodeStatistics(
episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32), episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32),
@ -350,39 +347,29 @@ class PPOTrainer:
iteration_time_start, iteration_time_start,
loss_info, loss_info,
): ):
metrics = {
"charts/avg_episodic_return": loss_info.avg_episodic_return,
"charts/avg_episodic_length": np.mean(
jax.device_get(episode_stats.returned_episode_lengths)
),
"charts/learning_rate": self.agent_state.opt_state[1]
.hyperparams["learning_rate"]
.item(),
"losses/value_loss": loss_info.v_loss[-1, -1].item(),
"losses/policy_loss": loss_info.pg_loss[-1, -1].item(),
"losses/entropy": loss_info.entropy_loss[-1, -1].item(),
"losses/approx_kl": loss_info.approx_kl[-1, -1].item(),
"losses/loss": loss_info.loss[-1, -1].item(),
"charts/SPS": int(global_step / (time.time() - start_time)),
"charts/SPS_update": int(
self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start)
),
}
self.logger.log(metrics, step=global_step)
self.writer.add_scalar( def _step(self, env_state, next_obs, next_done, is_tty: bool, iteration: int) -> tuple:
"charts/avg_episodic_return", loss_info.avg_episodic_return, global_step if not is_tty and iteration == 1:
) self.logger.info(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}")
self.writer.add_scalar(
"charts/avg_episodic_length",
np.mean(jax.device_get(episode_stats.returned_episode_lengths)),
global_step,
)
self.writer.add_scalar(
"charts/learning_rate",
self.agent_state.opt_state[1].hyperparams["learning_rate"].item(),
global_step,
)
self.writer.add_scalar("losses/value_loss", loss_info.v_loss[-1, -1].item(), global_step)
self.writer.add_scalar("losses/policy_loss", loss_info.pg_loss[-1, -1].item(), global_step)
self.writer.add_scalar("losses/entropy", loss_info.entropy_loss[-1, -1].item(), global_step)
self.writer.add_scalar("losses/approx_kl", loss_info.approx_kl[-1, -1].item(), global_step)
self.writer.add_scalar("losses/loss", loss_info.loss[-1, -1].item(), global_step)
self.writer.add_scalar(
"charts/SPS", int(global_step / (time.time() - start_time)), global_step
)
self.writer.add_scalar(
"charts/SPS_update",
int(self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start)),
global_step,
)
def _step(
self, env_state, next_obs, next_done, is_tty: bool, iteration: int, log: bool = True
) -> tuple:
if log and not is_tty and iteration == 1:
print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True)
( (
self.agent_state, self.agent_state,
@ -394,20 +381,20 @@ class PPOTrainer:
next_env_state, next_env_state,
) = self._rollout(env_state, next_obs, next_done) ) = self._rollout(env_state, next_obs, next_done)
if log and not is_tty and iteration == 1: if not is_tty and iteration == 1:
print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True) self.logger.info(f">>> [HPC] First rollout completed: {time.ctime()}")
storage = self._compute_gae(storage, next_obs, next_done) storage = self._compute_gae(storage, next_obs, next_done)
if log and not is_tty and iteration == 1: if not is_tty and iteration == 1:
print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True) self.logger.info(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}")
self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = ( self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = (
self._ppo.update_ppo(self.agent_state, storage, self.key) self._ppo.update_ppo(self.agent_state, storage, self.key)
) )
if log and not is_tty and iteration == 1: if not is_tty and iteration == 1:
print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True) self.logger.info(f">>> [HPC] First PPO update completed: {time.ctime()}")
avg_episodic_return = float( avg_episodic_return = float(
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item() jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item()
@ -429,77 +416,45 @@ class PPOTrainer:
def _close(self): def _close(self):
self.env.close() self.env.close()
self.writer.close()
def _save_model(self, model_path: str, log: bool = True): def _save_model(self, model_path: str):
if log: self.logger.info("[SAVE]: Saving the final model...")
print(f"[SAVE]: Saving the model to: {model_path}...")
with open(model_path, "wb") as f: params = [
f.write( vars(self.args),
flax.serialization.to_bytes( [
[ self.agent_state.params["sensor_params"],
vars(self.args), self.agent_state.params["actor_params"],
[ self.agent_state.params["critic_params"],
self.agent_state.params["sensor_params"], self.agent_state.params["feature_extractor_params"],
self.agent_state.params["actor_params"], ],
self.agent_state.params["critic_params"], ]
self.agent_state.params["feature_extractor_params"], self.logger.save_final_model(params=params)
],
]
)
)
def train(self, log: bool = True): def train(self):
""" """
Train the PPO agent for a specified number of iterations Train the PPO agent for a specified number of iterations
(passed through PPOArgs in constructor). (passed through PPOArgs in constructor).
Closes the environment at the end of training. Closes the environment at the end of training.
""" """
if log: self.logger.info(f"running name: {self.run_name}")
print(f"running name: {self.run_name}")
is_tty = sys.stdout.isatty() is_tty = sys.stdout.isatty()
if log: self.logger.info("[TRAIN]: Resetting environment...")
print("[TRAIN]: Resetting environment...")
if not is_tty: if not is_tty:
print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True) self.logger.info(f">>> [HPC] Initial reset started: {time.ctime()}")
env_state = self.env.reset(seed=self.args.seed) env_state = self.env.reset(seed=self.args.seed)
next_obs = _convert_obs_dict_to_array(env_state.observations) next_obs = _convert_obs_dict_to_array(env_state.observations)
next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_) next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_)
if log and not is_tty: if not is_tty:
print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True) self.logger.info(f">>> [HPC] Initial reset completed: {time.ctime()}")
global_step = 0 global_step = 0
start_time = time.time() start_time = time.time()
if self.args.track:
import wandb
if log:
print("[TRAIN]: Initializing Weights and Biases...")
wandb.init(
project=self.args.wandb_project_name,
entity=self.args.wandb_entity,
sync_tensorboard=True,
config=vars(self.args),
name=self.run_name,
save_code=True,
)
if log:
print("[TRAIN]: Adding hyperparameters to TensorBoard...")
self.writer.add_text(
"hyperparameters",
"|param|value|\n|---|---|\n"
+ "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()),
)
iter_bar = tqdm.tqdm( iter_bar = tqdm.tqdm(
range(1, self.args.num_iterations + 1), range(1, self.args.num_iterations + 1),
disable=not is_tty, disable=not is_tty,
@ -514,19 +469,18 @@ class PPOTrainer:
global_step += self.args.num_steps * self.args.num_envs global_step += self.args.num_steps * self.args.num_envs
self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info) self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info)
if log and not is_tty: if not is_tty:
sps = int(global_step / (time.time() - start_time)) sps = int(global_step / (time.time() - start_time))
remaining_steps = self.args.total_timesteps - global_step remaining_steps = self.args.total_timesteps - global_step
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0 eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
eta_str = str(datetime.timedelta(seconds=eta_seconds)) eta_str = str(datetime.timedelta(seconds=eta_seconds))
print( self.logger.info(
f"Iteration {iteration}/{self.args.num_iterations} | " f"Iteration {iteration}/{self.args.num_iterations} | "
f"Step {global_step}/{self.args.total_timesteps} | " f"Step {global_step}/{self.args.total_timesteps} | "
f"SPS {sps} | " f"SPS {sps} | "
f"Return {loss_info.avg_episodic_return:.4f} | " f"Return {loss_info.avg_episodic_return:.4f} | "
f"ETA {eta_str}", f"ETA {eta_str}"
flush=True,
) )
if self.args.save_model: if self.args.save_model: