fix(merge): fixed merge conflict after popping local changes back in from stash
This commit is contained in:
parent
1cb0e2d95b
commit
24da8948c5
3 changed files with 31 additions and 16 deletions
3
experiments/plots/__init__.py
Normal file
3
experiments/plots/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
from .plot import simple_plot
|
||||||
|
|
||||||
|
__all__ = ["simple_plot"]
|
||||||
|
|
@ -6,6 +6,7 @@ def simple_plot(x: list, y: list, show_window: bool = False, filename: str = "pl
|
||||||
plt.savefig(filename)
|
plt.savefig(filename)
|
||||||
|
|
||||||
if show_window:
|
if show_window:
|
||||||
|
# blocks until window is closed
|
||||||
plt.show()
|
plt.show()
|
||||||
|
|
||||||
plt.close()
|
plt.close()
|
||||||
|
|
|
||||||
|
|
@ -25,6 +25,7 @@ from MLPs.mlps import (
|
||||||
AgentParams,
|
AgentParams,
|
||||||
Storage,
|
Storage,
|
||||||
)
|
)
|
||||||
|
from plots import simple_plot
|
||||||
from ppo import PPO
|
from ppo import PPO
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -42,6 +43,21 @@ def make_env(config_path: str | None, num_envs: int) -> Callable:
|
||||||
|
|
||||||
return thunk
|
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):
|
def train(args: PPOArgs):
|
||||||
args.batch_size = args.num_envs * args.num_steps
|
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)
|
value = critic.apply(agent_state.params["critic_params"], hidden_critic)
|
||||||
return action, logprob, value.squeeze(-1), key
|
return action, logprob, value.squeeze(-1), key
|
||||||
|
|
||||||
|
|
||||||
|
# GAE
|
||||||
@jax.jit
|
@jax.jit
|
||||||
def compute_gae_once(carry, inp, gamma, gae_lambda):
|
def compute_gae_once(carry, inp, gamma, gae_lambda):
|
||||||
advantages = carry
|
advantages = carry
|
||||||
|
|
@ -208,6 +226,7 @@ def train(args: PPOArgs):
|
||||||
reverse=True,
|
reverse=True,
|
||||||
)
|
)
|
||||||
return storage.replace(advantages=advantages, returns=advantages + storage.values)
|
return storage.replace(advantages=advantages, returns=advantages + storage.values)
|
||||||
|
# END GAE
|
||||||
|
|
||||||
# --- Main training loop ---
|
# --- Main training loop ---
|
||||||
global_step = 0
|
global_step = 0
|
||||||
|
|
@ -277,7 +296,7 @@ def train(args: PPOArgs):
|
||||||
f"global_step={global_step}, avg_episodic_return={avg_episodic_return}"
|
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("charts/avg_episodic_return", avg_episodic_return, global_step)
|
||||||
writer.add_scalar(
|
writer.add_scalar(
|
||||||
|
|
@ -307,27 +326,19 @@ def train(args: PPOArgs):
|
||||||
|
|
||||||
if args.save_model:
|
if args.save_model:
|
||||||
model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model"
|
model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model"
|
||||||
with open(model_path, "wb") as f:
|
save_model(model_path, agent_state, args)
|
||||||
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"],
|
|
||||||
],
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
print(f"model saved to {model_path}")
|
print(f"model saved to {model_path}")
|
||||||
|
|
||||||
env.close()
|
env.close()
|
||||||
writer.close()
|
writer.close()
|
||||||
|
|
||||||
print("Saving loss plot...")
|
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:
|
def main() -> None:
|
||||||
|
|
|
||||||
Reference in a new issue