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 flax
import jax import jax
import jax.numpy as jnp import jax.numpy as jnp
import matplotlib.pyplot as plt
import numpy as np import numpy as np
import optax import optax
import torch 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.dataclasses.EpisodeStatistics import EpisodeStatistics
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from brittle_star_project.rl import Actor, AgentParams, Critic, Network, Storage from brittle_star_project.rl import Actor, AgentParams, Critic, Network, Storage
from experiments.plots.plot import simple_plot
from ppo import PPO from ppo import PPO
@ -41,7 +41,8 @@ def make_env(config_path: str | None, num_envs: int) -> Callable:
def train(args: PPOArgs): def train(args: PPOArgs):
args.batch_size = args.num_envs * args.num_steps args.batch_size = args.num_envs * args.num_steps
args.minibatch_size = args.batch_size // args.num_minibatches 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())}" run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
print(f"running name: {run_name}") print(f"running name: {run_name}")
@ -313,10 +314,7 @@ def train(args: PPOArgs):
writer.close() writer.close()
print("Saving loss plot...") print("Saving loss plot...")
plt.plot(returns) simple_plot(range(len(returns)), returns, show_window=True, filename=f"runs/{run_name}/{args.exp_name}_losses.png")
plt.title("PPO Episodic Returns, mean over minibatches")
plt.savefig(f"runs/{run_name}/{args.exp_name}_losses.png")
plt.close()
def main() -> None: def main() -> None:

View file

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