diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index 6e8b6f7..aea69d2 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -9,7 +9,6 @@ import jax import jax.numpy as jnp import numpy as np import optax -import torch import tqdm from flax.training.train_state import TrainState from torch.utils.tensorboard import SummaryWriter @@ -381,7 +380,9 @@ class PPOTrainer: self._ppo.update_ppo(self.agent_state, storage, self.key) ) - avg_episodic_return = float(jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))) + avg_episodic_return = float( + jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)) + ) return ( next_env_state, diff --git a/experiments/_train_backup.py b/experiments/_train_backup.py index 04a3df4..6d6a181 100644 --- a/experiments/_train_backup.py +++ b/experiments/_train_backup.py @@ -43,6 +43,7 @@ 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( @@ -198,7 +199,6 @@ 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): @@ -226,6 +226,7 @@ def train(args: PPOArgs): reverse=True, ) return storage.replace(advantages=advantages, returns=advantages + storage.values) + # END GAE # --- Main training loop --- diff --git a/experiments/plots/__init__.py b/experiments/plots/__init__.py index 5abf51e..92f34cb 100644 --- a/experiments/plots/__init__.py +++ b/experiments/plots/__init__.py @@ -1,3 +1,3 @@ from .plot import simple_plot -__all__ = ["simple_plot"] \ No newline at end of file +__all__ = ["simple_plot"]