From c0e7569773c17c1ded9e70a7a2ec18b1cebad3e6 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Sat, 11 Apr 2026 13:42:52 +0200 Subject: [PATCH] feat: More low lvl logs, reward, advantage, returns --- configs/hpc/debug.yaml | 4 +- scripts/train.py | 5 +- src/brittle_star_project/MLPs/mlps.py | 4 ++ .../trainers/PPOTrainer.py | 57 +++++++++++++++---- 4 files changed, 54 insertions(+), 16 deletions(-) diff --git a/configs/hpc/debug.yaml b/configs/hpc/debug.yaml index 20c65e8..c387dc1 100644 --- a/configs/hpc/debug.yaml +++ b/configs/hpc/debug.yaml @@ -1,10 +1,10 @@ # Configuration for debug session -exp_name: "debug-experiment-10042026" # started on april 10 +exp_name: "debug-experiment" seed: 42 track: true wandb_project_name: "Let's-find-that-bug" wandb_entity: "SEL3-2026-Groep-4" - + num_envs: 32 num_steps: 32 total_timesteps: 102400 diff --git a/scripts/train.py b/scripts/train.py index 1f543f7..886bd54 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -50,7 +50,7 @@ if __name__ == "__main__": 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}" + run_dir = f"/data/gent/465/vsc46589/runs/{run_name}" else: run_dir = args.run_dir @@ -69,9 +69,6 @@ if __name__ == "__main__": env = make_env(args.env_config_path, args.num_envs) raw_env = env.raw - logger.log( - {"run_dir": run_dir} - ) torch.backends.cudnn.deterministic = args.torch_deterministic diff --git a/src/brittle_star_project/MLPs/mlps.py b/src/brittle_star_project/MLPs/mlps.py index 6abb540..9568a36 100644 --- a/src/brittle_star_project/MLPs/mlps.py +++ b/src/brittle_star_project/MLPs/mlps.py @@ -58,6 +58,10 @@ class Storage: returns: jnp.array rewards: jnp.array + raw_actions: jnp.ndarray = None # before clipping + means: jnp.ndarray = None # policy mean + stds: jnp.ndarray = None # policy std + def replace(self, **kwargs) -> "Storage": fs = fields(self) return Storage(**{f.name: kwargs.get(f.name, getattr(self, f.name)) for f in fs}) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 1e554f8..8a9b341 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -11,8 +11,6 @@ import numpy as np import optax from flax.training.train_state import TrainState -from experiment_logger import get_logger - from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.MLPs.mlps import ( @@ -23,6 +21,7 @@ from brittle_star_project.MLPs.mlps import ( Storage, ) from brittle_star_project.ppo import PPO +from experiment_logger import get_logger @jax.jit def _clip_action(action: jnp.ndarray, low: jnp.ndarray, high: jnp.ndarray) -> jnp.ndarray: @@ -67,15 +66,17 @@ def _get_action_and_value_noise( key, subkey = jax.random.split(key) noise = jax.random.normal(subkey, shape=mean.shape) std = jnp.exp(log_std) - action = mean + noise * std + raw_action = mean + noise * std clipped_action = _clip_action( - action, + raw_action, action_low, action_high ) - logprob = -0.5 * (((clipped_action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1) + logprob = -0.5 * (((raw_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 clipped_action, logprob, value.squeeze(-1), key + + return clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key + def _step_once( @@ -90,21 +91,24 @@ def _step_once( action_high ): agent_state, episode_stats, obs, done, key, env_state = carry - action, logprob, value, key = _get_action_and_value_noise( + clipped_action, raw_action, logprob, value, mean, std, key = _get_action_and_value_noise( sensor, feature_extractor, actor, critic, agent_state, obs, key, action_low, action_high ) episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn( - episode_stats, env_state, action + episode_stats, env_state, clipped_action ) storage = Storage( obs=obs, - actions=action, + actions=clipped_action, + raw_actions=raw_action, logprobs=logprob, dones=done, values=value, rewards=reward, + means=mean, + stds=std, returns=jnp.zeros_like(reward), advantages=jnp.zeros_like(reward), ) @@ -371,7 +375,37 @@ class PPOTrainer: start_time, iteration_time_start, training_measurements, + storage ): + data = jax.device_get({ + 'rewards': storage.rewards[0], # (num_steps,) + 'values': storage.values[0], + 'returns': storage.returns[0], # NEW + 'advantages': storage.advantages[0],# NEW + 'actions': storage.actions[0], + 'raw_actions': storage.raw_actions[0], + 'means': storage.means[0], + 'stds': storage.stds[0], + 'logprobs': storage.logprobs[0], + }) + + storage_metrics = { + "rollout/env0/return_mean": float(np.mean(data['returns'])), + + "rollout/env0/advantage_mean": float(np.mean(data['advantages'])), + + "rollout/env0/value_mean": float(np.mean(data['values'])), + "rollout/env0/value_vs_return_diff": float(np.mean(data['values'] - data['returns'])), + + "rollout/env0/reward_mean": float(np.mean(data['rewards'])), + + "rollout/env0/mean_mean": float(np.mean(data['means'])), + "rollout/env0/logprob_mean": float(np.mean(data['logprobs'])), + + "rollout/env0/action_mean": float(np.mean(data['actions'])), + "rollout/env0/raw_action_mean": float(np.mean(data['raw_actions'])), + } + metrics = { "charts/avg_episodic_return": training_measurements.avg_episodic_return, "charts/avg_episodic_length": np.mean( @@ -394,6 +428,7 @@ class PPOTrainer: "charts/SPS_update": int( self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start) ), + **storage_metrics } self.logger.log(metrics, step=global_step) @@ -464,6 +499,7 @@ class PPOTrainer: avg_terminated_length=avg_terminated_length, avg_truncated_length=avg_truncated_length, ), + storage ) def _close(self): @@ -507,7 +543,7 @@ class PPOTrainer: for iteration in iter_bar: iteration_time_start = time.time() - env_state, next_obs, next_done, training_measurements = self._step( + env_state, next_obs, next_done, training_measurements, storage = self._step( env_state, next_obs, next_done, iteration=iteration ) @@ -518,6 +554,7 @@ class PPOTrainer: start_time, iteration_time_start, training_measurements, + storage ) sps = int(global_step / (time.time() - start_time))