From b839024e6eb0c8ca2cc84b1242b7e8a385ccd54e Mon Sep 17 00:00:00 2001 From: vsc46589 vscuser Date: Thu, 9 Apr 2026 23:14:52 +0200 Subject: [PATCH 1/7] feat(logging): added explained variance chart --- configs/hpc/wandb_expand.yaml | 10 ++++++++++ scripts/hpc/install.sh | 2 ++ scripts/hpc/train.pbs | 4 ++-- src/brittle_star_project/trainers/PPOTrainer.py | 12 ++++++++++++ 4 files changed, 26 insertions(+), 2 deletions(-) create mode 100644 configs/hpc/wandb_expand.yaml diff --git a/configs/hpc/wandb_expand.yaml b/configs/hpc/wandb_expand.yaml new file mode 100644 index 0000000..432160f --- /dev/null +++ b/configs/hpc/wandb_expand.yaml @@ -0,0 +1,10 @@ +exp_name: "explained_var_fun_more_steps" +seed: 42 +track: true +wandb_project_name: "LET-THERE-BE-MORE-LOGGING" +wandb_entity: "SEL3-2026-Groep-4" + +num_envs: 16 +num_steps: 256 +total_timesteps: 50000 +cuda: true \ No newline at end of file diff --git a/scripts/hpc/install.sh b/scripts/hpc/install.sh index f88d081..3d48443 100644 --- a/scripts/hpc/install.sh +++ b/scripts/hpc/install.sh @@ -46,6 +46,8 @@ source vsc-venv --activate \ set -euo pipefail cd "$PBS_O_WORKDIR" +uv run pip install . + echo '>>> Installing ipykernel...' CLUSTER_ID="${VSC_INSTITUTE_CLUSTER:-generic}" python -m ipykernel install --user --name="sel3_${CLUSTER_ID}" \ diff --git a/scripts/hpc/train.pbs b/scripts/hpc/train.pbs index e7d21f0..b9c8519 100644 --- a/scripts/hpc/train.pbs +++ b/scripts/hpc/train.pbs @@ -65,8 +65,8 @@ fi # TODO Once experiments get serious, change the config python scripts/train.py \ - --env-config-path configs/hpc/smoke_test.yaml \ - --hyperparameter-config-path configs/hpc/smoke_test.yaml \ + --env-config-path configs/hpc/wandb_expand.yaml \ + --hyperparameter-config-path configs/hpc/wandb_expand.yaml \ --run-dir "$SCRATCH_RUNDIR" echo ">>> Staging out results to $DATA_RUNDIR..." diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 8c238f1..f2f2d39 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -24,6 +24,10 @@ from brittle_star_project.MLPs.mlps import ( ) from brittle_star_project.ppo import PPO +def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> float: + var_returns = jnp.var(returns) + explained_var = 1.0 - jnp.var(returns - values) / (var_returns + 1e-8) + return float(explained_var) @jax.jit def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): @@ -197,6 +201,8 @@ class LossInfo: entropy_loss: Any approx_kl: Any avg_episodic_return: Any + # Example of better typing + explained_variance: float class PPOTrainer: @@ -352,6 +358,7 @@ class PPOTrainer: "charts/learning_rate": self.agent_state.opt_state[1] .hyperparams["learning_rate"] .item(), + "charts/explained_variance": loss_info.explained_variance, "losses/value_loss": loss_info.v_loss[-1, -1].item(), "losses/policy_loss": loss_info.pg_loss[-1, -1].item(), "losses/entropy": loss_info.entropy_loss[-1, -1].item(), @@ -397,6 +404,10 @@ class PPOTrainer: jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item() ) + explained_var = _compute_explained_variance( + storage.values, storage.returns + ) + return ( next_env_state, next_obs, @@ -408,6 +419,7 @@ class PPOTrainer: entropy_loss=entropy_loss, approx_kl=approx_kl, avg_episodic_return=avg_episodic_return, + explained_variance=explained_var ), ) From f7428f9d90888d723f85fb65eeb398d83ca1bed8 Mon Sep 17 00:00:00 2001 From: vsc46589 vscuser Date: Thu, 9 Apr 2026 23:40:58 +0200 Subject: [PATCH 2/7] feat(logging): added amount of envs that truncated and terminated --- src/brittle_star_project/trainers/PPOTrainer.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index f2f2d39..ecc2842 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -203,6 +203,8 @@ class LossInfo: avg_episodic_return: Any # Example of better typing explained_variance: float + num_terminated: int + num_truncated: int class PPOTrainer: @@ -359,6 +361,8 @@ class PPOTrainer: .hyperparams["learning_rate"] .item(), "charts/explained_variance": loss_info.explained_variance, + "charts/num_terminated": loss_info.num_terminated, + "charts/num_truncated": loss_info.num_truncated, "losses/value_loss": loss_info.v_loss[-1, -1].item(), "losses/policy_loss": loss_info.pg_loss[-1, -1].item(), "losses/entropy": loss_info.entropy_loss[-1, -1].item(), @@ -408,6 +412,12 @@ class PPOTrainer: storage.values, storage.returns ) + terminated = next_env_state.terminated # (num_envs,) + truncated = next_env_state.truncated # (num_envs,) + + num_terminated = int(jnp.sum(terminated).item()) + num_truncated = int(jnp.sum(truncated).item()) + return ( next_env_state, next_obs, @@ -419,7 +429,9 @@ class PPOTrainer: entropy_loss=entropy_loss, approx_kl=approx_kl, avg_episodic_return=avg_episodic_return, - explained_variance=explained_var + explained_variance=explained_var, + num_terminated=num_terminated, + num_truncated=num_truncated ), ) From 71349af5b99be4d83ca8f67013c3f5d9772d4d82 Mon Sep 17 00:00:00 2001 From: vsc46589 vscuser Date: Fri, 10 Apr 2026 00:11:18 +0200 Subject: [PATCH 3/7] feat(logging): added avg episode at which terminated or truncated happened --- .../trainers/PPOTrainer.py | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index ecc2842..33ffe47 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -205,6 +205,8 @@ class LossInfo: explained_variance: float num_terminated: int num_truncated: int + avg_terminated_length: Any + avg_truncated_length: Any class PPOTrainer: @@ -363,6 +365,8 @@ class PPOTrainer: "charts/explained_variance": loss_info.explained_variance, "charts/num_terminated": loss_info.num_terminated, "charts/num_truncated": loss_info.num_truncated, + "charts/avg_terminated_ep_length": loss_info.avg_terminated_length, + "charts/avg_truncated_ep_length": loss_info.avg_truncated_length, "losses/value_loss": loss_info.v_loss[-1, -1].item(), "losses/policy_loss": loss_info.pg_loss[-1, -1].item(), "losses/entropy": loss_info.entropy_loss[-1, -1].item(), @@ -414,10 +418,21 @@ class PPOTrainer: terminated = next_env_state.terminated # (num_envs,) truncated = next_env_state.truncated # (num_envs,) + episode_lengths = self.episode_stats.returned_episode_lengths num_terminated = int(jnp.sum(terminated).item()) num_truncated = int(jnp.sum(truncated).item()) + avg_terminated_length = ( + jnp.sum(episode_lengths * terminated) / + jnp.maximum(jnp.sum(terminated), 1) + ) + + avg_truncated_length = ( + jnp.sum(episode_lengths * truncated) / + jnp.maximum(jnp.sum(truncated), 1) + ) + return ( next_env_state, next_obs, @@ -431,7 +446,9 @@ class PPOTrainer: avg_episodic_return=avg_episodic_return, explained_variance=explained_var, num_terminated=num_terminated, - num_truncated=num_truncated + num_truncated=num_truncated, + avg_terminated_length=avg_terminated_length, + avg_truncated_length=avg_truncated_length ), ) From bf6434ea6ceebc9893e1f77623f37c4030d7514b Mon Sep 17 00:00:00 2001 From: cedric Date: Thu, 9 Apr 2026 22:13:07 +0000 Subject: [PATCH 4/7] fix(cleanup): ruff format --- .../trainers/PPOTrainer.py | 22 +++++++++---------- 1 file changed, 10 insertions(+), 12 deletions(-) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 33ffe47..298e41c 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -24,11 +24,13 @@ from brittle_star_project.MLPs.mlps import ( ) from brittle_star_project.ppo import PPO + def _compute_explained_variance(values: jnp.ndarray, returns: jnp.ndarray) -> float: var_returns = jnp.var(returns) explained_var = 1.0 - jnp.var(returns - values) / (var_returns + 1e-8) return float(explained_var) + @jax.jit def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate): frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations @@ -412,25 +414,21 @@ class PPOTrainer: jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item() ) - explained_var = _compute_explained_variance( - storage.values, storage.returns - ) + explained_var = _compute_explained_variance(storage.values, storage.returns) terminated = next_env_state.terminated # (num_envs,) truncated = next_env_state.truncated # (num_envs,) episode_lengths = self.episode_stats.returned_episode_lengths num_terminated = int(jnp.sum(terminated).item()) - num_truncated = int(jnp.sum(truncated).item()) - - avg_terminated_length = ( - jnp.sum(episode_lengths * terminated) / - jnp.maximum(jnp.sum(terminated), 1) + num_truncated = int(jnp.sum(truncated).item()) + + avg_terminated_length = jnp.sum(episode_lengths * terminated) / jnp.maximum( + jnp.sum(terminated), 1 ) - avg_truncated_length = ( - jnp.sum(episode_lengths * truncated) / - jnp.maximum(jnp.sum(truncated), 1) + avg_truncated_length = jnp.sum(episode_lengths * truncated) / jnp.maximum( + jnp.sum(truncated), 1 ) return ( @@ -448,7 +446,7 @@ class PPOTrainer: num_terminated=num_terminated, num_truncated=num_truncated, avg_terminated_length=avg_terminated_length, - avg_truncated_length=avg_truncated_length + avg_truncated_length=avg_truncated_length, ), ) From 6e942d896bfa38487065c6f4453c2ccec48803e2 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Fri, 10 Apr 2026 08:23:34 +0200 Subject: [PATCH 5/7] fix(gae): fixes #31 --- src/brittle_star_project/trainers/PPOTrainer.py | 14 +++++++++++--- 1 file changed, 11 insertions(+), 3 deletions(-) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index 298e41c..a9e9b1b 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -175,11 +175,19 @@ def _compute_gae_once(carry, inp, gamma, gae_lambda): # jit applied on partial-wrapped wrapper method self._compute_gae_jit def _compute_gae_jit( - agent_state, storage, next_obs, next_done, gamma, gae_lambda, num_envs, sensor, critic + agent_state, + storage, + next_obs, + next_done, + gamma, + gae_lambda, + num_envs, + feature_extractor, + critic, ): next_value = critic.apply( agent_state.params["critic_params"], - sensor.apply(agent_state.params["sensor_params"], next_obs), + feature_extractor.apply(agent_state.params["feature_extractor_params"], next_obs), ).squeeze(-1) advantages = jnp.zeros((num_envs,)) @@ -244,7 +252,7 @@ class PPOTrainer: num_envs=self.args.num_envs, gamma=self.args.gamma, gae_lambda=self.args.gae_lambda, - sensor=self.sensor, + feature_extractor=self.feature_extractor, critic=self.critic, ) ) From b7e2beb65d6ceb4d4d2149db291e110fcadcd994 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Fri, 10 Apr 2026 08:31:52 +0200 Subject: [PATCH 6/7] fx(cleanup): removed meaningless comments, added typing and renamed lossinfo to trainingmeasurements --- .../trainers/PPOTrainer.py | 81 +++++++++---------- 1 file changed, 37 insertions(+), 44 deletions(-) diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index a9e9b1b..6bc94ef 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -44,7 +44,6 @@ def _convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray: ) -# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit def _get_action_and_value_noise( sensor: GenericDenseLayersWithActivation, feature_extractor: GenericDenseLayersWithActivation, @@ -59,7 +58,6 @@ def _get_action_and_value_noise( 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) @@ -70,7 +68,6 @@ def _get_action_and_value_noise( return action, logprob, value.squeeze(-1), key -# removed jit: used in _rollout_jit, so will be compiled with _rollout_jit def _step_once( carry, _, @@ -102,15 +99,13 @@ def _step_once( return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage -# 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 - 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,) + reward = next_env_state.reward + terminated = next_env_state.terminated + truncated = next_env_state.truncated + done = terminated | truncated new_episode_return = episode_stats.episode_returns + reward new_episode_length = episode_stats.episode_lengths + 1 @@ -132,7 +127,6 @@ def _step_env_wrapped(episode_stats, env_state, action, env_step_fn): ) -# jit applied in wrapper method self._rollout_jit using partial def _rollout_jit( agent_state, episode_stats, @@ -163,7 +157,6 @@ def _rollout_jit( return agent_state, episode_stats, next_obs, next_done, storage, key, env_state -# 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 @@ -173,7 +166,6 @@ def _compute_gae_once(carry, inp, gamma, gae_lambda): return advantages, advantages -# jit applied on partial-wrapped wrapper method self._compute_gae_jit def _compute_gae_jit( agent_state, storage, @@ -203,15 +195,13 @@ def _compute_gae_jit( @dataclass -class LossInfo: - # todo: better typing - loss: Any - pg_loss: Any - v_loss: Any - entropy_loss: Any - approx_kl: Any - avg_episodic_return: Any - # Example of better typing +class TrainingMeasurements: + loss: jnp.ndarray + pg_loss: jnp.ndarray + v_loss: jnp.ndarray + entropy_loss: jnp.ndarray + approx_kl: jnp.ndarray + avg_episodic_return: float explained_variance: float num_terminated: int num_truncated: int @@ -276,11 +266,8 @@ class PPOTrainer: sensor = GenericDenseLayersWithActivation() feature_extractor = GenericDenseLayersWithActivation() - actor = Actor( - action_dim=self.env.single_action_space.shape[0] - ) # continuous actions for MJX + actor = Actor(action_dim=self.env.single_action_space.shape[0]) critic = OneDenseLayerMLP() - # messenger = OneDenseLayerMLP() return sensor, feature_extractor, actor, critic def _init_agent_state(self) -> TrainState: @@ -331,7 +318,7 @@ class PPOTrainer: def _init_episode_stats(self) -> EpisodeStatistics: self.logger.info("[EPISODE STATS]: Initializing episode stats...") - return EpisodeStatistics( # type: ignore[call-arg] + 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), @@ -362,26 +349,26 @@ class PPOTrainer: episode_stats, start_time, iteration_time_start, - loss_info, + training_measurements, ): metrics = { - "charts/avg_episodic_return": loss_info.avg_episodic_return, + "charts/avg_episodic_return": training_measurements.avg_episodic_return, "charts/avg_episodic_length": np.mean( jax.device_get(episode_stats.returned_episode_lengths) ), "charts/learning_rate": self.agent_state.opt_state[1] .hyperparams["learning_rate"] .item(), - "charts/explained_variance": loss_info.explained_variance, - "charts/num_terminated": loss_info.num_terminated, - "charts/num_truncated": loss_info.num_truncated, - "charts/avg_terminated_ep_length": loss_info.avg_terminated_length, - "charts/avg_truncated_ep_length": loss_info.avg_truncated_length, - "losses/value_loss": loss_info.v_loss[-1, -1].item(), - "losses/policy_loss": loss_info.pg_loss[-1, -1].item(), - "losses/entropy": loss_info.entropy_loss[-1, -1].item(), - "losses/approx_kl": loss_info.approx_kl[-1, -1].item(), - "losses/loss": loss_info.loss[-1, -1].item(), + "charts/explained_variance": training_measurements.explained_variance, + "charts/num_terminated": training_measurements.num_terminated, + "charts/num_truncated": training_measurements.num_truncated, + "charts/avg_terminated_ep_length": training_measurements.avg_terminated_length, + "charts/avg_truncated_ep_length": training_measurements.avg_truncated_length, + "losses/value_loss": training_measurements.v_loss[-1, -1].item(), + "losses/policy_loss": training_measurements.pg_loss[-1, -1].item(), + "losses/entropy": training_measurements.entropy_loss[-1, -1].item(), + "losses/approx_kl": training_measurements.approx_kl[-1, -1].item(), + "losses/loss": training_measurements.loss[-1, -1].item(), "charts/SPS": int(global_step / (time.time() - start_time)), "charts/SPS_update": int( self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start) @@ -424,8 +411,8 @@ class PPOTrainer: explained_var = _compute_explained_variance(storage.values, storage.returns) - terminated = next_env_state.terminated # (num_envs,) - truncated = next_env_state.truncated # (num_envs,) + terminated = next_env_state.terminated + truncated = next_env_state.truncated episode_lengths = self.episode_stats.returned_episode_lengths num_terminated = int(jnp.sum(terminated).item()) @@ -443,7 +430,7 @@ class PPOTrainer: next_env_state, next_obs, next_done, - LossInfo( + TrainingMeasurements( loss=loss, pg_loss=pg_loss, v_loss=v_loss, @@ -499,12 +486,18 @@ class PPOTrainer: 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, training_measurements = self._step( env_state, next_obs, next_done, iteration=iteration ) global_step += self.args.num_steps * self.args.num_envs - self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info) + self._log( + global_step, + self.episode_stats, + start_time, + iteration_time_start, + training_measurements, + ) sps = int(global_step / (time.time() - start_time)) remaining_steps = self.args.total_timesteps - global_step @@ -515,7 +508,7 @@ class PPOTrainer: 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"Return {training_measurements.avg_episodic_return:.4f} | " f"ETA {eta_str}" ) From cee6b53ada810db4a7a23c65552cbdda50dcce16 Mon Sep 17 00:00:00 2001 From: JibrilExe Date: Fri, 10 Apr 2026 09:07:43 +0200 Subject: [PATCH 7/7] fix(hpc): removed local package install command as it was fixed in tibo's older pr --- scripts/hpc/install.sh | 2 -- 1 file changed, 2 deletions(-) diff --git a/scripts/hpc/install.sh b/scripts/hpc/install.sh index 3d48443..f88d081 100644 --- a/scripts/hpc/install.sh +++ b/scripts/hpc/install.sh @@ -46,8 +46,6 @@ source vsc-venv --activate \ set -euo pipefail cd "$PBS_O_WORKDIR" -uv run pip install . - echo '>>> Installing ipykernel...' CLUSTER_ID="${VSC_INSTITUTE_CLUSTER:-generic}" python -m ipykernel install --user --name="sel3_${CLUSTER_ID}" \