From 45587c5620d87f5274ea770bf6f0ccb370305357 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Thu, 2 Apr 2026 15:11:39 +0200 Subject: [PATCH 01/19] fix: losses --> episodic returns --- src/train.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/src/train.py b/src/train.py index ac6849e..1ee00a0 100644 --- a/src/train.py +++ b/src/train.py @@ -245,7 +245,7 @@ def train(args: PPOArgs): print("Starting training...") iters_bar = tqdm.tqdm(range(1, args.num_iterations + 1)) - losses = [] + returns = [] for _ in iters_bar: iteration_time_start = time.time() @@ -259,13 +259,13 @@ def train(args: PPOArgs): agent_state, storage, key ) - losses.append(jnp.mean(loss)) - avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns)) iters_bar.set_postfix_str( f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" ) + returns.append(jnp.mean(avg_episodic_return)) + writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) writer.add_scalar( "charts/avg_episodic_length", @@ -313,8 +313,8 @@ def train(args: PPOArgs): writer.close() print("Saving loss plot...") - plt.plot(losses) - plt.title("PPO Loss, mean over minibatches") + plt.plot(returns) + plt.title("PPO Episodic Returns, mean over minibatches") plt.savefig(f"runs/{run_name}/{args.exp_name}_losses.png") plt.close() From 97eaa5c06b55e9290d47c04265300d05ea4fb74d Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Thu, 2 Apr 2026 15:13:27 +0200 Subject: [PATCH 02/19] chore(train.py): moved from package to experiments directory --- {src => experiments}/train.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename {src => experiments}/train.py (100%) diff --git a/src/train.py b/experiments/train.py similarity index 100% rename from src/train.py rename to experiments/train.py From ba59361f97b96fb4a83fc844dcc5cf89c0f876ca Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Thu, 2 Apr 2026 15:20:47 +0200 Subject: [PATCH 03/19] chore(main.py): deleted redundant main.py, feat(plot.py): added plot file to group all visualization related code for the experiments --- experiments/plots/plot.py | 11 +++++++++++ experiments/train.py | 10 ++++------ src/main.py | 4 ---- 3 files changed, 15 insertions(+), 10 deletions(-) create mode 100644 experiments/plots/plot.py delete mode 100644 src/main.py diff --git a/experiments/plots/plot.py b/experiments/plots/plot.py new file mode 100644 index 0000000..482f86e --- /dev/null +++ b/experiments/plots/plot.py @@ -0,0 +1,11 @@ +import matplotlib.pyplot as plt + + +def simple_plot(x: list, y: list, show_window: bool = False, filename: str = "plot.png") -> None: + plt.plot(x, y) + plt.savefig(filename) + + if show_window: + plt.show() + + plt.close() diff --git a/experiments/train.py b/experiments/train.py index 1ee00a0..bf59b3b 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -7,7 +7,6 @@ from typing import Callable import flax import jax import jax.numpy as jnp -import matplotlib.pyplot as plt import numpy as np import optax import torch @@ -20,6 +19,7 @@ from brittle_star_project.dataclasses import PPOArgs from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.rl import Actor, AgentParams, Critic, Network, Storage +from experiments.plots.plot import simple_plot from ppo import PPO @@ -41,7 +41,8 @@ def make_env(config_path: str | None, num_envs: int) -> Callable: def train(args: PPOArgs): args.batch_size = args.num_envs * args.num_steps 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 + args.num_iterations = 5 run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}" print(f"running name: {run_name}") @@ -313,10 +314,7 @@ def train(args: PPOArgs): writer.close() print("Saving loss plot...") - plt.plot(returns) - plt.title("PPO Episodic Returns, mean over minibatches") - plt.savefig(f"runs/{run_name}/{args.exp_name}_losses.png") - plt.close() + simple_plot(range(len(returns)), returns, show_window=True, filename=f"runs/{run_name}/{args.exp_name}_losses.png") def main() -> None: diff --git a/src/main.py b/src/main.py deleted file mode 100644 index ef7c36e..0000000 --- a/src/main.py +++ /dev/null @@ -1,4 +0,0 @@ -import jax - -if __name__ == "__main__": - print(jax.devices()) From 24da8948c50fe34cd25999bd573ad66b9403d2ea Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Fri, 3 Apr 2026 13:01:05 +0200 Subject: [PATCH 04/19] fix(merge): fixed merge conflict after popping local changes back in from stash --- experiments/plots/__init__.py | 3 +++ experiments/plots/plot.py | 1 + experiments/train.py | 43 ++++++++++++++++++++++------------- 3 files changed, 31 insertions(+), 16 deletions(-) create mode 100644 experiments/plots/__init__.py diff --git a/experiments/plots/__init__.py b/experiments/plots/__init__.py new file mode 100644 index 0000000..5abf51e --- /dev/null +++ b/experiments/plots/__init__.py @@ -0,0 +1,3 @@ +from .plot import simple_plot + +__all__ = ["simple_plot"] \ No newline at end of file diff --git a/experiments/plots/plot.py b/experiments/plots/plot.py index 482f86e..ab8f4ad 100644 --- a/experiments/plots/plot.py +++ b/experiments/plots/plot.py @@ -6,6 +6,7 @@ def simple_plot(x: list, y: list, show_window: bool = False, filename: str = "pl plt.savefig(filename) if show_window: + # blocks until window is closed plt.show() plt.close() diff --git a/experiments/train.py b/experiments/train.py index 877a882..04a3df4 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -25,6 +25,7 @@ from MLPs.mlps import ( AgentParams, Storage, ) +from plots import simple_plot from ppo import PPO @@ -42,6 +43,21 @@ def make_env(config_path: str | None, num_envs: int) -> Callable: return thunk +def save_model(model_path: str, agent_state: TrainState, args: PPOArgs): + with open(model_path, "wb") as f: + f.write( + flax.serialization.to_bytes( + [ + vars(args), + [ + agent_state.params["network_params"], + agent_state.params["actor_params"], + agent_state.params["critic_params"], + ], + ] + ) + ) + def train(args: PPOArgs): args.batch_size = args.num_envs * args.num_steps @@ -182,6 +198,8 @@ def train(args: PPOArgs): value = critic.apply(agent_state.params["critic_params"], hidden_critic) return action, logprob, value.squeeze(-1), key + + # GAE @jax.jit def compute_gae_once(carry, inp, gamma, gae_lambda): advantages = carry @@ -208,6 +226,7 @@ def train(args: PPOArgs): reverse=True, ) return storage.replace(advantages=advantages, returns=advantages + storage.values) + # END GAE # --- Main training loop --- global_step = 0 @@ -277,7 +296,7 @@ def train(args: PPOArgs): f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" ) - returns.append(jnp.mean(avg_episodic_return)) + returns.append(avg_episodic_return) writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) writer.add_scalar( @@ -307,27 +326,19 @@ def train(args: PPOArgs): if args.save_model: model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model" - with open(model_path, "wb") as f: - f.write( - flax.serialization.to_bytes( - [ - vars(args), - [ - agent_state.params["sensor_params"], - agent_state.params["actor_params"], - agent_state.params["critic_params"], - agent_state.params["feature_extractor_params"], - ], - ] - ) - ) + save_model(model_path, agent_state, args) print(f"model saved to {model_path}") env.close() writer.close() print("Saving loss plot...") - simple_plot(range(len(returns)), returns, show_window=True, filename=f"runs/{run_name}/{args.exp_name}_losses.png") + simple_plot( + list(range(len(returns))), + returns, + show_window=True, + filename=f"runs/{run_name}/{args.exp_name}_losses.png", + ) def main() -> None: From 4396c9ac4b3bed08519c536f1b79fe25b82e87f5 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Fri, 3 Apr 2026 14:23:09 +0200 Subject: [PATCH 05/19] feat(train.py, PPOTrainer.py): cleaned up training loop to specialized class --- experiments/PPOTrainer.py | 383 ++++++++++++++++++++++++++++++++++++++ experiments/train.py | 343 +--------------------------------- experiments/train.py.back | 350 ++++++++++++++++++++++++++++++++++ 3 files changed, 743 insertions(+), 333 deletions(-) create mode 100644 experiments/PPOTrainer.py create mode 100644 experiments/train.py.back diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py new file mode 100644 index 0000000..347b724 --- /dev/null +++ b/experiments/PPOTrainer.py @@ -0,0 +1,383 @@ +import time +from dataclasses import asdict, dataclass +from functools import partial +from typing import Any + +import optax +import tqdm +from flax.metrics.tensorboard import SummaryWriter +from flax.training.train_state import TrainState + +from MLPs.mlps import ( + GenericDenseLayersWithActivation, + Actor, + OneDenseLayerMLP, + AgentParams, + Storage, +) +from brittle_star_project.dataclasses import PPOArgs, EpisodeStatistics +from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +import jax +import jax.numpy as jnp +import numpy as np + +from ppo import PPO + + +@jax.jit +def linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): + frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations + return learning_rate * frac + + +@jax.jit +def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: + return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( + obs_dict + ) + + +@jax.jit +def get_action_and_value_noise( + sensor: GenericDenseLayersWithActivation, + feature_extractor: GenericDenseLayersWithActivation, + actor: Actor, + critic: OneDenseLayerMLP, + agent_state: TrainState, + next_obs: jnp.ndarray, + key: jax.random.PRNGKey, +): + hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) + hidden_critic = feature_extractor.apply( + agent_state.params["feature_extractor_params"], next_obs + ) + + # Continuous actions: sample from a Gaussian parameterized by the actor + mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) + key, subkey = jax.random.split(key) + noise = jax.random.normal(subkey, shape=mean.shape) + std = jnp.exp(log_std) + action = mean + noise * std + logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) + value = critic.apply(agent_state.params["critic_params"], hidden_critic) + return action, logprob, value.squeeze(-1), key + + +@jax.jit +def _step_once( + carry, + _, + env_step_fn, + sensor: GenericDenseLayersWithActivation, + feature_extractor: GenericDenseLayersWithActivation, + actor: Actor, + critic: OneDenseLayerMLP, +): + agent_state, episode_stats, obs, done, key, env_state = carry + action, logprob, value, key = get_action_and_value_noise( + sensor, feature_extractor, actor, critic, agent_state, obs, key + ) + + episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( + episode_stats, env_state, action + ) + + storage = Storage( + obs=obs, + actions=action, + logprobs=logprob, + dones=done, + values=value, + rewards=reward, + returns=jnp.zeros_like(reward), + advantages=jnp.zeros_like(reward), + ) + return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage + + +@jax.jit +def _rollout_jit( + agent_state, + episode_stats, + env_state, + next_obs, + next_done, + key, + args, + env_step_fn, + sensor: GenericDenseLayersWithActivation, + feature_extractor: GenericDenseLayersWithActivation, + actor: Actor, + critic: OneDenseLayerMLP, +): + (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( + partial( + _step_once, + sensor=sensor, + feature_extractor=feature_extractor, + actor=actor, + critic=critic, + env_step_fn=partial(_step_env_wrapped, env_step_fn=env_step_fn), + ), + (agent_state, episode_stats, next_obs, next_done, key, env_state), + (), + args.max_steps, + ) + return agent_state, episode_stats, next_obs, next_done, storage, key, env_state + + +@jax.jit +def _step_env_wrapped(env_step_fn, env_state, action, episode_stats): + next_env_state = env_step_fn(env_state, action) + + # Extract per-environment signals from the state object + reward = next_env_state.reward # (num_envs,) + terminated = next_env_state.terminated # (num_envs,) + truncated = next_env_state.truncated # (num_envs,) + done = terminated | truncated # (num_envs,) + + new_episode_return = episode_stats.episode_returns + reward + new_episode_length = episode_stats.episode_lengths + 1 + + episode_stats = episode_stats.replace( + episode_returns=new_episode_return * (1 - done), + episode_lengths=new_episode_length * (1 - done), + returned_episode_returns=jnp.where( + done, new_episode_return, episode_stats.returned_episode_returns + ), + returned_episode_lengths=jnp.where( + done, new_episode_length, episode_stats.returned_episode_lengths + ), + ) + return ( + episode_stats, + next_env_state, + (convert_obs_dict_to_array(next_env_state.observations), reward, done), + ) + + +@jax.jit +def compute_gae_once(carry, inp, gamma, gae_lambda): + advantages = carry + nextdone, nextvalues, curvalues, reward = inp + nextnonterminal = 1.0 - nextdone + delta = reward + gamma * nextvalues * nextnonterminal - curvalues + advantages = delta + gamma * gae_lambda * nextnonterminal * advantages + return advantages, advantages + + +@jax.jit +def compute_gae_jit(agent_state, storage, next_obs, next_done, sensor, critic, args): + next_value = critic.apply( + agent_state.params["critic_params"], + sensor.apply(agent_state.params["sensor_params"], next_obs), + ).squeeze(-1) + + advantages = jnp.zeros((args.num_envs,)) + dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0) + values = jnp.concatenate([storage.values, next_value[None, :]], axis=0) + _, advantages = jax.lax.scan( + partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda), + advantages, + (dones[1:], values[1:], values[:-1], storage.rewards), + reverse=True, + ) + return storage.replace(advantages=advantages, returns=advantages + storage.values) + + +@dataclass +class LossInfo: + # todo: better typing + loss: Any + pg_loss: Any + v_loss: Any + entropy_loss: Any + approx_kl: Any + avg_episodic_return: Any + + +class PPOTrainer: + def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_name: str): + self.args = args + self.env = env + self.writer = SummaryWriter(f"runs/{run_name}") + + self.key = jax.random.PRNGKey(args.seed) + + self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent() + self.sensor.apply = jax.jit(self.sensor.apply) + self.feature_extractor.apply = jax.jit(self.feature_extractor.apply) + self.actor.apply = jax.jit(self.actor.apply) + self.critic.apply = jax.jit(self.critic.apply) + + self._ppo = PPO(self.args, self.sensor, self.actor, self.critic, self.feature_extractor) + + self.agent_state = self._init_agent_state() + + self.episode_stats = self._init_episode_stats() + + def _init_agent(self): + sensor = GenericDenseLayersWithActivation() + feature_extractor = GenericDenseLayersWithActivation() + actor = Actor( + action_dim=self.env.single_action_space.shape[0] + ) # continuous actions for MJX + critic = OneDenseLayerMLP() + # messenger = OneDenseLayerMLP() + return sensor, feature_extractor, actor, critic + + def _init_agent_state(self) -> TrainState: + self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split( + self.key, 5 + ) + + sample_obs = jnp.concatenate( + [ + v.flatten() + for v in self.env.single_observation_space.sample( + rng=jax.random.PRNGKey(0) + ).values() + if v.size > 0 + ] + ) + sensor_params = self.sensor.init(sensor_key, sample_obs) + feature_extractor_params = self.feature_extractor.init(feature_extractor_key, sample_obs) + actor_params = self.actor.init(actor_key, self.sensor.apply(sensor_params, sample_obs)) + critic_params = self.critic.init( + critic_key, self.feature_extractor.apply(feature_extractor_params, sample_obs) + ) + + return TrainState.create( + apply_fn=None, + params=asdict( + AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) + ), + tx=optax.chain( + optax.clip_by_global_norm(self.args.max_grad_norm), + optax.inject_hyperparams(optax.adam)( + learning_rate=linear_schedule + if self.args.anneal_lr + else self.args.learning_rate, + eps=1e-5, + ), + ), + ) + + def _init_episode_stats(self) -> EpisodeStatistics: + return EpisodeStatistics( + episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32), + episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32), + returned_episode_returns=jnp.zeros(self.args.num_envs, jnp.float32), + returned_episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32), + ) + + def _rollout(self, env_state, next_obs, next_done) -> tuple[Storage, ...]: + return _rollout_jit( + self.agent_state, + self.episode_stats, + env_state, + next_obs, + next_done, + self.key, + self.args, + self.env.step, + self.sensor, + self.feature_extractor, + self.actor, + self.critic, + ) + + def _compute_gae(self, storage, next_obs, next_done) -> Storage: + return compute_gae_jit( + self.agent_state, storage, next_obs, next_done, self.sensor, self.critic, self.args + ) + + def _log( + self, + global_step, + episode_stats, + avg_episodic_return, + start_time, + iteration_time_start, + loss_info, + ): + self.writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) + 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) + + # iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}") + + 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) -> tuple: + storage, next_obs, next_done, env_state = self._rollout(env_state, next_obs, next_done) + storage = self._compute_gae(storage, next_obs, next_done) + self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = ( + self._ppo.update_ppo(self.agent_state, storage, self.key) + ) + + avg_episodic_return = jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)) + + return ( + env_state, + next_obs, + next_done, + LossInfo( + loss=loss, + pg_loss=pg_loss, + v_loss=v_loss, + entropy_loss=entropy_loss, + approx_kl=approx_kl, + avg_episodic_return=avg_episodic_return, + ), + ) + + def close(self): + self.env.close() + self.writer.close() + + def train(self): + """ + Train the PPO agent for a specified number of iterations + (passed through PPOArgs in constructor). + Closes the environment at the end of training. + """ + env_state = self.env.reset(seed=self.args.seed) + next_obs = convert_obs_dict_to_array(env_state.observations) + next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_) + global_step = 0 + start_time = time.time() + + for _ in tqdm.tqdm(range(self.args.num_iterations)): + iteration_time_start = time.time() + + env_state, next_obs, next_done, loss_info = self._step(env_state, next_obs, next_done) + global_step += self.args.num_steps * self.args.num_envs + self._log( + global_step, self.episode_stats, 0, start_time, iteration_time_start, loss_info + ) + + if self.args.save_model: + self._save_model(...) + + self.close() diff --git a/experiments/train.py b/experiments/train.py index 04a3df4..ba9dfcb 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -1,350 +1,27 @@ -import random import time -from dataclasses import asdict -from functools import partial -from typing import Callable -import flax -import jax -import jax.numpy as jnp -import numpy as np -import optax -import torch -import tqdm import tyro -from flax.training.train_state import TrainState -from torch.utils.tensorboard import SummaryWriter from brittle_star_project.dataclasses import PPOArgs -from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics +from PPOTrainer import PPOTrainer from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper -from MLPs.mlps import ( - GenericDenseLayersWithActivation, - Actor, - OneDenseLayerMLP, - AgentParams, - Storage, -) -from plots import simple_plot -from ppo import PPO -def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: - return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( - obs_dict - ) +def make_env(config_path: str | None, num_envs: int) -> BrittleStarJaxEnvWrapper: + if config_path is None: + return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) + return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs) -def make_env(config_path: str | None, num_envs: int) -> Callable: - def thunk(): - if config_path is None: - return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) - return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs) +if __name__ == "__main__": + args = tyro.cli(PPOArgs) - return thunk - -def save_model(model_path: str, agent_state: TrainState, args: PPOArgs): - with open(model_path, "wb") as f: - f.write( - flax.serialization.to_bytes( - [ - vars(args), - [ - agent_state.params["network_params"], - agent_state.params["actor_params"], - agent_state.params["critic_params"], - ], - ] - ) - ) - - -def train(args: PPOArgs): args.batch_size = args.num_envs * args.num_steps args.minibatch_size = args.batch_size // args.num_minibatches # args.num_iterations = args.total_timesteps // args.batch_size args.num_iterations = 5 run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}" - print(f"running name: {run_name}") + env = make_env(args.config_path, args.num_envs) - if args.track: - import wandb - - wandb.init( - project=args.wandb_project_name, - entity=args.wandb_entity, - sync_tensorboard=True, - config=vars(args), - name=run_name, - save_code=True, - ) - - writer = SummaryWriter(f"runs/{run_name}") - writer.add_text( - "hyperparameters", - "|param|value|\n|---|---|\n" + "\n".join(f"|{k}|{v}|" for k, v in vars(args).items()), - ) - - random.seed(args.seed) - np.random.seed(args.seed) - key = jax.random.PRNGKey(args.seed) - key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5) - - torch.backends.cudnn.deterministic = args.torch_deterministic - device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") - print(f"Running on device: {device}") - - print("Creating the environment...") - env = make_env(config_path=args.config_path, num_envs=args.num_envs)() - print(f"Environment: {env}") - - episode_stats = EpisodeStatistics( - episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32), - episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), - returned_episode_returns=jnp.zeros(args.num_envs, jnp.float32), - returned_episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), - ) - - def step_env_wrapped(episode_stats: EpisodeStatistics, env_state, action): - next_env_state = env.step(env_state, action) - - # Extract per-environment signals from the state object - reward = next_env_state.reward # (num_envs,) - terminated = next_env_state.terminated # (num_envs,) - truncated = next_env_state.truncated # (num_envs,) - done = terminated | truncated # (num_envs,) - - new_episode_return = episode_stats.episode_returns + reward - new_episode_length = episode_stats.episode_lengths + 1 - - episode_stats = episode_stats.replace( - episode_returns=new_episode_return * (1 - done), - episode_lengths=new_episode_length * (1 - done), - returned_episode_returns=jnp.where( - done, new_episode_return, episode_stats.returned_episode_returns - ), - returned_episode_lengths=jnp.where( - done, new_episode_length, episode_stats.returned_episode_lengths - ), - ) - return ( - episode_stats, - next_env_state, - (convert_obs_dict_to_array(next_env_state.observations), reward, done), - ) - - def linear_schedule(count): - frac = 1.0 - (count // (args.num_minibatches * args.update_epochs)) / args.num_iterations - return args.learning_rate * frac - - print("Initializing the models...") - sensor = GenericDenseLayersWithActivation() - feature_extractor = GenericDenseLayersWithActivation() - actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX - critic = OneDenseLayerMLP() - # messager = OneDenseLayerMLP() - - sample_obs = jnp.concatenate( - [ - v.flatten() - for v in env.single_observation_space.sample(rng=jax.random.PRNGKey(0)).values() - if v.size > 0 - ] - ) - sensor_params = sensor.init(sensor_key, sample_obs) - feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs) - actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs)) - critic_params = critic.init( - critic_key, feature_extractor.apply(feature_extractor_params, sample_obs) - ) - - agent_state = TrainState.create( - apply_fn=None, - params=asdict( - AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) - ), - tx=optax.chain( - optax.clip_by_global_norm(args.max_grad_norm), - optax.inject_hyperparams(optax.adam)( - learning_rate=linear_schedule if args.anneal_lr else args.learning_rate, eps=1e-5 - ), - ), - ) - - sensor.apply = jax.jit(sensor.apply) - feature_extractor.apply = jax.jit(feature_extractor.apply) - actor.apply = jax.jit(actor.apply) - critic.apply = jax.jit(critic.apply) - ppo_instance = PPO(args, sensor, actor, critic, feature_extractor) - - @jax.jit - def get_action_and_value_noise( - agent_state: TrainState, - next_obs: jnp.ndarray, - key: jax.random.PRNGKey, - ): - hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) - hidden_critic = feature_extractor.apply( - agent_state.params["feature_extractor_params"], next_obs - ) - - # Continuous actions: sample from a Gaussian parameterized by the actor - mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) - key, subkey = jax.random.split(key) - noise = jax.random.normal(subkey, shape=mean.shape) - std = jnp.exp(log_std) - action = mean + noise * std - logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) - value = critic.apply(agent_state.params["critic_params"], hidden_critic) - return action, logprob, value.squeeze(-1), key - - - # GAE - @jax.jit - def compute_gae_once(carry, inp, gamma, gae_lambda): - advantages = carry - nextdone, nextvalues, curvalues, reward = inp - nextnonterminal = 1.0 - nextdone - delta = reward + gamma * nextvalues * nextnonterminal - curvalues - advantages = delta + gamma * gae_lambda * nextnonterminal * advantages - return advantages, advantages - - @jax.jit - def compute_gae(agent_state, next_obs, next_done, storage): - next_value = critic.apply( - agent_state.params["critic_params"], - sensor.apply(agent_state.params["sensor_params"], next_obs), - ).squeeze(-1) - - advantages = jnp.zeros((args.num_envs,)) - dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0) - values = jnp.concatenate([storage.values, next_value[None, :]], axis=0) - _, advantages = jax.lax.scan( - partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda), - advantages, - (dones[1:], values[1:], values[:-1], storage.rewards), - reverse=True, - ) - return storage.replace(advantages=advantages, returns=advantages + storage.values) - # END GAE - - # --- Main training loop --- - global_step = 0 - start_time = time.time() - - # Reset once to get initial state - print("Resetting the environment...") - next_env_state = env.reset(seed=args.seed) - next_obs = convert_obs_dict_to_array(next_env_state.observations) - next_done = jnp.zeros(args.num_envs, dtype=jnp.bool_) - - def step_once(carry, _, env_step_fn): - agent_state, episode_stats, obs, done, key, env_state = carry - action, logprob, value, key = get_action_and_value_noise(agent_state, obs, key) - - episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( - episode_stats, env_state, action - ) - - storage = Storage( - obs=obs, - actions=action, - logprobs=logprob, - dones=done, - values=value, - rewards=reward, - returns=jnp.zeros_like(reward), - advantages=jnp.zeros_like(reward), - ) - return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage - - def rollout( - agent_state, episode_stats, next_obs, next_done, key, env_state, step_once_fn, max_steps - ): - (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( - step_once_fn, - (agent_state, episode_stats, next_obs, next_done, key, env_state), - (), - max_steps, - ) - return agent_state, episode_stats, next_obs, next_done, storage, key, env_state - - rollout = partial( - rollout, - step_once_fn=partial(step_once, env_step_fn=step_env_wrapped), - max_steps=args.num_steps, - ) - - print("Starting training...") - iters_bar = tqdm.tqdm(range(1, args.num_iterations + 1)) - returns = [] - for _ in iters_bar: - iteration_time_start = time.time() - - agent_state, episode_stats, next_obs, next_done, storage, key, next_env_state = rollout( - agent_state, episode_stats, next_obs, next_done, key, next_env_state - ) - - global_step += args.num_steps * args.num_envs - storage = compute_gae(agent_state, next_obs, next_done, storage) - agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo( - agent_state, storage, key - ) - - avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns)) - iters_bar.set_postfix_str( - f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" - ) - - returns.append(avg_episodic_return) - - writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) - writer.add_scalar( - "charts/avg_episodic_length", - np.mean(jax.device_get(episode_stats.returned_episode_lengths)), - global_step, - ) - writer.add_scalar( - "charts/learning_rate", - agent_state.opt_state[1].hyperparams["learning_rate"].item(), - global_step, - ) - writer.add_scalar("losses/value_loss", v_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/policy_loss", pg_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/entropy", entropy_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/approx_kl", approx_kl[-1, -1].item(), global_step) - writer.add_scalar("losses/loss", loss[-1, -1].item(), global_step) - - # iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}") - - writer.add_scalar("charts/SPS", int(global_step / (time.time() - start_time)), global_step) - writer.add_scalar( - "charts/SPS_update", - int(args.num_envs * args.num_steps / (time.time() - iteration_time_start)), - global_step, - ) - - if args.save_model: - model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model" - save_model(model_path, agent_state, args) - print(f"model saved to {model_path}") - - env.close() - writer.close() - - print("Saving loss plot...") - simple_plot( - list(range(len(returns))), - returns, - show_window=True, - filename=f"runs/{run_name}/{args.exp_name}_losses.png", - ) - - -def main() -> None: - args = tyro.cli(PPOArgs) - train(args) - - -if __name__ == "__main__": - main() + ppo_trainer = PPOTrainer(args, env, run_name) + ppo_trainer.train() diff --git a/experiments/train.py.back b/experiments/train.py.back new file mode 100644 index 0000000..04a3df4 --- /dev/null +++ b/experiments/train.py.back @@ -0,0 +1,350 @@ +import random +import time +from dataclasses import asdict +from functools import partial +from typing import Callable + +import flax +import jax +import jax.numpy as jnp +import numpy as np +import optax +import torch +import tqdm +import tyro +from flax.training.train_state import TrainState +from torch.utils.tensorboard import SummaryWriter + +from brittle_star_project.dataclasses import PPOArgs +from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics +from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +from MLPs.mlps import ( + GenericDenseLayersWithActivation, + Actor, + OneDenseLayerMLP, + AgentParams, + Storage, +) +from plots import simple_plot +from ppo import PPO + + +def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: + return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( + obs_dict + ) + + +def make_env(config_path: str | None, num_envs: int) -> Callable: + def thunk(): + if config_path is None: + return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) + return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs) + + return thunk + +def save_model(model_path: str, agent_state: TrainState, args: PPOArgs): + with open(model_path, "wb") as f: + f.write( + flax.serialization.to_bytes( + [ + vars(args), + [ + agent_state.params["network_params"], + agent_state.params["actor_params"], + agent_state.params["critic_params"], + ], + ] + ) + ) + + +def train(args: PPOArgs): + args.batch_size = args.num_envs * args.num_steps + args.minibatch_size = args.batch_size // args.num_minibatches + # args.num_iterations = args.total_timesteps // args.batch_size + args.num_iterations = 5 + run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}" + print(f"running name: {run_name}") + + if args.track: + import wandb + + wandb.init( + project=args.wandb_project_name, + entity=args.wandb_entity, + sync_tensorboard=True, + config=vars(args), + name=run_name, + save_code=True, + ) + + writer = SummaryWriter(f"runs/{run_name}") + writer.add_text( + "hyperparameters", + "|param|value|\n|---|---|\n" + "\n".join(f"|{k}|{v}|" for k, v in vars(args).items()), + ) + + random.seed(args.seed) + np.random.seed(args.seed) + key = jax.random.PRNGKey(args.seed) + key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5) + + torch.backends.cudnn.deterministic = args.torch_deterministic + device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") + print(f"Running on device: {device}") + + print("Creating the environment...") + env = make_env(config_path=args.config_path, num_envs=args.num_envs)() + print(f"Environment: {env}") + + episode_stats = EpisodeStatistics( + episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32), + episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), + returned_episode_returns=jnp.zeros(args.num_envs, jnp.float32), + returned_episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), + ) + + def step_env_wrapped(episode_stats: EpisodeStatistics, env_state, action): + next_env_state = env.step(env_state, action) + + # Extract per-environment signals from the state object + reward = next_env_state.reward # (num_envs,) + terminated = next_env_state.terminated # (num_envs,) + truncated = next_env_state.truncated # (num_envs,) + done = terminated | truncated # (num_envs,) + + new_episode_return = episode_stats.episode_returns + reward + new_episode_length = episode_stats.episode_lengths + 1 + + episode_stats = episode_stats.replace( + episode_returns=new_episode_return * (1 - done), + episode_lengths=new_episode_length * (1 - done), + returned_episode_returns=jnp.where( + done, new_episode_return, episode_stats.returned_episode_returns + ), + returned_episode_lengths=jnp.where( + done, new_episode_length, episode_stats.returned_episode_lengths + ), + ) + return ( + episode_stats, + next_env_state, + (convert_obs_dict_to_array(next_env_state.observations), reward, done), + ) + + def linear_schedule(count): + frac = 1.0 - (count // (args.num_minibatches * args.update_epochs)) / args.num_iterations + return args.learning_rate * frac + + print("Initializing the models...") + sensor = GenericDenseLayersWithActivation() + feature_extractor = GenericDenseLayersWithActivation() + actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX + critic = OneDenseLayerMLP() + # messager = OneDenseLayerMLP() + + sample_obs = jnp.concatenate( + [ + v.flatten() + for v in env.single_observation_space.sample(rng=jax.random.PRNGKey(0)).values() + if v.size > 0 + ] + ) + sensor_params = sensor.init(sensor_key, sample_obs) + feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs) + actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs)) + critic_params = critic.init( + critic_key, feature_extractor.apply(feature_extractor_params, sample_obs) + ) + + agent_state = TrainState.create( + apply_fn=None, + params=asdict( + AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) + ), + tx=optax.chain( + optax.clip_by_global_norm(args.max_grad_norm), + optax.inject_hyperparams(optax.adam)( + learning_rate=linear_schedule if args.anneal_lr else args.learning_rate, eps=1e-5 + ), + ), + ) + + sensor.apply = jax.jit(sensor.apply) + feature_extractor.apply = jax.jit(feature_extractor.apply) + actor.apply = jax.jit(actor.apply) + critic.apply = jax.jit(critic.apply) + ppo_instance = PPO(args, sensor, actor, critic, feature_extractor) + + @jax.jit + def get_action_and_value_noise( + agent_state: TrainState, + next_obs: jnp.ndarray, + key: jax.random.PRNGKey, + ): + hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) + hidden_critic = feature_extractor.apply( + agent_state.params["feature_extractor_params"], next_obs + ) + + # Continuous actions: sample from a Gaussian parameterized by the actor + mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) + key, subkey = jax.random.split(key) + noise = jax.random.normal(subkey, shape=mean.shape) + std = jnp.exp(log_std) + action = mean + noise * std + logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) + value = critic.apply(agent_state.params["critic_params"], hidden_critic) + return action, logprob, value.squeeze(-1), key + + + # GAE + @jax.jit + def compute_gae_once(carry, inp, gamma, gae_lambda): + advantages = carry + nextdone, nextvalues, curvalues, reward = inp + nextnonterminal = 1.0 - nextdone + delta = reward + gamma * nextvalues * nextnonterminal - curvalues + advantages = delta + gamma * gae_lambda * nextnonterminal * advantages + return advantages, advantages + + @jax.jit + def compute_gae(agent_state, next_obs, next_done, storage): + next_value = critic.apply( + agent_state.params["critic_params"], + sensor.apply(agent_state.params["sensor_params"], next_obs), + ).squeeze(-1) + + advantages = jnp.zeros((args.num_envs,)) + dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0) + values = jnp.concatenate([storage.values, next_value[None, :]], axis=0) + _, advantages = jax.lax.scan( + partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda), + advantages, + (dones[1:], values[1:], values[:-1], storage.rewards), + reverse=True, + ) + return storage.replace(advantages=advantages, returns=advantages + storage.values) + # END GAE + + # --- Main training loop --- + global_step = 0 + start_time = time.time() + + # Reset once to get initial state + print("Resetting the environment...") + next_env_state = env.reset(seed=args.seed) + next_obs = convert_obs_dict_to_array(next_env_state.observations) + next_done = jnp.zeros(args.num_envs, dtype=jnp.bool_) + + def step_once(carry, _, env_step_fn): + agent_state, episode_stats, obs, done, key, env_state = carry + action, logprob, value, key = get_action_and_value_noise(agent_state, obs, key) + + episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( + episode_stats, env_state, action + ) + + storage = Storage( + obs=obs, + actions=action, + logprobs=logprob, + dones=done, + values=value, + rewards=reward, + returns=jnp.zeros_like(reward), + advantages=jnp.zeros_like(reward), + ) + return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage + + def rollout( + agent_state, episode_stats, next_obs, next_done, key, env_state, step_once_fn, max_steps + ): + (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( + step_once_fn, + (agent_state, episode_stats, next_obs, next_done, key, env_state), + (), + max_steps, + ) + return agent_state, episode_stats, next_obs, next_done, storage, key, env_state + + rollout = partial( + rollout, + step_once_fn=partial(step_once, env_step_fn=step_env_wrapped), + max_steps=args.num_steps, + ) + + print("Starting training...") + iters_bar = tqdm.tqdm(range(1, args.num_iterations + 1)) + returns = [] + for _ in iters_bar: + iteration_time_start = time.time() + + agent_state, episode_stats, next_obs, next_done, storage, key, next_env_state = rollout( + agent_state, episode_stats, next_obs, next_done, key, next_env_state + ) + + global_step += args.num_steps * args.num_envs + storage = compute_gae(agent_state, next_obs, next_done, storage) + agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo( + agent_state, storage, key + ) + + avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns)) + iters_bar.set_postfix_str( + f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" + ) + + returns.append(avg_episodic_return) + + writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) + writer.add_scalar( + "charts/avg_episodic_length", + np.mean(jax.device_get(episode_stats.returned_episode_lengths)), + global_step, + ) + writer.add_scalar( + "charts/learning_rate", + agent_state.opt_state[1].hyperparams["learning_rate"].item(), + global_step, + ) + writer.add_scalar("losses/value_loss", v_loss[-1, -1].item(), global_step) + writer.add_scalar("losses/policy_loss", pg_loss[-1, -1].item(), global_step) + writer.add_scalar("losses/entropy", entropy_loss[-1, -1].item(), global_step) + writer.add_scalar("losses/approx_kl", approx_kl[-1, -1].item(), global_step) + writer.add_scalar("losses/loss", loss[-1, -1].item(), global_step) + + # iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}") + + writer.add_scalar("charts/SPS", int(global_step / (time.time() - start_time)), global_step) + writer.add_scalar( + "charts/SPS_update", + int(args.num_envs * args.num_steps / (time.time() - iteration_time_start)), + global_step, + ) + + if args.save_model: + model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model" + save_model(model_path, agent_state, args) + print(f"model saved to {model_path}") + + env.close() + writer.close() + + print("Saving loss plot...") + simple_plot( + list(range(len(returns))), + returns, + show_window=True, + filename=f"runs/{run_name}/{args.exp_name}_losses.png", + ) + + +def main() -> None: + args = tyro.cli(PPOArgs) + train(args) + + +if __name__ == "__main__": + main() From a9302cfd4eb4c031479e7b1bd36dece07511b3e2 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Fri, 3 Apr 2026 14:54:57 +0200 Subject: [PATCH 06/19] fix(PPOTrainer.py): cleaned up + bug fixes regarding misuse of variable/wrong returns --- docs/api/simulate.md | 14 ++ docs/api/train.md | 1 + docs/api/train_simulate.md | 30 --- experiments/PPOTrainer.py | 173 ++++++++++++------ .../{train.py.back => _train_backup.py} | 0 experiments/train.py | 3 + 6 files changed, 133 insertions(+), 88 deletions(-) create mode 100644 docs/api/simulate.md create mode 100644 docs/api/train.md delete mode 100644 docs/api/train_simulate.md rename experiments/{train.py.back => _train_backup.py} (100%) diff --git a/docs/api/simulate.md b/docs/api/simulate.md new file mode 100644 index 0000000..9c4b21b --- /dev/null +++ b/docs/api/simulate.md @@ -0,0 +1,14 @@ +# Training and Simulation for Brittle Star Models + +## Simulating a model + +In order to simulate and view the behavior of a trained model, you can use the `simulate.py` script. This script allows you to specify the path to a trained model and will launch a simulation using that model. This script has the following parameters: + +- `--model`: The path to the trained model artifact to simulate. +- `--model-type`: The type of model to simulate (e.g., `random`, ...) +- `--task`: The task to simulate (e.g., `directed_locomotion`, ...) +- `--seed`: The random seed for reproducibility. + +```bash +python simulate.py --model artifacts/my_model --model-type random --task directed_locomotion --seed 0 +``` \ No newline at end of file diff --git a/docs/api/train.md b/docs/api/train.md new file mode 100644 index 0000000..f87f5c1 --- /dev/null +++ b/docs/api/train.md @@ -0,0 +1 @@ +# TODO \ No newline at end of file diff --git a/docs/api/train_simulate.md b/docs/api/train_simulate.md deleted file mode 100644 index 9b299a9..0000000 --- a/docs/api/train_simulate.md +++ /dev/null @@ -1,30 +0,0 @@ -# Training and Simulation for Brittle Star Models - -## Training a model - -To train a model, you can use the `train.py` script. This script allows to pass some parameters to customize the training process: - -- `--out`: The output path where the trained model will be saved. -- `--model_type`: The type of model to train (e.g., `random`, ...) -- `--task`: The task to train on (e.g., `directed_locomotion`, ...) -- `--seed`: The random seed for reproducibility. -- `--epochs`: The number of epochs to train for. - -This will then train the specified model on the specified task for the given number of epochs and save the trained model to the specified output path. - -```bash -python train.py --out artifacts/my_model --model-type random --task directed_locomotion --seed 0 --epochs 50 -``` - -## Simulating a model - -In order to simulate and view the behavior of a trained model, you can use the `simulate.py` script. This script allows you to specify the path to a trained model and will launch a simulation using that model. This script has the following parameters: - -- `--model`: The path to the trained model artifact to simulate. -- `--model-type`: The type of model to simulate (e.g., `random`, ...) -- `--task`: The task to simulate (e.g., `directed_locomotion`, ...) -- `--seed`: The random seed for reproducibility. - -```bash -python simulate.py --model artifacts/my_model --model-type random --task directed_locomotion --seed 0 -``` \ No newline at end of file diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index 347b724..42d064f 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -1,26 +1,28 @@ +import random import time from dataclasses import asdict, dataclass from functools import partial from typing import Any -import optax -import tqdm -from flax.metrics.tensorboard import SummaryWriter -from flax.training.train_state import TrainState - -from MLPs.mlps import ( - GenericDenseLayersWithActivation, - Actor, - OneDenseLayerMLP, - AgentParams, - Storage, -) -from brittle_star_project.dataclasses import PPOArgs, EpisodeStatistics -from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +import flax import jax import jax.numpy as jnp import numpy as np +import optax +import torch +import tqdm +from flax.training.train_state import TrainState +from torch.utils.tensorboard import SummaryWriter +from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs +from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +from MLPs.mlps import ( + Actor, + AgentParams, + GenericDenseLayersWithActivation, + OneDenseLayerMLP, + Storage, +) from ppo import PPO @@ -95,6 +97,36 @@ def _step_once( return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage +@jax.jit +def _step_env_wrapped(env_step_fn, env_state, action, episode_stats): + next_env_state = env_step_fn(env_state, action) + + # Extract per-environment signals from the state object + reward = next_env_state.reward # (num_envs,) + terminated = next_env_state.terminated # (num_envs,) + truncated = next_env_state.truncated # (num_envs,) + done = terminated | truncated # (num_envs,) + + new_episode_return = episode_stats.episode_returns + reward + new_episode_length = episode_stats.episode_lengths + 1 + + episode_stats = episode_stats.replace( + episode_returns=new_episode_return * (1 - done), + episode_lengths=new_episode_length * (1 - done), + returned_episode_returns=jnp.where( + done, new_episode_return, episode_stats.returned_episode_returns + ), + returned_episode_lengths=jnp.where( + done, new_episode_length, episode_stats.returned_episode_lengths + ), + ) + return ( + episode_stats, + next_env_state, + (convert_obs_dict_to_array(next_env_state.observations), reward, done), + ) + + @jax.jit def _rollout_jit( agent_state, @@ -126,36 +158,6 @@ def _rollout_jit( return agent_state, episode_stats, next_obs, next_done, storage, key, env_state -@jax.jit -def _step_env_wrapped(env_step_fn, env_state, action, episode_stats): - next_env_state = env_step_fn(env_state, action) - - # Extract per-environment signals from the state object - reward = next_env_state.reward # (num_envs,) - terminated = next_env_state.terminated # (num_envs,) - truncated = next_env_state.truncated # (num_envs,) - done = terminated | truncated # (num_envs,) - - new_episode_return = episode_stats.episode_returns + reward - new_episode_length = episode_stats.episode_lengths + 1 - - episode_stats = episode_stats.replace( - episode_returns=new_episode_return * (1 - done), - episode_lengths=new_episode_length * (1 - done), - returned_episode_returns=jnp.where( - done, new_episode_return, episode_stats.returned_episode_returns - ), - returned_episode_lengths=jnp.where( - done, new_episode_length, episode_stats.returned_episode_lengths - ), - ) - return ( - episode_stats, - next_env_state, - (convert_obs_dict_to_array(next_env_state.observations), reward, done), - ) - - @jax.jit def compute_gae_once(carry, inp, gamma, gae_lambda): advantages = carry @@ -200,7 +202,8 @@ class PPOTrainer: def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_name: str): self.args = args self.env = env - self.writer = SummaryWriter(f"runs/{run_name}") + self.run_name = run_name + self.writer = SummaryWriter(f"runs/{self.run_name}") self.key = jax.random.PRNGKey(args.seed) @@ -216,6 +219,12 @@ class PPOTrainer: self.episode_stats = self._init_episode_stats() + self._init_random() + + def _init_random(self): + random.seed(self.args.seed) + np.random.seed(self.args.seed) + def _init_agent(self): sensor = GenericDenseLayersWithActivation() feature_extractor = GenericDenseLayersWithActivation() @@ -255,7 +264,13 @@ class PPOTrainer: tx=optax.chain( optax.clip_by_global_norm(self.args.max_grad_norm), optax.inject_hyperparams(optax.adam)( - learning_rate=linear_schedule + learning_rate=partial( + linear_schedule, + minibatch_count=self.args.num_minibatches, + update_epochs=self.args.update_epochs, + num_iterations=self.args.num_iterations, + learning_rate=self.args.learning_rate, + ) if self.args.anneal_lr else self.args.learning_rate, eps=1e-5, @@ -296,12 +311,13 @@ class PPOTrainer: self, global_step, episode_stats, - avg_episodic_return, start_time, iteration_time_start, loss_info, ): - self.writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) + self.writer.add_scalar( + "charts/avg_episodic_return", loss_info.avg_episodic_return, global_step + ) self.writer.add_scalar( "charts/avg_episodic_length", np.mean(jax.device_get(episode_stats.returned_episode_lengths)), @@ -318,8 +334,6 @@ class PPOTrainer: 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) - # iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}") - self.writer.add_scalar( "charts/SPS", int(global_step / (time.time() - start_time)), global_step ) @@ -330,8 +344,18 @@ class PPOTrainer: ) def _step(self, env_state, next_obs, next_done) -> tuple: - storage, next_obs, next_done, env_state = self._rollout(env_state, next_obs, next_done) + ( + self.agent_state, + self.episode_stats, + next_obs, + next_done, + storage, + self.key, + next_env_state, + ) = self._rollout(env_state, next_obs, next_done) + storage = self._compute_gae(storage, next_obs, next_done) + self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = ( self._ppo.update_ppo(self.agent_state, storage, self.key) ) @@ -339,7 +363,7 @@ class PPOTrainer: avg_episodic_return = jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)) return ( - env_state, + next_env_state, next_obs, next_done, LossInfo( @@ -352,10 +376,26 @@ class PPOTrainer: ), ) - def close(self): + def _close(self): self.env.close() self.writer.close() + def _save_model(self, model_path: str): + with open(model_path, "wb") as f: + f.write( + flax.serialization.to_bytes( + [ + vars(self.args), + [ + self.agent_state.params["sensor_params"], + self.agent_state.params["actor_params"], + self.agent_state.params["critic_params"], + self.agent_state.params["feature_extractor_params"], + ], + ] + ) + ) + def train(self): """ Train the PPO agent for a specified number of iterations @@ -368,16 +408,33 @@ class PPOTrainer: global_step = 0 start_time = time.time() + if self.args.track: + import wandb + + 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, + ) + + self.writer.add_text( + "hyperparameters", + "|param|value|\n|---|---|\n" + + "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()), + ) + for _ in tqdm.tqdm(range(self.args.num_iterations)): iteration_time_start = time.time() env_state, next_obs, next_done, loss_info = self._step(env_state, next_obs, next_done) global_step += self.args.num_steps * self.args.num_envs - self._log( - global_step, self.episode_stats, 0, start_time, iteration_time_start, loss_info - ) + self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info) if self.args.save_model: - self._save_model(...) + model_path = f"runs/{self.run_name}/{self.args.exp_name}.cleanrl_model" + self._save_model(model_path=model_path) - self.close() + self._close() diff --git a/experiments/train.py.back b/experiments/_train_backup.py similarity index 100% rename from experiments/train.py.back rename to experiments/_train_backup.py diff --git a/experiments/train.py b/experiments/train.py index ba9dfcb..716a917 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -1,5 +1,6 @@ import time +import torch import tyro from brittle_star_project.dataclasses import PPOArgs @@ -23,5 +24,7 @@ if __name__ == "__main__": run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}" env = make_env(args.config_path, args.num_envs) + torch.backends.cudnn.deterministic = args.torch_deterministic + ppo_trainer = PPOTrainer(args, env, run_name) ppo_trainer.train() From 62e11ce2430f87ce6ed9b6efddf5fe03e0f2d81c Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Fri, 3 Apr 2026 15:22:51 +0200 Subject: [PATCH 07/19] fix(PPOTrainer.py): bug fixes regarding jax.jit --- experiments/PPOTrainer.py | 59 +++++++++++++------ .../dataclasses/PPOArgs.py | 3 + 2 files changed, 44 insertions(+), 18 deletions(-) diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index 42d064f..fcb49a5 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -40,7 +40,7 @@ def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: @jax.jit -def get_action_and_value_noise( +def _get_action_and_value_noise( sensor: GenericDenseLayersWithActivation, feature_extractor: GenericDenseLayersWithActivation, actor: Actor, @@ -76,7 +76,7 @@ def _step_once( critic: OneDenseLayerMLP, ): agent_state, episode_stats, obs, done, key, env_state = carry - action, logprob, value, key = get_action_and_value_noise( + action, logprob, value, key = _get_action_and_value_noise( sensor, feature_extractor, actor, critic, agent_state, obs, key ) @@ -127,7 +127,10 @@ def _step_env_wrapped(env_step_fn, env_state, action, episode_stats): ) -@jax.jit +@partial( + jax.jit, + static_argnames=("env_step_fn", "sensor", "feature_extractor", "actor", "critic"), +) def _rollout_jit( agent_state, episode_stats, @@ -135,7 +138,7 @@ def _rollout_jit( next_obs, next_done, key, - args, + max_steps, env_step_fn, sensor: GenericDenseLayersWithActivation, feature_extractor: GenericDenseLayersWithActivation, @@ -153,13 +156,13 @@ def _rollout_jit( ), (agent_state, episode_stats, next_obs, next_done, key, env_state), (), - args.max_steps, + max_steps, ) return agent_state, episode_stats, next_obs, next_done, storage, key, env_state @jax.jit -def compute_gae_once(carry, inp, gamma, gae_lambda): +def _compute_gae_once(carry, inp, gamma, gae_lambda): advantages = carry nextdone, nextvalues, curvalues, reward = inp nextnonterminal = 1.0 - nextdone @@ -168,18 +171,23 @@ def compute_gae_once(carry, inp, gamma, gae_lambda): return advantages, advantages -@jax.jit -def compute_gae_jit(agent_state, storage, next_obs, next_done, sensor, critic, args): +@partial( + jax.jit, + static_argnames=("sensor", "critic"), +) +def _compute_gae_jit( + agent_state, storage, next_obs, next_done, sensor, critic, gamma, gae_lambda, num_envs +): next_value = critic.apply( agent_state.params["critic_params"], sensor.apply(agent_state.params["sensor_params"], next_obs), ).squeeze(-1) - advantages = jnp.zeros((args.num_envs,)) + advantages = jnp.zeros((num_envs,)) dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0) values = jnp.concatenate([storage.values, next_value[None, :]], axis=0) _, advantages = jax.lax.scan( - partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda), + partial(_compute_gae_once, gamma=gamma, gae_lambda=gae_lambda), advantages, (dones[1:], values[1:], values[:-1], storage.rewards), reverse=True, @@ -213,6 +221,19 @@ class PPOTrainer: self.actor.apply = jax.jit(self.actor.apply) self.critic.apply = jax.jit(self.critic.apply) + self._rollout_jit = partial( + _rollout_jit, + sensor=self.sensor, + feature_extractor=self.feature_extractor, + actor=self.actor, + critic=self.critic, + ) + self._compute_gae_jit = partial( + _compute_gae_jit, + sensor=self.sensor, + critic=self.critic, + ) + self._ppo = PPO(self.args, self.sensor, self.actor, self.critic, self.feature_extractor) self.agent_state = self._init_agent_state() @@ -287,24 +308,26 @@ class PPOTrainer: ) def _rollout(self, env_state, next_obs, next_done) -> tuple[Storage, ...]: - return _rollout_jit( + return self._rollout_jit( self.agent_state, self.episode_stats, env_state, next_obs, next_done, self.key, - self.args, + self.args.num_steps, self.env.step, - self.sensor, - self.feature_extractor, - self.actor, - self.critic, ) def _compute_gae(self, storage, next_obs, next_done) -> Storage: - return compute_gae_jit( - self.agent_state, storage, next_obs, next_done, self.sensor, self.critic, self.args + return self._compute_gae_jit( + self.agent_state, + storage, + next_obs, + next_done, + self.args.gamma, + self.args.gae_lambda, + self.args.num_envs, ) def _log( diff --git a/src/brittle_star_project/dataclasses/PPOArgs.py b/src/brittle_star_project/dataclasses/PPOArgs.py index 03c220e..8d13148 100644 --- a/src/brittle_star_project/dataclasses/PPOArgs.py +++ b/src/brittle_star_project/dataclasses/PPOArgs.py @@ -1,6 +1,9 @@ from dataclasses import dataclass +import jax + +@jax.tree_util.register_dataclass @dataclass class PPOArgs: """ From ca175742724ffd1e8f097a80da3beb9ec3a7f6ef Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 11:53:22 +0200 Subject: [PATCH 08/19] fix(PPOTrainer.py): finished all jit-related bugs --- experiments/PPOTrainer.py | 62 +++++++++++++++++++-------------------- 1 file changed, 30 insertions(+), 32 deletions(-) diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index fcb49a5..6e8b6f7 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -39,7 +39,7 @@ def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: ) -@jax.jit +# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit def _get_action_and_value_noise( sensor: GenericDenseLayersWithActivation, feature_extractor: GenericDenseLayersWithActivation, @@ -65,7 +65,7 @@ def _get_action_and_value_noise( return action, logprob, value.squeeze(-1), key -@jax.jit +# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit def _step_once( carry, _, @@ -97,8 +97,8 @@ def _step_once( return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage -@jax.jit -def _step_env_wrapped(env_step_fn, env_state, action, episode_stats): +# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit +def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): next_env_state = env_step_fn(env_state, action) # Extract per-environment signals from the state object @@ -127,10 +127,7 @@ def _step_env_wrapped(env_step_fn, env_state, action, episode_stats): ) -@partial( - jax.jit, - static_argnames=("env_step_fn", "sensor", "feature_extractor", "actor", "critic"), -) +# jit applied in wrapper method self._rollout_jit using partial def _rollout_jit( agent_state, episode_stats, @@ -139,7 +136,7 @@ def _rollout_jit( next_done, key, max_steps, - env_step_fn, + step_env_fn, sensor: GenericDenseLayersWithActivation, feature_extractor: GenericDenseLayersWithActivation, actor: Actor, @@ -152,7 +149,7 @@ def _rollout_jit( feature_extractor=feature_extractor, actor=actor, critic=critic, - env_step_fn=partial(_step_env_wrapped, env_step_fn=env_step_fn), + env_step_fn=step_env_fn, ), (agent_state, episode_stats, next_obs, next_done, key, env_state), (), @@ -161,7 +158,7 @@ def _rollout_jit( return agent_state, episode_stats, next_obs, next_done, storage, key, env_state -@jax.jit +# removed jit: used in _compute_gae_jit, so will be compiled with _compute_gae_jit def _compute_gae_once(carry, inp, gamma, gae_lambda): advantages = carry nextdone, nextvalues, curvalues, reward = inp @@ -171,12 +168,9 @@ def _compute_gae_once(carry, inp, gamma, gae_lambda): return advantages, advantages -@partial( - jax.jit, - static_argnames=("sensor", "critic"), -) +# jit applied on partial-wrapped wrapper method self._compute_gae_jit def _compute_gae_jit( - agent_state, storage, next_obs, next_done, sensor, critic, gamma, gae_lambda, num_envs + agent_state, storage, next_obs, next_done, gamma, gae_lambda, num_envs, sensor, critic ): next_value = critic.apply( agent_state.params["critic_params"], @@ -221,17 +215,26 @@ class PPOTrainer: self.actor.apply = jax.jit(self.actor.apply) self.critic.apply = jax.jit(self.critic.apply) - self._rollout_jit = partial( - _rollout_jit, - sensor=self.sensor, - feature_extractor=self.feature_extractor, - actor=self.actor, - critic=self.critic, + self._rollout_jit = jax.jit( + partial( + _rollout_jit, + max_steps=self.args.num_steps, + step_env_fn=partial(_step_env_wrapped, env_step_fn=self.env.step), + sensor=self.sensor, + feature_extractor=self.feature_extractor, + actor=self.actor, + critic=self.critic, + ) ) - self._compute_gae_jit = partial( - _compute_gae_jit, - sensor=self.sensor, - critic=self.critic, + self._compute_gae_jit = jax.jit( + partial( + _compute_gae_jit, + num_envs=self.args.num_envs, + gamma=self.args.gamma, + gae_lambda=self.args.gae_lambda, + sensor=self.sensor, + critic=self.critic, + ) ) self._ppo = PPO(self.args, self.sensor, self.actor, self.critic, self.feature_extractor) @@ -315,8 +318,6 @@ class PPOTrainer: next_obs, next_done, self.key, - self.args.num_steps, - self.env.step, ) def _compute_gae(self, storage, next_obs, next_done) -> Storage: @@ -325,9 +326,6 @@ class PPOTrainer: storage, next_obs, next_done, - self.args.gamma, - self.args.gae_lambda, - self.args.num_envs, ) def _log( @@ -383,7 +381,7 @@ class PPOTrainer: self._ppo.update_ppo(self.agent_state, storage, self.key) ) - avg_episodic_return = jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)) + avg_episodic_return = float(jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))) return ( next_env_state, From 3c7670e1b4f3641d8c3d6464c4a05b2fcf607d2f Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 11:54:07 +0200 Subject: [PATCH 09/19] fix(train.py): replaced debugging iteration count with dynamically calculated count --- experiments/train.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/experiments/train.py b/experiments/train.py index 716a917..2e4bca7 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -19,8 +19,7 @@ if __name__ == "__main__": args.batch_size = args.num_envs * args.num_steps args.minibatch_size = args.batch_size // args.num_minibatches - # args.num_iterations = args.total_timesteps // args.batch_size - args.num_iterations = 5 + 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) From 9c5203fc7d66af0e52d6e0a42be3f5302263d85b Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 11:55:05 +0200 Subject: [PATCH 10/19] fix: formatting + linting --- experiments/PPOTrainer.py | 5 +++-- experiments/_train_backup.py | 3 ++- experiments/plots/__init__.py | 2 +- 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index 6e8b6f7..aea69d2 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -9,7 +9,6 @@ import jax import jax.numpy as jnp import numpy as np import optax -import torch import tqdm from flax.training.train_state import TrainState from torch.utils.tensorboard import SummaryWriter @@ -381,7 +380,9 @@ class PPOTrainer: self._ppo.update_ppo(self.agent_state, storage, self.key) ) - avg_episodic_return = float(jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))) + avg_episodic_return = float( + jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)) + ) return ( next_env_state, diff --git a/experiments/_train_backup.py b/experiments/_train_backup.py index 04a3df4..6d6a181 100644 --- a/experiments/_train_backup.py +++ b/experiments/_train_backup.py @@ -43,6 +43,7 @@ def make_env(config_path: str | None, num_envs: int) -> Callable: return thunk + def save_model(model_path: str, agent_state: TrainState, args: PPOArgs): with open(model_path, "wb") as f: f.write( @@ -198,7 +199,6 @@ def train(args: PPOArgs): value = critic.apply(agent_state.params["critic_params"], hidden_critic) return action, logprob, value.squeeze(-1), key - # GAE @jax.jit def compute_gae_once(carry, inp, gamma, gae_lambda): @@ -226,6 +226,7 @@ def train(args: PPOArgs): reverse=True, ) return storage.replace(advantages=advantages, returns=advantages + storage.values) + # END GAE # --- Main training loop --- diff --git a/experiments/plots/__init__.py b/experiments/plots/__init__.py index 5abf51e..92f34cb 100644 --- a/experiments/plots/__init__.py +++ b/experiments/plots/__init__.py @@ -1,3 +1,3 @@ from .plot import simple_plot -__all__ = ["simple_plot"] \ No newline at end of file +__all__ = ["simple_plot"] From b97fc30d5d05d98a3ca8b336851417568b96b89f Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 13:56:54 +0200 Subject: [PATCH 11/19] feat(PPOTrainer.py): added logging messages --- configs/hpc/smoke_test.yaml | 2 +- experiments/PPOTrainer.py | 97 +++++++++++++++++++++++++++++++----- experiments/_train_backup.py | 9 ---- experiments/train.py | 48 ++++++++++++++++-- uv.lock | 6 ++- 5 files changed, 135 insertions(+), 27 deletions(-) diff --git a/configs/hpc/smoke_test.yaml b/configs/hpc/smoke_test.yaml index 720f829..c892e3e 100644 --- a/configs/hpc/smoke_test.yaml +++ b/configs/hpc/smoke_test.yaml @@ -1,5 +1,5 @@ # 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" seed: 0 track: false # Test WandB integration diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index aea69d2..cfadb54 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -1,4 +1,6 @@ +import datetime import random +import sys import time from dataclasses import asdict, dataclass from functools import partial @@ -200,11 +202,12 @@ class LossInfo: 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.env = env + self.run_dir = run_dir 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) @@ -244,11 +247,17 @@ class PPOTrainer: 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) 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() feature_extractor = GenericDenseLayersWithActivation() actor = Actor( @@ -258,7 +267,10 @@ class PPOTrainer: # messenger = OneDenseLayerMLP() 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, 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( episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32), episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32), @@ -335,6 +350,7 @@ class PPOTrainer: iteration_time_start, loss_info, ): + self.writer.add_scalar( "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/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 ) @@ -363,7 +378,10 @@ class PPOTrainer: 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.episode_stats, @@ -374,12 +392,21 @@ class PPOTrainer: next_env_state, ) = 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) + 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._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( jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)) ) @@ -402,7 +429,10 @@ class PPOTrainer: self.env.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: f.write( 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 (passed through PPOArgs in constructor). 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) next_obs = convert_obs_dict_to_array(env_state.observations) 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 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, @@ -442,21 +489,47 @@ class PPOTrainer: 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()), ) - 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() 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 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: - 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._close() diff --git a/experiments/_train_backup.py b/experiments/_train_backup.py index 2cc38ab..a46e500 100644 --- a/experiments/_train_backup.py +++ b/experiments/_train_backup.py @@ -81,7 +81,6 @@ def train(args: PPOArgs): 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 = 5 - run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}" print(f"running name: {run_name}") if args.run_dir is None: @@ -385,8 +384,6 @@ def train(args: PPOArgs): ) 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" with open(model_path, "wb") as f: f.write( @@ -408,12 +405,6 @@ def train(args: PPOArgs): writer.close() 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.title("PPO Loss, mean over minibatches") plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png") diff --git a/experiments/train.py b/experiments/train.py index 2e4bca7..596dad9 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -1,7 +1,10 @@ +import subprocess import time import torch import tyro +import yaml +import os from brittle_star_project.dataclasses import PPOArgs 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) +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__": - args = tyro.cli(PPOArgs) + args = parse_args() args.batch_size = args.num_envs * args.num_steps args.minibatch_size = args.batch_size // args.num_minibatches 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 - ppo_trainer = PPOTrainer(args, env, run_name) + ppo_trainer = PPOTrainer(args, env, run_dir, run_name) ppo_trainer.train() diff --git a/uv.lock b/uv.lock index 7fbaacb..fd45cfc 100644 --- a/uv.lock +++ b/uv.lock @@ -35,6 +35,9 @@ dependencies = [ ] [package.optional-dependencies] +analysis = [ + { name = "tensorboard" }, +] cuda = [ { name = "jax", extra = ["cuda13"] }, ] @@ -64,12 +67,13 @@ requires-dist = [ { name = "protobuf", specifier = ">=5.0.0" }, { name = "pyopengl", specifier = ">=3.1.10" }, { name = "pyopengl-accelerate", specifier = ">=3.1.10" }, + { name = "tensorboard", marker = "extra == 'analysis'" }, { name = "torch", specifier = ">=2.4.0" }, { name = "tyro", specifier = ">=1.0.10" }, { name = "wandb", specifier = "==0.24.2" }, { name = "warp-lang" }, ] -provides-extras = ["cuda"] +provides-extras = ["cuda", "analysis"] [package.metadata.requires-dev] dev = [ From d5b911fa32af4afe7fa602dba34752bae032fd32 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 13:58:08 +0200 Subject: [PATCH 12/19] feat(experiments): renamed _train_backup.py to _train_backup.py.back to remove duplicated code warning in pycharm --- experiments/_train_backup.py | 435 ----------------------------------- 1 file changed, 435 deletions(-) delete mode 100644 experiments/_train_backup.py diff --git a/experiments/_train_backup.py b/experiments/_train_backup.py deleted file mode 100644 index a46e500..0000000 --- a/experiments/_train_backup.py +++ /dev/null @@ -1,435 +0,0 @@ -import datetime -import random -import yaml -import subprocess -import sys -import time -from dataclasses import asdict -from functools import partial -from typing import Callable - -import flax -import jax -import jax.numpy as jnp -import matplotlib.pyplot as plt -import numpy as np -import optax -import torch -import tqdm -import tyro -from flax.training.train_state import TrainState -from torch.utils.tensorboard import SummaryWriter - -from brittle_star_project.dataclasses import PPOArgs -from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics -from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper -from MLPs.mlps import ( - GenericDenseLayersWithActivation, - Actor, - OneDenseLayerMLP, - AgentParams, - Storage, -) -from plots import simple_plot -from ppo import PPO - - -def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: - return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( - obs_dict - ) - - -def make_env(env_config_path: str | None, num_envs: int) -> Callable: - def thunk(): - if env_config_path is None: - return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) - return BrittleStarJaxEnvWrapper.from_config(env_config_path, num_envs=num_envs) - - return thunk - - -def save_model(model_path: str, agent_state: TrainState, args: PPOArgs): - with open(model_path, "wb") as f: - f.write( - flax.serialization.to_bytes( - [ - vars(args), - [ - agent_state.params["network_params"], - agent_state.params["actor_params"], - agent_state.params["critic_params"], - ], - ] - ) - ) - - -def train(args: PPOArgs): - args.batch_size = args.num_envs * args.num_steps - args.minibatch_size = args.batch_size // args.num_minibatches - args.num_iterations = args.total_timesteps // args.batch_size - - # Try to get git short hash - try: - git_hash = ( - subprocess.check_output(["git", "rev-parse", "--short", "HEAD"]).decode("ascii").strip() - ) - except Exception: - git_hash = "none" - - 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 = 5 - print(f"running name: {run_name}") - - if args.run_dir is None: - args.run_dir = f"runs/{run_name}" - - import os - - os.makedirs(args.run_dir, exist_ok=True) - - if args.track: - import wandb - - wandb.init( - project=args.wandb_project_name, - entity=args.wandb_entity, - sync_tensorboard=True, - config=vars(args), - name=run_name, - save_code=True, - ) - - writer = SummaryWriter(args.run_dir) - writer.add_text( - "hyperparameters", - "|param|value|\n|---|---|\n" + "\n".join(f"|{k}|{v}|" for k, v in vars(args).items()), - ) - - random.seed(args.seed) - np.random.seed(args.seed) - key = jax.random.PRNGKey(args.seed) - key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5) - - torch.backends.cudnn.deterministic = args.torch_deterministic - device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") - print(f"Running on device: {device}") - - print("Creating the environment...") - env = make_env(env_config_path=args.env_config_path, num_envs=args.num_envs)() - print(f"Environment: {env}") - - episode_stats = EpisodeStatistics( - episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32), - episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), - returned_episode_returns=jnp.zeros(args.num_envs, jnp.float32), - returned_episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), - ) - - def step_env_wrapped(episode_stats: EpisodeStatistics, env_state, action): - next_env_state = env.step(env_state, action) - - # Extract per-environment signals from the state object - reward = next_env_state.reward # (num_envs,) - terminated = next_env_state.terminated # (num_envs,) - truncated = next_env_state.truncated # (num_envs,) - done = terminated | truncated # (num_envs,) - - new_episode_return = episode_stats.episode_returns + reward - new_episode_length = episode_stats.episode_lengths + 1 - - episode_stats = episode_stats.replace( - episode_returns=new_episode_return * (1 - done), - episode_lengths=new_episode_length * (1 - done), - returned_episode_returns=jnp.where( - done, new_episode_return, episode_stats.returned_episode_returns - ), - returned_episode_lengths=jnp.where( - done, new_episode_length, episode_stats.returned_episode_lengths - ), - ) - return ( - episode_stats, - next_env_state, - (convert_obs_dict_to_array(next_env_state.observations), reward, done), - ) - - def linear_schedule(count): - frac = 1.0 - (count // (args.num_minibatches * args.update_epochs)) / args.num_iterations - return args.learning_rate * frac - - print("Initializing the models...") - sensor = GenericDenseLayersWithActivation() - feature_extractor = GenericDenseLayersWithActivation() - actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX - critic = OneDenseLayerMLP() - # messager = OneDenseLayerMLP() - - sample_obs = jnp.concatenate( - [ - v.flatten() - for v in env.single_observation_space.sample(rng=jax.random.PRNGKey(0)).values() - if v.size > 0 - ] - ) - sensor_params = sensor.init(sensor_key, sample_obs) - feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs) - actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs)) - critic_params = critic.init( - critic_key, feature_extractor.apply(feature_extractor_params, sample_obs) - ) - - agent_state = TrainState.create( - apply_fn=None, - params=asdict( - AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) - ), - tx=optax.chain( - optax.clip_by_global_norm(args.max_grad_norm), - optax.inject_hyperparams(optax.adam)( - learning_rate=linear_schedule if args.anneal_lr else args.learning_rate, eps=1e-5 - ), - ), - ) - - sensor.apply = jax.jit(sensor.apply) - feature_extractor.apply = jax.jit(feature_extractor.apply) - actor.apply = jax.jit(actor.apply) - critic.apply = jax.jit(critic.apply) - ppo_instance = PPO(args, sensor, actor, critic, feature_extractor) - - @jax.jit - def get_action_and_value_noise( - agent_state: TrainState, - next_obs: jnp.ndarray, - key: jax.random.PRNGKey, - ): - hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) - hidden_critic = feature_extractor.apply( - agent_state.params["feature_extractor_params"], next_obs - ) - - # Continuous actions: sample from a Gaussian parameterized by the actor - mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) - key, subkey = jax.random.split(key) - noise = jax.random.normal(subkey, shape=mean.shape) - std = jnp.exp(log_std) - action = mean + noise * std - logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) - value = critic.apply(agent_state.params["critic_params"], hidden_critic) - return action, logprob, value.squeeze(-1), key - - # GAE - @jax.jit - def compute_gae_once(carry, inp, gamma, gae_lambda): - advantages = carry - nextdone, nextvalues, curvalues, reward = inp - nextnonterminal = 1.0 - nextdone - delta = reward + gamma * nextvalues * nextnonterminal - curvalues - advantages = delta + gamma * gae_lambda * nextnonterminal * advantages - return advantages, advantages - - @jax.jit - def compute_gae(agent_state, next_obs, next_done, storage): - next_value = critic.apply( - agent_state.params["critic_params"], - sensor.apply(agent_state.params["sensor_params"], next_obs), - ).squeeze(-1) - - advantages = jnp.zeros((args.num_envs,)) - dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0) - values = jnp.concatenate([storage.values, next_value[None, :]], axis=0) - _, advantages = jax.lax.scan( - partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda), - advantages, - (dones[1:], values[1:], values[:-1], storage.rewards), - reverse=True, - ) - return storage.replace(advantages=advantages, returns=advantages + storage.values) - - # END GAE - - # --- Main training loop --- - global_step = 0 - start_time = time.time() - - # Reset once to get initial state - print("Resetting the environment...") - if not sys.stdout.isatty(): - print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True) - next_env_state = env.reset(seed=args.seed) - next_obs = convert_obs_dict_to_array(next_env_state.observations) - next_done = jnp.zeros(args.num_envs, dtype=jnp.bool_) - if not sys.stdout.isatty(): - print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True) - - def step_once(carry, _, env_step_fn): - agent_state, episode_stats, obs, done, key, env_state = carry - action, logprob, value, key = get_action_and_value_noise(agent_state, obs, key) - - episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( - episode_stats, env_state, action - ) - - storage = Storage( - obs=obs, - actions=action, - logprobs=logprob, - dones=done, - values=value, - rewards=reward, - returns=jnp.zeros_like(reward), - advantages=jnp.zeros_like(reward), - ) - return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage - - def rollout( - agent_state, episode_stats, next_obs, next_done, key, env_state, step_once_fn, max_steps - ): - (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( - step_once_fn, - (agent_state, episode_stats, next_obs, next_done, key, env_state), - (), - max_steps, - ) - return agent_state, episode_stats, next_obs, next_done, storage, key, env_state - - rollout = partial( - rollout, - step_once_fn=partial(step_once, env_step_fn=step_env_wrapped), - max_steps=args.num_steps, - ) - - print("Starting training...") - iters_bar = tqdm.tqdm( - range(1, args.num_iterations + 1), - disable=not sys.stdout.isatty(), - ) - returns = [] - is_tty = sys.stdout.isatty() - for iteration in iters_bar: - iteration_time_start = time.time() - - if not is_tty and iteration == 1: - print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True) - - agent_state, episode_stats, next_obs, next_done, storage, key, next_env_state = rollout( - agent_state, episode_stats, next_obs, next_done, key, next_env_state - ) - - if not is_tty and iteration == 1: - print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True) - - global_step += args.num_steps * args.num_envs - storage = compute_gae(agent_state, next_obs, next_done, storage) - - if not is_tty and iteration == 1: - print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True) - - agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo( - agent_state, storage, key - ) - - if not is_tty and iteration == 1: - print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True) - - losses.append(jnp.mean(loss)) - - avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns)) - iters_bar.set_postfix_str( - f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" - ) - - writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) - writer.add_scalar( - "charts/avg_episodic_length", - np.mean(jax.device_get(episode_stats.returned_episode_lengths)), - global_step, - ) - writer.add_scalar( - "charts/learning_rate", - agent_state.opt_state[1].hyperparams["learning_rate"].item(), - global_step, - ) - writer.add_scalar("losses/value_loss", v_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/policy_loss", pg_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/entropy", entropy_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/approx_kl", approx_kl[-1, -1].item(), global_step) - writer.add_scalar("losses/loss", loss[-1, -1].item(), global_step) - - # iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}") - - writer.add_scalar("charts/SPS", int(global_step / (time.time() - start_time)), global_step) - writer.add_scalar( - "charts/SPS_update", - int(args.num_envs * args.num_steps / (time.time() - iteration_time_start)), - global_step, - ) - - if not is_tty: - sps = int(global_step / (time.time() - start_time)) - remaining_steps = 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}/{args.num_iterations} | " - f"Step {global_step}/{args.total_timesteps} | " - f"SPS {sps} | " - f"Return {avg_episodic_return:.4f} | " - f"ETA {eta_str}", - flush=True, - ) - - if args.save_model: - model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model" - with open(model_path, "wb") as f: - f.write( - flax.serialization.to_bytes( - [ - vars(args), - [ - agent_state.params["sensor_params"], - agent_state.params["actor_params"], - agent_state.params["critic_params"], - agent_state.params["feature_extractor_params"], - ], - ] - ) - ) - print(f"model saved to {model_path}") - - env.close() - writer.close() - - print("Saving loss plot...") - plt.plot(losses) - plt.title("PPO Loss, mean over minibatches") - plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png") - plt.close() - - -def main() -> None: - 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) - - # Re-parse CLI to ensure they OVERRIDE the yaml - args = tyro.cli(PPOArgs, default=temp_args) - else: - args = temp_args - - train(args) - - -if __name__ == "__main__": - main() From ab7fa1584f489d67486f0701530a0ed17d2ab75b Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 13:58:26 +0200 Subject: [PATCH 13/19] other: actually add the renamed file to git --- experiments/_train_backup.py.back | 435 ++++++++++++++++++++++++++++++ 1 file changed, 435 insertions(+) create mode 100644 experiments/_train_backup.py.back diff --git a/experiments/_train_backup.py.back b/experiments/_train_backup.py.back new file mode 100644 index 0000000..a46e500 --- /dev/null +++ b/experiments/_train_backup.py.back @@ -0,0 +1,435 @@ +import datetime +import random +import yaml +import subprocess +import sys +import time +from dataclasses import asdict +from functools import partial +from typing import Callable + +import flax +import jax +import jax.numpy as jnp +import matplotlib.pyplot as plt +import numpy as np +import optax +import torch +import tqdm +import tyro +from flax.training.train_state import TrainState +from torch.utils.tensorboard import SummaryWriter + +from brittle_star_project.dataclasses import PPOArgs +from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics +from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper +from MLPs.mlps import ( + GenericDenseLayersWithActivation, + Actor, + OneDenseLayerMLP, + AgentParams, + Storage, +) +from plots import simple_plot +from ppo import PPO + + +def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: + return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( + obs_dict + ) + + +def make_env(env_config_path: str | None, num_envs: int) -> Callable: + def thunk(): + if env_config_path is None: + return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) + return BrittleStarJaxEnvWrapper.from_config(env_config_path, num_envs=num_envs) + + return thunk + + +def save_model(model_path: str, agent_state: TrainState, args: PPOArgs): + with open(model_path, "wb") as f: + f.write( + flax.serialization.to_bytes( + [ + vars(args), + [ + agent_state.params["network_params"], + agent_state.params["actor_params"], + agent_state.params["critic_params"], + ], + ] + ) + ) + + +def train(args: PPOArgs): + args.batch_size = args.num_envs * args.num_steps + args.minibatch_size = args.batch_size // args.num_minibatches + args.num_iterations = args.total_timesteps // args.batch_size + + # Try to get git short hash + try: + git_hash = ( + subprocess.check_output(["git", "rev-parse", "--short", "HEAD"]).decode("ascii").strip() + ) + except Exception: + git_hash = "none" + + 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 = 5 + print(f"running name: {run_name}") + + if args.run_dir is None: + args.run_dir = f"runs/{run_name}" + + import os + + os.makedirs(args.run_dir, exist_ok=True) + + if args.track: + import wandb + + wandb.init( + project=args.wandb_project_name, + entity=args.wandb_entity, + sync_tensorboard=True, + config=vars(args), + name=run_name, + save_code=True, + ) + + writer = SummaryWriter(args.run_dir) + writer.add_text( + "hyperparameters", + "|param|value|\n|---|---|\n" + "\n".join(f"|{k}|{v}|" for k, v in vars(args).items()), + ) + + random.seed(args.seed) + np.random.seed(args.seed) + key = jax.random.PRNGKey(args.seed) + key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5) + + torch.backends.cudnn.deterministic = args.torch_deterministic + device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") + print(f"Running on device: {device}") + + print("Creating the environment...") + env = make_env(env_config_path=args.env_config_path, num_envs=args.num_envs)() + print(f"Environment: {env}") + + episode_stats = EpisodeStatistics( + episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32), + episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), + returned_episode_returns=jnp.zeros(args.num_envs, jnp.float32), + returned_episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), + ) + + def step_env_wrapped(episode_stats: EpisodeStatistics, env_state, action): + next_env_state = env.step(env_state, action) + + # Extract per-environment signals from the state object + reward = next_env_state.reward # (num_envs,) + terminated = next_env_state.terminated # (num_envs,) + truncated = next_env_state.truncated # (num_envs,) + done = terminated | truncated # (num_envs,) + + new_episode_return = episode_stats.episode_returns + reward + new_episode_length = episode_stats.episode_lengths + 1 + + episode_stats = episode_stats.replace( + episode_returns=new_episode_return * (1 - done), + episode_lengths=new_episode_length * (1 - done), + returned_episode_returns=jnp.where( + done, new_episode_return, episode_stats.returned_episode_returns + ), + returned_episode_lengths=jnp.where( + done, new_episode_length, episode_stats.returned_episode_lengths + ), + ) + return ( + episode_stats, + next_env_state, + (convert_obs_dict_to_array(next_env_state.observations), reward, done), + ) + + def linear_schedule(count): + frac = 1.0 - (count // (args.num_minibatches * args.update_epochs)) / args.num_iterations + return args.learning_rate * frac + + print("Initializing the models...") + sensor = GenericDenseLayersWithActivation() + feature_extractor = GenericDenseLayersWithActivation() + actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX + critic = OneDenseLayerMLP() + # messager = OneDenseLayerMLP() + + sample_obs = jnp.concatenate( + [ + v.flatten() + for v in env.single_observation_space.sample(rng=jax.random.PRNGKey(0)).values() + if v.size > 0 + ] + ) + sensor_params = sensor.init(sensor_key, sample_obs) + feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs) + actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs)) + critic_params = critic.init( + critic_key, feature_extractor.apply(feature_extractor_params, sample_obs) + ) + + agent_state = TrainState.create( + apply_fn=None, + params=asdict( + AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) + ), + tx=optax.chain( + optax.clip_by_global_norm(args.max_grad_norm), + optax.inject_hyperparams(optax.adam)( + learning_rate=linear_schedule if args.anneal_lr else args.learning_rate, eps=1e-5 + ), + ), + ) + + sensor.apply = jax.jit(sensor.apply) + feature_extractor.apply = jax.jit(feature_extractor.apply) + actor.apply = jax.jit(actor.apply) + critic.apply = jax.jit(critic.apply) + ppo_instance = PPO(args, sensor, actor, critic, feature_extractor) + + @jax.jit + def get_action_and_value_noise( + agent_state: TrainState, + next_obs: jnp.ndarray, + key: jax.random.PRNGKey, + ): + hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) + hidden_critic = feature_extractor.apply( + agent_state.params["feature_extractor_params"], next_obs + ) + + # Continuous actions: sample from a Gaussian parameterized by the actor + mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) + key, subkey = jax.random.split(key) + noise = jax.random.normal(subkey, shape=mean.shape) + std = jnp.exp(log_std) + action = mean + noise * std + logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) + value = critic.apply(agent_state.params["critic_params"], hidden_critic) + return action, logprob, value.squeeze(-1), key + + # GAE + @jax.jit + def compute_gae_once(carry, inp, gamma, gae_lambda): + advantages = carry + nextdone, nextvalues, curvalues, reward = inp + nextnonterminal = 1.0 - nextdone + delta = reward + gamma * nextvalues * nextnonterminal - curvalues + advantages = delta + gamma * gae_lambda * nextnonterminal * advantages + return advantages, advantages + + @jax.jit + def compute_gae(agent_state, next_obs, next_done, storage): + next_value = critic.apply( + agent_state.params["critic_params"], + sensor.apply(agent_state.params["sensor_params"], next_obs), + ).squeeze(-1) + + advantages = jnp.zeros((args.num_envs,)) + dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0) + values = jnp.concatenate([storage.values, next_value[None, :]], axis=0) + _, advantages = jax.lax.scan( + partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda), + advantages, + (dones[1:], values[1:], values[:-1], storage.rewards), + reverse=True, + ) + return storage.replace(advantages=advantages, returns=advantages + storage.values) + + # END GAE + + # --- Main training loop --- + global_step = 0 + start_time = time.time() + + # Reset once to get initial state + print("Resetting the environment...") + if not sys.stdout.isatty(): + print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True) + next_env_state = env.reset(seed=args.seed) + next_obs = convert_obs_dict_to_array(next_env_state.observations) + next_done = jnp.zeros(args.num_envs, dtype=jnp.bool_) + if not sys.stdout.isatty(): + print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True) + + def step_once(carry, _, env_step_fn): + agent_state, episode_stats, obs, done, key, env_state = carry + action, logprob, value, key = get_action_and_value_noise(agent_state, obs, key) + + episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( + episode_stats, env_state, action + ) + + storage = Storage( + obs=obs, + actions=action, + logprobs=logprob, + dones=done, + values=value, + rewards=reward, + returns=jnp.zeros_like(reward), + advantages=jnp.zeros_like(reward), + ) + return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage + + def rollout( + agent_state, episode_stats, next_obs, next_done, key, env_state, step_once_fn, max_steps + ): + (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( + step_once_fn, + (agent_state, episode_stats, next_obs, next_done, key, env_state), + (), + max_steps, + ) + return agent_state, episode_stats, next_obs, next_done, storage, key, env_state + + rollout = partial( + rollout, + step_once_fn=partial(step_once, env_step_fn=step_env_wrapped), + max_steps=args.num_steps, + ) + + print("Starting training...") + iters_bar = tqdm.tqdm( + range(1, args.num_iterations + 1), + disable=not sys.stdout.isatty(), + ) + returns = [] + is_tty = sys.stdout.isatty() + for iteration in iters_bar: + iteration_time_start = time.time() + + if not is_tty and iteration == 1: + print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True) + + agent_state, episode_stats, next_obs, next_done, storage, key, next_env_state = rollout( + agent_state, episode_stats, next_obs, next_done, key, next_env_state + ) + + if not is_tty and iteration == 1: + print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True) + + global_step += args.num_steps * args.num_envs + storage = compute_gae(agent_state, next_obs, next_done, storage) + + if not is_tty and iteration == 1: + print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True) + + agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo( + agent_state, storage, key + ) + + if not is_tty and iteration == 1: + print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True) + + losses.append(jnp.mean(loss)) + + avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns)) + iters_bar.set_postfix_str( + f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" + ) + + writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) + writer.add_scalar( + "charts/avg_episodic_length", + np.mean(jax.device_get(episode_stats.returned_episode_lengths)), + global_step, + ) + writer.add_scalar( + "charts/learning_rate", + agent_state.opt_state[1].hyperparams["learning_rate"].item(), + global_step, + ) + writer.add_scalar("losses/value_loss", v_loss[-1, -1].item(), global_step) + writer.add_scalar("losses/policy_loss", pg_loss[-1, -1].item(), global_step) + writer.add_scalar("losses/entropy", entropy_loss[-1, -1].item(), global_step) + writer.add_scalar("losses/approx_kl", approx_kl[-1, -1].item(), global_step) + writer.add_scalar("losses/loss", loss[-1, -1].item(), global_step) + + # iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}") + + writer.add_scalar("charts/SPS", int(global_step / (time.time() - start_time)), global_step) + writer.add_scalar( + "charts/SPS_update", + int(args.num_envs * args.num_steps / (time.time() - iteration_time_start)), + global_step, + ) + + if not is_tty: + sps = int(global_step / (time.time() - start_time)) + remaining_steps = 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}/{args.num_iterations} | " + f"Step {global_step}/{args.total_timesteps} | " + f"SPS {sps} | " + f"Return {avg_episodic_return:.4f} | " + f"ETA {eta_str}", + flush=True, + ) + + if args.save_model: + model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model" + with open(model_path, "wb") as f: + f.write( + flax.serialization.to_bytes( + [ + vars(args), + [ + agent_state.params["sensor_params"], + agent_state.params["actor_params"], + agent_state.params["critic_params"], + agent_state.params["feature_extractor_params"], + ], + ] + ) + ) + print(f"model saved to {model_path}") + + env.close() + writer.close() + + print("Saving loss plot...") + plt.plot(losses) + plt.title("PPO Loss, mean over minibatches") + plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png") + plt.close() + + +def main() -> None: + 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) + + # Re-parse CLI to ensure they OVERRIDE the yaml + args = tyro.cli(PPOArgs, default=temp_args) + else: + args = temp_args + + train(args) + + +if __name__ == "__main__": + main() From 8a2f8b14bcc4d5c4fdfbaeecdc56732ee3ce6a46 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 14:07:46 +0200 Subject: [PATCH 14/19] fix(PPOTrainer.py, PPOArgs.py): added hyperparam config path to cli args, fixed missing arguments error in _step call --- experiments/PPOTrainer.py | 6 ++++-- experiments/train.py | 12 +++++++++--- src/brittle_star_project/dataclasses/PPOArgs.py | 3 +++ 3 files changed, 16 insertions(+), 5 deletions(-) diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index cfadb54..d947fb7 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -500,12 +500,14 @@ class PPOTrainer: iter_bar = tqdm.tqdm( range(1, self.args.num_iterations + 1), - disable=not sys.stdout.isatty(), + disable=not is_tty, ) for iteration in iter_bar: 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, is_tty=is_tty, iteration=iteration + ) if not is_tty and iteration == 1: print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True) diff --git a/experiments/train.py b/experiments/train.py index 596dad9..0a65493 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -17,11 +17,14 @@ def make_env(config_path: str | None, num_envs: int) -> BrittleStarJaxEnvWrapper return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs) -def parse_args() -> PPOArgs: +def parse_args(log: bool = True) -> 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: + if temp_args.hyperparameter_config_path is not None: + if log: + print(f"Loading hyperparameter config from {temp_args.hyperparameter_config_path}") + + with open(temp_args.hyperparameter_config_path, "r") as f: config = yaml.safe_load(f) if config: # parse PPOArgs with defaults from yaml. @@ -32,6 +35,9 @@ def parse_args() -> PPOArgs: # Reparse CLI to ensure they OVERRIDE the yaml args = tyro.cli(PPOArgs, default=temp_args) else: + if log: + print("No hyperparameter config provided, using default config") + args = temp_args return args diff --git a/src/brittle_star_project/dataclasses/PPOArgs.py b/src/brittle_star_project/dataclasses/PPOArgs.py index f6b20bf..3f04c16 100644 --- a/src/brittle_star_project/dataclasses/PPOArgs.py +++ b/src/brittle_star_project/dataclasses/PPOArgs.py @@ -13,6 +13,9 @@ class PPOArgs: # path to environment config file, if None, use default config env_config_path: str | None = None + # path to hyperparameter config file (yaml), if None, use default config + hyperparameter_config_path: str | None = None + # the name of this experiment exp_name: str = "brittle_star_ppo" From 7a0d00a071fbf51ade995c0b97be1081b096f6f6 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 14:10:17 +0200 Subject: [PATCH 15/19] fix(PPOTrainer.py): added log toggle to _step --- experiments/PPOTrainer.py | 19 +++++++++---------- 1 file changed, 9 insertions(+), 10 deletions(-) diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index d947fb7..365191c 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -378,8 +378,10 @@ class PPOTrainer: global_step, ) - def _step(self, env_state, next_obs, next_done, is_tty: bool, iteration: int) -> tuple: - if not is_tty and iteration == 1: + 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) ( @@ -392,23 +394,23 @@ class PPOTrainer: next_env_state, ) = self._rollout(env_state, next_obs, next_done) - if not is_tty and iteration == 1: + if log and 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) - if not is_tty and iteration == 1: + if log and 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._ppo.update_ppo(self.agent_state, storage, self.key) ) - if not is_tty and iteration == 1: + if log and not is_tty and iteration == 1: print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True) 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)).item() ) return ( @@ -509,13 +511,10 @@ class PPOTrainer: env_state, next_obs, next_done, is_tty=is_tty, iteration=iteration ) - 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 self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info) - if not is_tty: + if log and 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 From 5d2db286ff4a2f50bdc8b91c554e0ad1b80f7e5f Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 16:23:59 +0200 Subject: [PATCH 16/19] feat(ruff.toml): ignore unused imports (F401) in __init__.py --- ruff.toml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/ruff.toml b/ruff.toml index d665ebb..db1bcd3 100644 --- a/ruff.toml +++ b/ruff.toml @@ -394,3 +394,6 @@ extend-ignore = [ "PLW0603", # global-statement # "PLW1404", # implicit-str-concat ] + +[lint.per-file-ignores] +"__init__.py" = ["F401"] From b892e3777ef0e2a7b46f809f447185facdff94d0 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 16:30:38 +0200 Subject: [PATCH 17/19] feat: removed backup train code --- experiments/_train_backup.py.back | 435 ------------------------------ 1 file changed, 435 deletions(-) delete mode 100644 experiments/_train_backup.py.back diff --git a/experiments/_train_backup.py.back b/experiments/_train_backup.py.back deleted file mode 100644 index a46e500..0000000 --- a/experiments/_train_backup.py.back +++ /dev/null @@ -1,435 +0,0 @@ -import datetime -import random -import yaml -import subprocess -import sys -import time -from dataclasses import asdict -from functools import partial -from typing import Callable - -import flax -import jax -import jax.numpy as jnp -import matplotlib.pyplot as plt -import numpy as np -import optax -import torch -import tqdm -import tyro -from flax.training.train_state import TrainState -from torch.utils.tensorboard import SummaryWriter - -from brittle_star_project.dataclasses import PPOArgs -from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics -from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper -from MLPs.mlps import ( - GenericDenseLayersWithActivation, - Actor, - OneDenseLayerMLP, - AgentParams, - Storage, -) -from plots import simple_plot -from ppo import PPO - - -def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: - return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( - obs_dict - ) - - -def make_env(env_config_path: str | None, num_envs: int) -> Callable: - def thunk(): - if env_config_path is None: - return BrittleStarJaxEnvWrapper.default(num_envs=num_envs) - return BrittleStarJaxEnvWrapper.from_config(env_config_path, num_envs=num_envs) - - return thunk - - -def save_model(model_path: str, agent_state: TrainState, args: PPOArgs): - with open(model_path, "wb") as f: - f.write( - flax.serialization.to_bytes( - [ - vars(args), - [ - agent_state.params["network_params"], - agent_state.params["actor_params"], - agent_state.params["critic_params"], - ], - ] - ) - ) - - -def train(args: PPOArgs): - args.batch_size = args.num_envs * args.num_steps - args.minibatch_size = args.batch_size // args.num_minibatches - args.num_iterations = args.total_timesteps // args.batch_size - - # Try to get git short hash - try: - git_hash = ( - subprocess.check_output(["git", "rev-parse", "--short", "HEAD"]).decode("ascii").strip() - ) - except Exception: - git_hash = "none" - - 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 = 5 - print(f"running name: {run_name}") - - if args.run_dir is None: - args.run_dir = f"runs/{run_name}" - - import os - - os.makedirs(args.run_dir, exist_ok=True) - - if args.track: - import wandb - - wandb.init( - project=args.wandb_project_name, - entity=args.wandb_entity, - sync_tensorboard=True, - config=vars(args), - name=run_name, - save_code=True, - ) - - writer = SummaryWriter(args.run_dir) - writer.add_text( - "hyperparameters", - "|param|value|\n|---|---|\n" + "\n".join(f"|{k}|{v}|" for k, v in vars(args).items()), - ) - - random.seed(args.seed) - np.random.seed(args.seed) - key = jax.random.PRNGKey(args.seed) - key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5) - - torch.backends.cudnn.deterministic = args.torch_deterministic - device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu") - print(f"Running on device: {device}") - - print("Creating the environment...") - env = make_env(env_config_path=args.env_config_path, num_envs=args.num_envs)() - print(f"Environment: {env}") - - episode_stats = EpisodeStatistics( - episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32), - episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), - returned_episode_returns=jnp.zeros(args.num_envs, jnp.float32), - returned_episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32), - ) - - def step_env_wrapped(episode_stats: EpisodeStatistics, env_state, action): - next_env_state = env.step(env_state, action) - - # Extract per-environment signals from the state object - reward = next_env_state.reward # (num_envs,) - terminated = next_env_state.terminated # (num_envs,) - truncated = next_env_state.truncated # (num_envs,) - done = terminated | truncated # (num_envs,) - - new_episode_return = episode_stats.episode_returns + reward - new_episode_length = episode_stats.episode_lengths + 1 - - episode_stats = episode_stats.replace( - episode_returns=new_episode_return * (1 - done), - episode_lengths=new_episode_length * (1 - done), - returned_episode_returns=jnp.where( - done, new_episode_return, episode_stats.returned_episode_returns - ), - returned_episode_lengths=jnp.where( - done, new_episode_length, episode_stats.returned_episode_lengths - ), - ) - return ( - episode_stats, - next_env_state, - (convert_obs_dict_to_array(next_env_state.observations), reward, done), - ) - - def linear_schedule(count): - frac = 1.0 - (count // (args.num_minibatches * args.update_epochs)) / args.num_iterations - return args.learning_rate * frac - - print("Initializing the models...") - sensor = GenericDenseLayersWithActivation() - feature_extractor = GenericDenseLayersWithActivation() - actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX - critic = OneDenseLayerMLP() - # messager = OneDenseLayerMLP() - - sample_obs = jnp.concatenate( - [ - v.flatten() - for v in env.single_observation_space.sample(rng=jax.random.PRNGKey(0)).values() - if v.size > 0 - ] - ) - sensor_params = sensor.init(sensor_key, sample_obs) - feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs) - actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs)) - critic_params = critic.init( - critic_key, feature_extractor.apply(feature_extractor_params, sample_obs) - ) - - agent_state = TrainState.create( - apply_fn=None, - params=asdict( - AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params) - ), - tx=optax.chain( - optax.clip_by_global_norm(args.max_grad_norm), - optax.inject_hyperparams(optax.adam)( - learning_rate=linear_schedule if args.anneal_lr else args.learning_rate, eps=1e-5 - ), - ), - ) - - sensor.apply = jax.jit(sensor.apply) - feature_extractor.apply = jax.jit(feature_extractor.apply) - actor.apply = jax.jit(actor.apply) - critic.apply = jax.jit(critic.apply) - ppo_instance = PPO(args, sensor, actor, critic, feature_extractor) - - @jax.jit - def get_action_and_value_noise( - agent_state: TrainState, - next_obs: jnp.ndarray, - key: jax.random.PRNGKey, - ): - hidden = sensor.apply(agent_state.params["sensor_params"], next_obs) - hidden_critic = feature_extractor.apply( - agent_state.params["feature_extractor_params"], next_obs - ) - - # Continuous actions: sample from a Gaussian parameterized by the actor - mean, log_std = actor.apply(agent_state.params["actor_params"], hidden) - key, subkey = jax.random.split(key) - noise = jax.random.normal(subkey, shape=mean.shape) - std = jnp.exp(log_std) - action = mean + noise * std - logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) - value = critic.apply(agent_state.params["critic_params"], hidden_critic) - return action, logprob, value.squeeze(-1), key - - # GAE - @jax.jit - def compute_gae_once(carry, inp, gamma, gae_lambda): - advantages = carry - nextdone, nextvalues, curvalues, reward = inp - nextnonterminal = 1.0 - nextdone - delta = reward + gamma * nextvalues * nextnonterminal - curvalues - advantages = delta + gamma * gae_lambda * nextnonterminal * advantages - return advantages, advantages - - @jax.jit - def compute_gae(agent_state, next_obs, next_done, storage): - next_value = critic.apply( - agent_state.params["critic_params"], - sensor.apply(agent_state.params["sensor_params"], next_obs), - ).squeeze(-1) - - advantages = jnp.zeros((args.num_envs,)) - dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0) - values = jnp.concatenate([storage.values, next_value[None, :]], axis=0) - _, advantages = jax.lax.scan( - partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda), - advantages, - (dones[1:], values[1:], values[:-1], storage.rewards), - reverse=True, - ) - return storage.replace(advantages=advantages, returns=advantages + storage.values) - - # END GAE - - # --- Main training loop --- - global_step = 0 - start_time = time.time() - - # Reset once to get initial state - print("Resetting the environment...") - if not sys.stdout.isatty(): - print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True) - next_env_state = env.reset(seed=args.seed) - next_obs = convert_obs_dict_to_array(next_env_state.observations) - next_done = jnp.zeros(args.num_envs, dtype=jnp.bool_) - if not sys.stdout.isatty(): - print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True) - - def step_once(carry, _, env_step_fn): - agent_state, episode_stats, obs, done, key, env_state = carry - action, logprob, value, key = get_action_and_value_noise(agent_state, obs, key) - - episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( - episode_stats, env_state, action - ) - - storage = Storage( - obs=obs, - actions=action, - logprobs=logprob, - dones=done, - values=value, - rewards=reward, - returns=jnp.zeros_like(reward), - advantages=jnp.zeros_like(reward), - ) - return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage - - def rollout( - agent_state, episode_stats, next_obs, next_done, key, env_state, step_once_fn, max_steps - ): - (agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan( - step_once_fn, - (agent_state, episode_stats, next_obs, next_done, key, env_state), - (), - max_steps, - ) - return agent_state, episode_stats, next_obs, next_done, storage, key, env_state - - rollout = partial( - rollout, - step_once_fn=partial(step_once, env_step_fn=step_env_wrapped), - max_steps=args.num_steps, - ) - - print("Starting training...") - iters_bar = tqdm.tqdm( - range(1, args.num_iterations + 1), - disable=not sys.stdout.isatty(), - ) - returns = [] - is_tty = sys.stdout.isatty() - for iteration in iters_bar: - iteration_time_start = time.time() - - if not is_tty and iteration == 1: - print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True) - - agent_state, episode_stats, next_obs, next_done, storage, key, next_env_state = rollout( - agent_state, episode_stats, next_obs, next_done, key, next_env_state - ) - - if not is_tty and iteration == 1: - print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True) - - global_step += args.num_steps * args.num_envs - storage = compute_gae(agent_state, next_obs, next_done, storage) - - if not is_tty and iteration == 1: - print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True) - - agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo( - agent_state, storage, key - ) - - if not is_tty and iteration == 1: - print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True) - - losses.append(jnp.mean(loss)) - - avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns)) - iters_bar.set_postfix_str( - f"global_step={global_step}, avg_episodic_return={avg_episodic_return}" - ) - - writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step) - writer.add_scalar( - "charts/avg_episodic_length", - np.mean(jax.device_get(episode_stats.returned_episode_lengths)), - global_step, - ) - writer.add_scalar( - "charts/learning_rate", - agent_state.opt_state[1].hyperparams["learning_rate"].item(), - global_step, - ) - writer.add_scalar("losses/value_loss", v_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/policy_loss", pg_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/entropy", entropy_loss[-1, -1].item(), global_step) - writer.add_scalar("losses/approx_kl", approx_kl[-1, -1].item(), global_step) - writer.add_scalar("losses/loss", loss[-1, -1].item(), global_step) - - # iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}") - - writer.add_scalar("charts/SPS", int(global_step / (time.time() - start_time)), global_step) - writer.add_scalar( - "charts/SPS_update", - int(args.num_envs * args.num_steps / (time.time() - iteration_time_start)), - global_step, - ) - - if not is_tty: - sps = int(global_step / (time.time() - start_time)) - remaining_steps = 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}/{args.num_iterations} | " - f"Step {global_step}/{args.total_timesteps} | " - f"SPS {sps} | " - f"Return {avg_episodic_return:.4f} | " - f"ETA {eta_str}", - flush=True, - ) - - if args.save_model: - model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model" - with open(model_path, "wb") as f: - f.write( - flax.serialization.to_bytes( - [ - vars(args), - [ - agent_state.params["sensor_params"], - agent_state.params["actor_params"], - agent_state.params["critic_params"], - agent_state.params["feature_extractor_params"], - ], - ] - ) - ) - print(f"model saved to {model_path}") - - env.close() - writer.close() - - print("Saving loss plot...") - plt.plot(losses) - plt.title("PPO Loss, mean over minibatches") - plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png") - plt.close() - - -def main() -> None: - 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) - - # Re-parse CLI to ensure they OVERRIDE the yaml - args = tyro.cli(PPOArgs, default=temp_args) - else: - args = temp_args - - train(args) - - -if __name__ == "__main__": - main() From 431ecdf4b482df9f5d712550bdb6caff68170f86 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 17:02:40 +0200 Subject: [PATCH 18/19] restructure --- .github/workflows/update_hpc_requirements.yml | 2 +- docs/HPC.md | 2 +- docs/api/train.md | 1 - experiments/plots/__init__.py | 3 --- experiments/plots/plot.py | 12 ------------ .../export_requirements.py} | 3 --- {experiments => scripts}/simulate.py | 0 {experiments => scripts}/train.py | 2 +- .../brittle_star_project/trainers}/PPOTrainer.py | 10 +++++----- src/brittle_star_project/trainers/__init__.py | 0 10 files changed, 8 insertions(+), 27 deletions(-) delete mode 100644 docs/api/train.md delete mode 100644 experiments/plots/__init__.py delete mode 100644 experiments/plots/plot.py rename scripts/{export_hpc_requirements.py => hpc/export_requirements.py} (98%) rename {experiments => scripts}/simulate.py (100%) rename {experiments => scripts}/train.py (97%) rename {experiments => src/brittle_star_project/trainers}/PPOTrainer.py (98%) create mode 100644 src/brittle_star_project/trainers/__init__.py diff --git a/.github/workflows/update_hpc_requirements.yml b/.github/workflows/update_hpc_requirements.yml index 7e09b14..38dbc42 100644 --- a/.github/workflows/update_hpc_requirements.yml +++ b/.github/workflows/update_hpc_requirements.yml @@ -28,7 +28,7 @@ jobs: uses: astral-sh/setup-uv@v5 - name: Regenerate env/hpc/requirements.txt - run: uv run scripts/export_hpc_requirements.py + run: uv run scripts/hpc/export_requirements.py - name: Commit updated requirements if changed uses: stefanzweifel/git-auto-commit-action@v5 diff --git a/docs/HPC.md b/docs/HPC.md index ff18a5f..3e0e8d1 100644 --- a/docs/HPC.md +++ b/docs/HPC.md @@ -79,7 +79,7 @@ After installation, run these commands to ensure your environment is set up corr `env/hpc/requirements.txt` is auto-generated from `pyproject.toml`. To regenerate: ```bash -uv run scripts/export_hpc_requirements.py +uv run scripts/hpc/export_requirements.py ``` Modules listed in `env/hpc/modules.txt` are automatically excluded from the pip requirements to save space and use HPC-optimized binaries. diff --git a/docs/api/train.md b/docs/api/train.md deleted file mode 100644 index f87f5c1..0000000 --- a/docs/api/train.md +++ /dev/null @@ -1 +0,0 @@ -# TODO \ No newline at end of file diff --git a/experiments/plots/__init__.py b/experiments/plots/__init__.py deleted file mode 100644 index 92f34cb..0000000 --- a/experiments/plots/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .plot import simple_plot - -__all__ = ["simple_plot"] diff --git a/experiments/plots/plot.py b/experiments/plots/plot.py deleted file mode 100644 index ab8f4ad..0000000 --- a/experiments/plots/plot.py +++ /dev/null @@ -1,12 +0,0 @@ -import matplotlib.pyplot as plt - - -def simple_plot(x: list, y: list, show_window: bool = False, filename: str = "plot.png") -> None: - plt.plot(x, y) - plt.savefig(filename) - - if show_window: - # blocks until window is closed - plt.show() - - plt.close() diff --git a/scripts/export_hpc_requirements.py b/scripts/hpc/export_requirements.py similarity index 98% rename from scripts/export_hpc_requirements.py rename to scripts/hpc/export_requirements.py index 9ff3de7..5908ad3 100644 --- a/scripts/export_hpc_requirements.py +++ b/scripts/hpc/export_requirements.py @@ -5,9 +5,6 @@ This is a LOCAL DEVELOPER UTILITY — run it on your own machine before pushing code whenever pyproject.toml dependencies change. It reads the modules from env/hpc/modules.txt and the full dependency list from pyproject.toml, then writes the remainder to env/hpc/requirements.txt. - -Usage: - uv run scripts/export_hpc_requirements.py """ from __future__ import annotations diff --git a/experiments/simulate.py b/scripts/simulate.py similarity index 100% rename from experiments/simulate.py rename to scripts/simulate.py diff --git a/experiments/train.py b/scripts/train.py similarity index 97% rename from experiments/train.py rename to scripts/train.py index 0a65493..9ee6da6 100644 --- a/experiments/train.py +++ b/scripts/train.py @@ -7,7 +7,7 @@ import yaml import os from brittle_star_project.dataclasses import PPOArgs -from PPOTrainer import PPOTrainer +from brittle_star_project.trainers.PPOTrainer import PPOTrainer from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper diff --git a/experiments/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py similarity index 98% rename from experiments/PPOTrainer.py rename to src/brittle_star_project/trainers/PPOTrainer.py index 365191c..7df57f0 100644 --- a/experiments/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -28,13 +28,13 @@ from ppo import PPO @jax.jit -def linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): +def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations return learning_rate * frac @jax.jit -def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: +def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))( obs_dict ) @@ -124,7 +124,7 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): return ( episode_stats, next_env_state, - (convert_obs_dict_to_array(next_env_state.observations), reward, done), + (_convert_obs_dict_to_array(next_env_state.observations), reward, done), ) @@ -300,7 +300,7 @@ class PPOTrainer: optax.clip_by_global_norm(self.args.max_grad_norm), optax.inject_hyperparams(optax.adam)( learning_rate=partial( - linear_schedule, + _linear_schedule, minibatch_count=self.args.num_minibatches, update_epochs=self.args.update_epochs, num_iterations=self.args.num_iterations, @@ -467,7 +467,7 @@ class PPOTrainer: print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True) 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_) if log and not is_tty: diff --git a/src/brittle_star_project/trainers/__init__.py b/src/brittle_star_project/trainers/__init__.py new file mode 100644 index 0000000..e69de29 From ea851ebdbc7eaf8d5ccd4e09c71413cb7ba17071 Mon Sep 17 00:00:00 2001 From: RobinMeersman <77965843+RobinMeersman@users.noreply.github.com> Date: Mon, 6 Apr 2026 17:08:48 +0200 Subject: [PATCH 19/19] Apply suggestion from @tdpeuter Co-authored-by: Tibo De Peuter --- configs/hpc/smoke_test.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/configs/hpc/smoke_test.yaml b/configs/hpc/smoke_test.yaml index c892e3e..1dd2fc0 100644 --- a/configs/hpc/smoke_test.yaml +++ b/configs/hpc/smoke_test.yaml @@ -1,5 +1,5 @@ # Minimal config to verify HPC setup is functional. -# Run with: python experiments/train.py --config-path configs/hpc/smoke_test.yaml +# Run with: python scripts/train.py --config-path configs/hpc/smoke_test.yaml exp_name: "hpc_smoke_test" seed: 0 track: false # Test WandB integration