1
Fork 0

fix: formatting + linting

This commit is contained in:
Robin Meersman 2026-04-06 11:55:05 +02:00
parent 3c7670e1b4
commit 9c5203fc7d
3 changed files with 6 additions and 4 deletions

View file

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

View file

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

View file

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