1
Fork 0

chore(main.py): deleted redundant main.py, feat(plot.py): added plot file to group all visualization related code for the experiments

This commit is contained in:
Robin Meersman 2026-04-02 15:20:47 +02:00
parent 97eaa5c06b
commit ba59361f97
3 changed files with 15 additions and 10 deletions

11
experiments/plots/plot.py Normal file
View file

@ -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()

View file

@ -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:

View file

@ -1,4 +0,0 @@
import jax
if __name__ == "__main__":
print(jax.devices())