1
Fork 0

feat(PPOTrainer.py): added logging messages

This commit is contained in:
Robin Meersman 2026-04-06 13:56:54 +02:00
parent 84f6fbd76e
commit b97fc30d5d
5 changed files with 135 additions and 27 deletions

View file

@ -1,5 +1,5 @@
# Minimal config to verify HPC setup is functional. # Minimal config to verify HPC setup is functional.
# Run with: python src/train.py --config-path configs/hpc/smoke_test.yaml # Run with: python experiments/train.py --config-path configs/hpc/smoke_test.yaml
exp_name: "hpc_smoke_test" exp_name: "hpc_smoke_test"
seed: 0 seed: 0
track: false # Test WandB integration track: false # Test WandB integration

View file

@ -1,4 +1,6 @@
import datetime
import random import random
import sys
import time import time
from dataclasses import asdict, dataclass from dataclasses import asdict, dataclass
from functools import partial from functools import partial
@ -200,11 +202,12 @@ class LossInfo:
class PPOTrainer: class PPOTrainer:
def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_name: str): def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_dir: str, run_name: str):
self.args = args self.args = args
self.env = env self.env = env
self.run_dir = run_dir
self.run_name = run_name self.run_name = run_name
self.writer = SummaryWriter(f"runs/{self.run_name}") self.writer = SummaryWriter(self.run_dir)
self.key = jax.random.PRNGKey(args.seed) self.key = jax.random.PRNGKey(args.seed)
@ -244,11 +247,17 @@ class PPOTrainer:
self._init_random() self._init_random()
def _init_random(self): def _init_random(self, log: bool = True):
if log:
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): def _init_agent(self, log: bool = True):
if log:
print("[AGENT]: Initializing agent...")
sensor = GenericDenseLayersWithActivation() sensor = GenericDenseLayersWithActivation()
feature_extractor = GenericDenseLayersWithActivation() feature_extractor = GenericDenseLayersWithActivation()
actor = Actor( actor = Actor(
@ -258,7 +267,10 @@ class PPOTrainer:
# messenger = OneDenseLayerMLP() # messenger = OneDenseLayerMLP()
return sensor, feature_extractor, actor, critic return sensor, feature_extractor, actor, critic
def _init_agent_state(self) -> TrainState: def _init_agent_state(self, log: bool = True) -> TrainState:
if log:
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
) )
@ -301,7 +313,10 @@ class PPOTrainer:
), ),
) )
def _init_episode_stats(self) -> EpisodeStatistics: def _init_episode_stats(self, log: bool = True) -> EpisodeStatistics:
if log:
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),
episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32), episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32),
@ -335,6 +350,7 @@ class PPOTrainer:
iteration_time_start, iteration_time_start,
loss_info, loss_info,
): ):
self.writer.add_scalar( self.writer.add_scalar(
"charts/avg_episodic_return", loss_info.avg_episodic_return, global_step "charts/avg_episodic_return", loss_info.avg_episodic_return, global_step
) )
@ -353,7 +369,6 @@ class PPOTrainer:
self.writer.add_scalar("losses/entropy", loss_info.entropy_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/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("losses/loss", loss_info.loss[-1, -1].item(), global_step)
self.writer.add_scalar( self.writer.add_scalar(
"charts/SPS", int(global_step / (time.time() - start_time)), global_step "charts/SPS", int(global_step / (time.time() - start_time)), global_step
) )
@ -363,7 +378,10 @@ class PPOTrainer:
global_step, global_step,
) )
def _step(self, env_state, next_obs, next_done) -> tuple: def _step(self, env_state, next_obs, next_done, is_tty: bool, iteration: int) -> tuple:
if not is_tty and iteration == 1:
print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True)
( (
self.agent_state, self.agent_state,
self.episode_stats, self.episode_stats,
@ -374,12 +392,21 @@ 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 not is_tty and iteration == 1:
print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True)
storage = self._compute_gae(storage, next_obs, next_done) storage = self._compute_gae(storage, next_obs, next_done)
if not is_tty and iteration == 1:
print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True)
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 not is_tty and iteration == 1:
print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True)
avg_episodic_return = float( avg_episodic_return = float(
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)) jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))
) )
@ -402,7 +429,10 @@ class PPOTrainer:
self.env.close() self.env.close()
self.writer.close() self.writer.close()
def _save_model(self, model_path: str): def _save_model(self, model_path: str, log: bool = True):
if log:
print(f"[SAVE]: Saving the model to: {model_path}...")
with open(model_path, "wb") as f: with open(model_path, "wb") as f:
f.write( f.write(
flax.serialization.to_bytes( flax.serialization.to_bytes(
@ -418,21 +448,38 @@ class PPOTrainer:
) )
) )
def train(self): def train(self, log: bool = True):
""" """
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:
print(f"running name: {self.run_name}")
is_tty = sys.stdout.isatty()
if log:
print("[TRAIN]: Resetting environment...")
if not is_tty:
print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True)
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:
print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True)
global_step = 0 global_step = 0
start_time = time.time() start_time = time.time()
if self.args.track: if self.args.track:
import wandb import wandb
if log:
print("[TRAIN]: Initializing Weights and Biases...")
wandb.init( wandb.init(
project=self.args.wandb_project_name, project=self.args.wandb_project_name,
entity=self.args.wandb_entity, entity=self.args.wandb_entity,
@ -442,21 +489,47 @@ class PPOTrainer:
save_code=True, save_code=True,
) )
if log:
print("[TRAIN]: Adding hyperparameters to TensorBoard...")
self.writer.add_text( self.writer.add_text(
"hyperparameters", "hyperparameters",
"|param|value|\n|---|---|\n" "|param|value|\n|---|---|\n"
+ "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()), + "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()),
) )
for _ in tqdm.tqdm(range(self.args.num_iterations)): iter_bar = tqdm.tqdm(
range(1, self.args.num_iterations + 1),
disable=not sys.stdout.isatty(),
)
for iteration in iter_bar:
iteration_time_start = time.time() iteration_time_start = time.time()
env_state, next_obs, next_done, loss_info = self._step(env_state, next_obs, next_done) env_state, next_obs, next_done, loss_info = self._step(env_state, next_obs, next_done)
if not is_tty and iteration == 1:
print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True)
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 not is_tty:
sps = int(global_step / (time.time() - start_time))
remaining_steps = self.args.total_timesteps - global_step
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
eta_str = str(datetime.timedelta(seconds=eta_seconds))
print(
f"Iteration {iteration}/{self.args.num_iterations} | "
f"Step {global_step}/{self.args.total_timesteps} | "
f"SPS {sps} | "
f"Return {loss_info.avg_episodic_return:.4f} | "
f"ETA {eta_str}",
flush=True,
)
if self.args.save_model: if self.args.save_model:
model_path = f"runs/{self.run_name}/{self.args.exp_name}.cleanrl_model" model_path = f"{self.run_dir}/{self.args.exp_name}.cleanrl_model"
self._save_model(model_path=model_path) self._save_model(model_path=model_path)
self._close() self._close()

View file

@ -81,7 +81,6 @@ def train(args: PPOArgs):
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}" run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
# args.num_iterations = args.total_timesteps // args.batch_size # args.num_iterations = args.total_timesteps // args.batch_size
args.num_iterations = 5 args.num_iterations = 5
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
print(f"running name: {run_name}") print(f"running name: {run_name}")
if args.run_dir is None: if args.run_dir is None:
@ -385,8 +384,6 @@ def train(args: PPOArgs):
) )
if args.save_model: if args.save_model:
model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model"
save_model(model_path, agent_state, args)
model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model" model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model"
with open(model_path, "wb") as f: with open(model_path, "wb") as f:
f.write( f.write(
@ -408,12 +405,6 @@ def train(args: PPOArgs):
writer.close() writer.close()
print("Saving loss plot...") print("Saving loss plot...")
simple_plot(
list(range(len(returns))),
returns,
show_window=True,
filename=f"runs/{run_name}/{args.exp_name}_losses.png",
)
plt.plot(losses) plt.plot(losses)
plt.title("PPO Loss, mean over minibatches") plt.title("PPO Loss, mean over minibatches")
plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png") plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png")

View file

@ -1,7 +1,10 @@
import subprocess
import time import time
import torch import torch
import tyro import tyro
import yaml
import os
from brittle_star_project.dataclasses import PPOArgs from brittle_star_project.dataclasses import PPOArgs
from PPOTrainer import PPOTrainer from PPOTrainer import PPOTrainer
@ -14,16 +17,53 @@ def make_env(config_path: str | None, num_envs: int) -> BrittleStarJaxEnvWrapper
return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs) return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs)
def parse_args() -> PPOArgs:
temp_args = tyro.cli(PPOArgs)
if temp_args.env_config_path is not None:
with open(temp_args.env_config_path, "r") as f:
config = yaml.safe_load(f)
if config:
# parse PPOArgs with defaults from yaml.
for key, value in config.items():
if hasattr(temp_args, key):
setattr(temp_args, key, value)
# Reparse CLI to ensure they OVERRIDE the yaml
args = tyro.cli(PPOArgs, default=temp_args)
else:
args = temp_args
return args
def get_git_hash() -> str:
try:
return (
subprocess.check_output(["git", "rev-parse", "--short", "HEAD"]).decode("ascii").strip()
)
except subprocess.CalledProcessError | UnicodeDecodeError:
return "none"
if __name__ == "__main__": if __name__ == "__main__":
args = tyro.cli(PPOArgs) args = parse_args()
args.batch_size = args.num_envs * args.num_steps args.batch_size = args.num_envs * args.num_steps
args.minibatch_size = args.batch_size // args.num_minibatches args.minibatch_size = args.batch_size // args.num_minibatches
args.num_iterations = args.total_timesteps // args.batch_size args.num_iterations = args.total_timesteps // args.batch_size
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
env = make_env(args.config_path, args.num_envs) git_hash = get_git_hash()
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
if args.run_dir is None:
run_dir = f"runs/{run_name}"
else:
run_dir = args.run_dir
os.makedirs(run_dir, exist_ok=True)
env = make_env(args.env_config_path, args.num_envs)
torch.backends.cudnn.deterministic = args.torch_deterministic torch.backends.cudnn.deterministic = args.torch_deterministic
ppo_trainer = PPOTrainer(args, env, run_name) ppo_trainer = PPOTrainer(args, env, run_dir, run_name)
ppo_trainer.train() ppo_trainer.train()

