1
Fork 0

fix(PPOTrainer.py): cleaned up + bug fixes regarding misuse of variable/wrong returns

This commit is contained in:
Robin Meersman 2026-04-03 14:54:57 +02:00
parent 4396c9ac4b
commit a9302cfd4e
6 changed files with 133 additions and 88 deletions

14
docs/api/simulate.md Normal file
View file

@ -0,0 +1,14 @@
# Training and Simulation for Brittle Star Models
## Simulating a model
In order to simulate and view the behavior of a trained model, you can use the `simulate.py` script. This script allows you to specify the path to a trained model and will launch a simulation using that model. This script has the following parameters:
- `--model`: The path to the trained model artifact to simulate.
- `--model-type`: The type of model to simulate (e.g., `random`, ...)
- `--task`: The task to simulate (e.g., `directed_locomotion`, ...)
- `--seed`: The random seed for reproducibility.
```bash
python simulate.py --model artifacts/my_model --model-type random --task directed_locomotion --seed 0
```

1
docs/api/train.md Normal file
View file

@ -0,0 +1 @@
# TODO

View file

@ -1,30 +0,0 @@
# Training and Simulation for Brittle Star Models
## Training a model
To train a model, you can use the `train.py` script. This script allows to pass some parameters to customize the training process:
- `--out`: The output path where the trained model will be saved.
- `--model_type`: The type of model to train (e.g., `random`, ...)
- `--task`: The task to train on (e.g., `directed_locomotion`, ...)
- `--seed`: The random seed for reproducibility.
- `--epochs`: The number of epochs to train for.
This will then train the specified model on the specified task for the given number of epochs and save the trained model to the specified output path.
```bash
python train.py --out artifacts/my_model --model-type random --task directed_locomotion --seed 0 --epochs 50
```
## Simulating a model
In order to simulate and view the behavior of a trained model, you can use the `simulate.py` script. This script allows you to specify the path to a trained model and will launch a simulation using that model. This script has the following parameters:
- `--model`: The path to the trained model artifact to simulate.
- `--model-type`: The type of model to simulate (e.g., `random`, ...)
- `--task`: The task to simulate (e.g., `directed_locomotion`, ...)
- `--seed`: The random seed for reproducibility.
```bash
python simulate.py --model artifacts/my_model --model-type random --task directed_locomotion --seed 0
```

View file

