1
Fork 0

feat(PPOTrainer.py): added logging messages

This commit is contained in:
Robin Meersman 2026-04-06 13:56:54 +02:00
parent 84f6fbd76e
commit b97fc30d5d
5 changed files with 135 additions and 27 deletions

View file

@ -81,7 +81,6 @@ def train(args: PPOArgs):
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
# 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}")
if args.run_dir is None:
@ -385,8 +384,6 @@ def train(args: PPOArgs):
)
if args.save_model:
model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model"
save_model(model_path, agent_state, args)
model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model"
with open(model_path, "wb") as f:
f.write(
@ -408,12 +405,6 @@ def train(args: PPOArgs):
writer.close()
print("Saving loss plot...")
simple_plot(
list(range(len(returns))),
returns,
show_window=True,
filename=f"runs/{run_name}/{args.exp_name}_losses.png",
)
plt.plot(losses)
plt.title("PPO Loss, mean over minibatches")
plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png")