fix(PPOTrainer.py): cleaned up + bug fixes regarding misuse of variable/wrong returns
This commit is contained in:
parent
4396c9ac4b
commit
a9302cfd4e
6 changed files with 133 additions and 88 deletions
14
docs/api/simulate.md
Normal file
14
docs/api/simulate.md
Normal 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
1
docs/api/train.md
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
# TODO
|
||||||
|
|
@ -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
|
|
||||||
```
|
|
||||||
|
|
@ -1,26 +1,28 @@
|
||||||
|
import random
|
||||||
import time
|
import time
|
||||||
from dataclasses import asdict, dataclass
|
from dataclasses import asdict, dataclass
|
||||||
from functools import partial
|
from functools import partial
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import optax
|
import flax
|
||||||
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 jax
|
import jax
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
import numpy as np
|
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
|
from ppo import PPO
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -95,6 +97,36 @@ def _step_once(
|
||||||
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
|
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
|
@jax.jit
|
||||||
def _rollout_jit(
|
def _rollout_jit(
|
||||||
agent_state,
|
agent_state,
|
||||||
|
|
@ -126,36 +158,6 @@ def _rollout_jit(
|
||||||
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
|
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
|
@jax.jit
|
||||||
def compute_gae_once(carry, inp, gamma, gae_lambda):
|
def compute_gae_once(carry, inp, gamma, gae_lambda):
|
||||||
advantages = carry
|
advantages = carry
|
||||||
|
|
@ -200,7 +202,8 @@ class PPOTrainer:
|
||||||
def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_name: str):
|
def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_name: str):
|
||||||
self.args = args
|
self.args = args
|
||||||
self.env = env
|
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)
|
self.key = jax.random.PRNGKey(args.seed)
|
||||||
|
|
||||||
|
|
@ -216,6 +219,12 @@ class PPOTrainer:
|
||||||
|
|
||||||
self.episode_stats = self._init_episode_stats()
|
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):
|
def _init_agent(self):
|
||||||
sensor = GenericDenseLayersWithActivation()
|
sensor = GenericDenseLayersWithActivation()
|
||||||
feature_extractor = GenericDenseLayersWithActivation()
|
feature_extractor = GenericDenseLayersWithActivation()
|
||||||
|
|
@ -255,7 +264,13 @@ class PPOTrainer:
|
||||||
tx=optax.chain(
|
tx=optax.chain(
|
||||||
optax.clip_by_global_norm(self.args.max_grad_norm),
|
optax.clip_by_global_norm(self.args.max_grad_norm),
|
||||||
optax.inject_hyperparams(optax.adam)(
|
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
|
if self.args.anneal_lr
|
||||||
else self.args.learning_rate,
|
else self.args.learning_rate,
|
||||||
eps=1e-5,
|
eps=1e-5,
|
||||||
|
|
@ -296,12 +311,13 @@ class PPOTrainer:
|
||||||
self,
|
self,
|
||||||
global_step,
|
global_step,
|
||||||
episode_stats,
|
episode_stats,
|
||||||
avg_episodic_return,
|
|
||||||
start_time,
|
start_time,
|
||||||
iteration_time_start,
|
iteration_time_start,
|
||||||
loss_info,
|
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(
|
self.writer.add_scalar(
|
||||||
"charts/avg_episodic_length",
|
"charts/avg_episodic_length",
|
||||||
np.mean(jax.device_get(episode_stats.returned_episode_lengths)),
|
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/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)
|
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(
|
self.writer.add_scalar(
|
||||||
"charts/SPS", int(global_step / (time.time() - start_time)), global_step
|
"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:
|
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)
|
storage = self._compute_gae(storage, next_obs, next_done)
|
||||||
|
|
||||||
self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = (
|
self.agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, self.key = (
|
||||||
self._ppo.update_ppo(self.agent_state, storage, 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))
|
avg_episodic_return = jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))
|
||||||
|
|
||||||
return (
|
return (
|
||||||
env_state,
|
next_env_state,
|
||||||
next_obs,
|
next_obs,
|
||||||
next_done,
|
next_done,
|
||||||
LossInfo(
|
LossInfo(
|
||||||
|
|
@ -352,10 +376,26 @@ class PPOTrainer:
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def close(self):
|
def _close(self):
|
||||||
self.env.close()
|
self.env.close()
|
||||||
self.writer.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):
|
def train(self):
|
||||||
"""
|
"""
|
||||||
Train the PPO agent for a specified number of iterations
|
Train the PPO agent for a specified number of iterations
|
||||||
|
|
@ -368,16 +408,33 @@ class PPOTrainer:
|
||||||
global_step = 0
|
global_step = 0
|
||||||
start_time = time.time()
|
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)):
|
for _ in tqdm.tqdm(range(self.args.num_iterations)):
|
||||||
iteration_time_start = time.time()
|
iteration_time_start = time.time()
|
||||||
|
|
||||||
env_state, next_obs, next_done, loss_info = self._step(env_state, next_obs, next_done)
|
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
|
global_step += self.args.num_steps * self.args.num_envs
|
||||||
self._log(
|
self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info)
|
||||||
global_step, self.episode_stats, 0, start_time, iteration_time_start, loss_info
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.args.save_model:
|
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()
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
import time
|
import time
|
||||||
|
|
||||||
|
import torch
|
||||||
import tyro
|
import tyro
|
||||||
|
|
||||||
from brittle_star_project.dataclasses import PPOArgs
|
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())}"
|
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
|
||||||
env = make_env(args.config_path, args.num_envs)
|
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 = PPOTrainer(args, env, run_name)
|
||||||
ppo_trainer.train()
|
ppo_trainer.train()
|
||||||
|
|
|
||||||
Reference in a new issue