1
Fork 0

fix(merge): fixed merge conflict after popping local changes back in from stash

This commit is contained in:
Robin Meersman 2026-04-03 13:01:05 +02:00
parent 1cb0e2d95b
commit 24da8948c5
3 changed files with 31 additions and 16 deletions

View file

@ -0,0 +1,3 @@
from .plot import simple_plot
__all__ = ["simple_plot"]

View file

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

View file

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