From 24da8948c50fe34cd25999bd573ad66b9403d2ea Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Fri, 3 Apr 2026 13:01:05 +0200 Subject: [PATCH] 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: