fix: formatting + linting
This commit is contained in:
parent
3c7670e1b4
commit
9c5203fc7d
3 changed files with 6 additions and 4 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ---
|
||||
|
|
|
|||
|
|
@ -1,3 +1,3 @@
|
|||
from .plot import simple_plot
|
||||
|
||||
__all__ = ["simple_plot"]
|
||||
__all__ = ["simple_plot"]
|
||||
|
|
|
|||
Reference in a new issue