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 ), )