@ -1,26 +1,28 @@
import random
import time
from dataclasses import asdict, dataclass
from functools import partial
from typing import Any
import optax
import tqdm
from flax.metrics.tensorboard import SummaryWriter
from flax.training.train_state import TrainState
from MLPs.mlps import (
GenericDenseLayersWithActivation,
Actor,
OneDenseLayerMLP,
AgentParams,
Storage,
)
from brittle_star_project.dataclasses import PPOArgs, EpisodeStatistics
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
import flax
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
from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from MLPs.mlps import (
Actor,
AgentParams,
GenericDenseLayersWithActivation,
OneDenseLayerMLP,
Storage,
)
from ppo import PPO
@ -95,6 +97,36 @@ def _step_once(
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
@jax.jit
def _step_env_wrapped(env_step_fn, env_state, action, episode_stats):
next_env_state = env_step_fn(env_state, action)
# Extract per-environment signals from the state object
reward = next_env_state.reward # (num_envs,)
terminated = next_env_state.terminated # (num_envs,)
truncated = next_env_state.truncated # (num_envs,)
done = terminated | truncated # (num_envs,)
new_episode_return = episode_stats.episode_returns + reward
new_episode_length = episode_stats.episode_lengths + 1
episode_stats = episode_stats.replace(
episode_returns=new_episode_return * (1 - done),
episode_lengths=new_episode_length * (1 - done),
returned_episode_returns=jnp.where(
done, new_episode_return, episode_stats.returned_episode_returns
),
returned_episode_lengths=jnp.where(
done, new_episode_length, episode_stats.returned_episode_lengths
),
)
return (
episode_stats,
next_env_state,
(convert_obs_dict_to_array(next_env_state.observations), reward, done),
)
@jax.jit
def _rollout_jit(
agent_state,
@ -126,36 +158,6 @@ def _rollout_jit(
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
@jax.jit
def _step_env_wrapped(env_step_fn, env_state, action, episode_stats):
next_env_state = env_step_fn(env_state, action)
# Extract per-environment signals from the state object
reward = next_env_state.reward # (num_envs,)
terminated = next_env_state.terminated # (num_envs,)
truncated = next_env_state.truncated # (num_envs,)
done = terminated | truncated # (num_envs,)
new_episode_return = episode_stats.episode_returns + reward
new_episode_length = episode_stats.episode_lengths + 1
episode_stats = episode_stats.replace(
episode_returns=new_episode_return * (1 - done),
episode_lengths=new_episode_length * (1 - done),
returned_episode_returns=jnp.where(
done, new_episode_return, episode_stats.returned_episode_returns
),
returned_episode_lengths=jnp.where(
done, new_episode_length, episode_stats.returned_episode_lengths
),
)
return (
episode_stats,
next_env_state,
(convert_obs_dict_to_array(next_env_state.observations), reward, done),
)
@jax.jit
def compute_gae_once(carry, inp, gamma, gae_lambda):
advantages = carry
@ -200,7 +202,8 @@ class PPOTrainer:
def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_name: str):
self.args = args
self.env = env
self.writer = SummaryWriter(f"runs/{run_name}")
self.run_name = run_name
self.writer = SummaryWriter(f"runs/{self.run_name}")
self.key = jax.random.PRNGKey(args.seed)
@ -216,6 +219,12 @@ class PPOTrainer:
self.episode_stats = self._init_episode_stats()
self._init_random()
def _init_random(self):
random.seed(self.args.seed)
np.random.seed(self.args.seed)
def _init_agent(self):
sensor = GenericDenseLayersWithActivation()
feature_extractor = GenericDenseLayersWithActivation()
@ -255,7 +264,13 @@ class PPOTrainer:
tx=optax.chain(
optax.clip_by_global_norm(self.args.max_grad_norm),
optax.inject_hyperparams(optax.adam)(
learning_rate=linear_schedule
learning_rate=partial(
linear_schedule,
minibatch_count=self.args.num_minibatches,
update_epochs=self.args.update_epochs,
num_iterations=self.args.num_iterations,
learning_rate=self.args.learning_rate,
)
if self.args.anneal_lr
else self.args.learning_rate,
eps=1e-5,
@ -296,12 +311,13 @@ class PPOTrainer:
self,
global_step,
episode_stats,
avg_episodic_return,
start_time,
iteration_time_start,
loss_info,
):
self.writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step)
self.writer.add_scalar(
"charts/avg_episodic_return", loss_info.avg_episodic_return, global_step
)
self.writer.add_scalar(
"charts/avg_episodic_length",
np.mean(jax.device_get(episode_stats.returned_episode_lengths)),
@ -318,8 +334,6 @@ class PPOTrainer:
self.writer.add_scalar("losses/approx_kl", loss_info.approx_kl[-1, -1].item(), global_step)
self.writer.add_scalar("losses/loss", loss_info.loss[-1, -1].item(), global_step)
# iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}")
self.writer.add_scalar(
"charts/SPS", int(global_step / (time.time() - start_time)), global_step
)
@ -330,8 +344,18 @@ class PPOTrainer:
)
def _step(self, env_state, next_obs, next_done) -> tuple:
storage, next_obs, next_done, env_state = self._rollout(env_state, next_obs, next_done)
(
self.agent_state,
self.episode_stats,
next_obs,
next_done,
storage,
self.key,
next_env_state,
) = self._rollout(env_state, next_obs, next_done)
storage = self._compute_gae(storage, next_obs, next_done)
self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = (
self._ppo.update_ppo(self.agent_state, storage, self.key)
)
@ -339,7 +363,7 @@ class PPOTrainer:
avg_episodic_return = jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))
return (
env_state,
next_env_state,
next_obs,
next_done,
LossInfo(
@ -352,10 +376,26 @@ class PPOTrainer:
),
)
def close(self):
def _close(self):
self.env.close()
self.writer.close()
def _save_model(self, model_path: str):
with open(model_path, "wb") as f:
f.write(
flax.serialization.to_bytes(
[
vars(self.args),
[
self.agent_state.params["sensor_params"],
self.agent_state.params["actor_params"],
self.agent_state.params["critic_params"],
self.agent_state.params["feature_extractor_params"],
],
]
)
)
def train(self):
"""
Train the PPO agent for a specified number of iterations
@ -368,16 +408,33 @@ class PPOTrainer:
global_step = 0
start_time = time.time()
if self.args.track:
import wandb
wandb.init(
project=self.args.wandb_project_name,
entity=self.args.wandb_entity,
sync_tensorboard=True,
config=vars(self.args),
name=self.run_name,
save_code=True,
)
self.writer.add_text(
"hyperparameters",
"|param|value|\n|---|---|\n"
+ "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()),
)
for _ in tqdm.tqdm(range(self.args.num_iterations)):
iteration_time_start = time.time()
env_state, next_obs, next_done, loss_info = self._step(env_state, next_obs, next_done)
global_step += self.args.num_steps * self.args.num_envs
self._log(
global_step, self.episode_stats, 0, start_time, iteration_time_start, loss_info
)
self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info)
if self.args.save_model:
self._save_model(...)
model_path = f"runs/{self.run_name}/{self.args.exp_name}.cleanrl_model"
self._save_model(model_path=model_path)
self.close()
self._close()

View file

@ -1,5 +1,6 @@
import time
import torch
import tyro
from brittle_star_project.dataclasses import PPOArgs
@ -23,5 +24,7 @@ if __name__ == "__main__":
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
env = make_env(args.config_path, args.num_envs)
torch.backends.cudnn.deterministic = args.torch_deterministic
ppo_trainer = PPOTrainer(args, env, run_name)
ppo_trainer.train()