6
uv.lock generated
View file

@ -35,6 +35,9 @@ dependencies = [
] ]
[package.optional-dependencies] [package.optional-dependencies]
analysis = [
{ name = "tensorboard" },
]
cuda = [ cuda = [
{ name = "jax", extra = ["cuda13"] }, { name = "jax", extra = ["cuda13"] },
] ]
@ -64,12 +67,13 @@ requires-dist = [
{ name = "protobuf", specifier = ">=5.0.0" }, { name = "protobuf", specifier = ">=5.0.0" },
{ name = "pyopengl", specifier = ">=3.1.10" }, { name = "pyopengl", specifier = ">=3.1.10" },
{ name = "pyopengl-accelerate", specifier = ">=3.1.10" }, { name = "pyopengl-accelerate", specifier = ">=3.1.10" },
{ name = "tensorboard", marker = "extra == 'analysis'" },
{ name = "torch", specifier = ">=2.4.0" }, { name = "torch", specifier = ">=2.4.0" },
{ name = "tyro", specifier = ">=1.0.10" }, { name = "tyro", specifier = ">=1.0.10" },
{ name = "wandb", specifier = "==0.24.2" }, { name = "wandb", specifier = "==0.24.2" },
{ name = "warp-lang" }, { name = "warp-lang" },
] ]
provides-extras = ["cuda"] provides-extras = ["cuda", "analysis"]
[package.metadata.requires-dev] [package.metadata.requires-dev]
dev = [ dev = [