1
Fork 0

feat(train.py, PPOTrainer.py): cleaned up training loop to specialized class

This commit is contained in:
Robin Meersman 2026-04-03 14:23:09 +02:00
parent 24da8948c5
commit 4396c9ac4b
3 changed files with 743 additions and 333 deletions

383
experiments/PPOTrainer.py Normal file
View file

@ -0,0 +1,383 @@
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 jax
import jax.numpy as jnp
import numpy as np
from ppo import PPO
@jax.jit
def linear_schedule(count, minibatch_count, update_epochs, num_iterations, learning_rate):
frac = 1.0 - (count // (minibatch_count * update_epochs)) / num_iterations
return learning_rate * frac
@jax.jit
def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))(
obs_dict
)
@jax.jit
def get_action_and_value_noise(
sensor: GenericDenseLayersWithActivation,
feature_extractor: GenericDenseLayersWithActivation,
actor: Actor,
critic: OneDenseLayerMLP,
agent_state: TrainState,
next_obs: jnp.ndarray,
key: jax.random.PRNGKey,
):
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
hidden_critic = feature_extractor.apply(
agent_state.params["feature_extractor_params"], next_obs
)
# Continuous actions: sample from a Gaussian parameterized by the actor
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
key, subkey = jax.random.split(key)
noise = jax.random.normal(subkey, shape=mean.shape)
std = jnp.exp(log_std)
action = mean + noise * std
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
value = critic.apply(agent_state.params["critic_params"], hidden_critic)
return action, logprob, value.squeeze(-1), key
@jax.jit
def _step_once(
carry,
_,
env_step_fn,
sensor: GenericDenseLayersWithActivation,
feature_extractor: GenericDenseLayersWithActivation,
actor: Actor,
critic: OneDenseLayerMLP,
):
agent_state, episode_stats, obs, done, key, env_state = carry
action, logprob, value, key = get_action_and_value_noise(
sensor, feature_extractor, actor, critic, agent_state, obs, key
)
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
episode_stats, env_state, action
)
storage = Storage(
obs=obs,
actions=action,
logprobs=logprob,
dones=done,
values=value,
rewards=reward,
returns=jnp.zeros_like(reward),
advantages=jnp.zeros_like(reward),
)
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
@jax.jit
def _rollout_jit(
agent_state,
episode_stats,
env_state,
next_obs,
next_done,
key,
args,
env_step_fn,
sensor: GenericDenseLayersWithActivation,
feature_extractor: GenericDenseLayersWithActivation,
actor: Actor,
critic: OneDenseLayerMLP,
):
(agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan(
partial(
_step_once,
sensor=sensor,
feature_extractor=feature_extractor,
actor=actor,
critic=critic,
env_step_fn=partial(_step_env_wrapped, env_step_fn=env_step_fn),
),
(agent_state, episode_stats, next_obs, next_done, key, env_state),
(),
args.max_steps,
)
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
nextdone, nextvalues, curvalues, reward = inp
nextnonterminal = 1.0 - nextdone
delta = reward + gamma * nextvalues * nextnonterminal - curvalues
advantages = delta + gamma * gae_lambda * nextnonterminal * advantages
return advantages, advantages
@jax.jit
def compute_gae_jit(agent_state, storage, next_obs, next_done, sensor, critic, args):
next_value = critic.apply(
agent_state.params["critic_params"],
sensor.apply(agent_state.params["sensor_params"], next_obs),
).squeeze(-1)
advantages = jnp.zeros((args.num_envs,))
dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0)
values = jnp.concatenate([storage.values, next_value[None, :]], axis=0)
_, advantages = jax.lax.scan(
partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda),
advantages,
(dones[1:], values[1:], values[:-1], storage.rewards),
reverse=True,
)
return storage.replace(advantages=advantages, returns=advantages + storage.values)
@dataclass
class LossInfo:
# todo: better typing
loss: Any
pg_loss: Any
v_loss: Any
entropy_loss: Any
approx_kl: Any
avg_episodic_return: Any
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.key = jax.random.PRNGKey(args.seed)
self.sensor, self.feature_extractor, self.actor, self.critic = self._init_agent()
self.sensor.apply = jax.jit(self.sensor.apply)
self.feature_extractor.apply = jax.jit(self.feature_extractor.apply)
self.actor.apply = jax.jit(self.actor.apply)
self.critic.apply = jax.jit(self.critic.apply)
self._ppo = PPO(self.args, self.sensor, self.actor, self.critic, self.feature_extractor)
self.agent_state = self._init_agent_state()
self.episode_stats = self._init_episode_stats()
def _init_agent(self):
sensor = GenericDenseLayersWithActivation()
feature_extractor = GenericDenseLayersWithActivation()
actor = Actor(
action_dim=self.env.single_action_space.shape[0]
) # continuous actions for MJX
critic = OneDenseLayerMLP()
# messenger = OneDenseLayerMLP()
return sensor, feature_extractor, actor, critic
def _init_agent_state(self) -> TrainState:
self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(
self.key, 5
)
sample_obs = jnp.concatenate(
[
v.flatten()
for v in self.env.single_observation_space.sample(
rng=jax.random.PRNGKey(0)
).values()
if v.size > 0
]
)
sensor_params = self.sensor.init(sensor_key, sample_obs)
feature_extractor_params = self.feature_extractor.init(feature_extractor_key, sample_obs)
actor_params = self.actor.init(actor_key, self.sensor.apply(sensor_params, sample_obs))
critic_params = self.critic.init(
critic_key, self.feature_extractor.apply(feature_extractor_params, sample_obs)
)
return TrainState.create(
apply_fn=None,
params=asdict(
AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params)
),
tx=optax.chain(
optax.clip_by_global_norm(self.args.max_grad_norm),
optax.inject_hyperparams(optax.adam)(
learning_rate=linear_schedule
if self.args.anneal_lr
else self.args.learning_rate,
eps=1e-5,
),
),
)
def _init_episode_stats(self) -> EpisodeStatistics:
return EpisodeStatistics(
episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32),
episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32),
returned_episode_returns=jnp.zeros(self.args.num_envs, jnp.float32),
returned_episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32),
)
def _rollout(self, env_state, next_obs, next_done) -> tuple[Storage, ...]:
return _rollout_jit(
self.agent_state,
self.episode_stats,
env_state,
next_obs,
next_done,
self.key,
self.args,
self.env.step,
self.sensor,
self.feature_extractor,
self.actor,
self.critic,
)
def _compute_gae(self, storage, next_obs, next_done) -> Storage:
return compute_gae_jit(
self.agent_state, storage, next_obs, next_done, self.sensor, self.critic, self.args
)
def _log(
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_length",
np.mean(jax.device_get(episode_stats.returned_episode_lengths)),
global_step,
)
self.writer.add_scalar(
"charts/learning_rate",
self.agent_state.opt_state[1].hyperparams["learning_rate"].item(),
global_step,
)
self.writer.add_scalar("losses/value_loss", loss_info.v_loss[-1, -1].item(), global_step)
self.writer.add_scalar("losses/policy_loss", loss_info.pg_loss[-1, -1].item(), global_step)
self.writer.add_scalar("losses/entropy", loss_info.entropy_loss[-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)
# 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
)
self.writer.add_scalar(
"charts/SPS_update",
int(self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start)),
global_step,
)
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)
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)
)
avg_episodic_return = jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))
return (
env_state,
next_obs,
next_done,
LossInfo(
loss=loss,
pg_loss=pg_loss,
v_loss=v_loss,
entropy_loss=entropy_loss,
approx_kl=approx_kl,
avg_episodic_return=avg_episodic_return,
),
)
def close(self):
self.env.close()
self.writer.close()
def train(self):
"""
Train the PPO agent for a specified number of iterations
(passed through PPOArgs in constructor).
Closes the environment at the end of training.
"""
env_state = self.env.reset(seed=self.args.seed)
next_obs = convert_obs_dict_to_array(env_state.observations)
next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_)
global_step = 0
start_time = time.time()
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
)
if self.args.save_model:
self._save_model(...)
self.close()

View file

@ -1,350 +1,27 @@
import random
import time
from dataclasses import asdict
from functools import partial
from typing import Callable
import flax
import jax
import jax.numpy as jnp
import numpy as np
import optax
import torch
import tqdm
import tyro
from flax.training.train_state import TrainState
from torch.utils.tensorboard import SummaryWriter
from brittle_star_project.dataclasses import PPOArgs
from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics
from PPOTrainer import PPOTrainer
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from MLPs.mlps import (
GenericDenseLayersWithActivation,
Actor,
OneDenseLayerMLP,
AgentParams,
Storage,
)
from plots import simple_plot
from ppo import PPO
def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))(
obs_dict
)
def make_env(config_path: str | None, num_envs: int) -> BrittleStarJaxEnvWrapper:
if config_path is None:
return BrittleStarJaxEnvWrapper.default(num_envs=num_envs)
return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs)
def make_env(config_path: str | None, num_envs: int) -> Callable:
def thunk():
if config_path is None:
return BrittleStarJaxEnvWrapper.default(num_envs=num_envs)
return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs)
if __name__ == "__main__":
args = tyro.cli(PPOArgs)
return thunk
def save_model(model_path: str, agent_state: TrainState, args: PPOArgs):
with open(model_path, "wb") as f:
f.write(
flax.serialization.to_bytes(
[
vars(args),
[
agent_state.params["network_params"],
agent_state.params["actor_params"],
agent_state.params["critic_params"],
],
]
)
)
def train(args: PPOArgs):
args.batch_size = args.num_envs * args.num_steps
args.minibatch_size = args.batch_size // args.num_minibatches
# args.num_iterations = args.total_timesteps // args.batch_size
args.num_iterations = 5
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
print(f"running name: {run_name}")
env = make_env(args.config_path, args.num_envs)
if args.track:
import wandb
wandb.init(
project=args.wandb_project_name,
entity=args.wandb_entity,
sync_tensorboard=True,
config=vars(args),
name=run_name,
save_code=True,
)
writer = SummaryWriter(f"runs/{run_name}")
writer.add_text(
"hyperparameters",
"|param|value|\n|---|---|\n" + "\n".join(f"|{k}|{v}|" for k, v in vars(args).items()),
)
random.seed(args.seed)
np.random.seed(args.seed)
key = jax.random.PRNGKey(args.seed)
key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5)
torch.backends.cudnn.deterministic = args.torch_deterministic
device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu")
print(f"Running on device: {device}")
print("Creating the environment...")
env = make_env(config_path=args.config_path, num_envs=args.num_envs)()
print(f"Environment: {env}")
episode_stats = EpisodeStatistics(
episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32),
episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32),
returned_episode_returns=jnp.zeros(args.num_envs, jnp.float32),
returned_episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32),
)
def step_env_wrapped(episode_stats: EpisodeStatistics, env_state, action):
next_env_state = env.step(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),
)
def linear_schedule(count):
frac = 1.0 - (count // (args.num_minibatches * args.update_epochs)) / args.num_iterations
return args.learning_rate * frac
print("Initializing the models...")
sensor = GenericDenseLayersWithActivation()
feature_extractor = GenericDenseLayersWithActivation()
actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX
critic = OneDenseLayerMLP()
# messager = OneDenseLayerMLP()
sample_obs = jnp.concatenate(
[
v.flatten()
for v in env.single_observation_space.sample(rng=jax.random.PRNGKey(0)).values()
if v.size > 0
]
)
sensor_params = sensor.init(sensor_key, sample_obs)
feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs)
actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs))
critic_params = critic.init(
critic_key, feature_extractor.apply(feature_extractor_params, sample_obs)
)
agent_state = TrainState.create(
apply_fn=None,
params=asdict(
AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params)
),
tx=optax.chain(
optax.clip_by_global_norm(args.max_grad_norm),
optax.inject_hyperparams(optax.adam)(
learning_rate=linear_schedule if args.anneal_lr else args.learning_rate, eps=1e-5
),
),
)
sensor.apply = jax.jit(sensor.apply)
feature_extractor.apply = jax.jit(feature_extractor.apply)
actor.apply = jax.jit(actor.apply)
critic.apply = jax.jit(critic.apply)
ppo_instance = PPO(args, sensor, actor, critic, feature_extractor)
@jax.jit
def get_action_and_value_noise(
agent_state: TrainState,
next_obs: jnp.ndarray,
key: jax.random.PRNGKey,
):
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
hidden_critic = feature_extractor.apply(
agent_state.params["feature_extractor_params"], next_obs
)
# Continuous actions: sample from a Gaussian parameterized by the actor
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
key, subkey = jax.random.split(key)
noise = jax.random.normal(subkey, shape=mean.shape)
std = jnp.exp(log_std)
action = mean + noise * std
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
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):
advantages = carry
nextdone, nextvalues, curvalues, reward = inp
nextnonterminal = 1.0 - nextdone
delta = reward + gamma * nextvalues * nextnonterminal - curvalues
advantages = delta + gamma * gae_lambda * nextnonterminal * advantages
return advantages, advantages
@jax.jit
def compute_gae(agent_state, next_obs, next_done, storage):
next_value = critic.apply(
agent_state.params["critic_params"],
sensor.apply(agent_state.params["sensor_params"], next_obs),
).squeeze(-1)
advantages = jnp.zeros((args.num_envs,))
dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0)
values = jnp.concatenate([storage.values, next_value[None, :]], axis=0)
_, advantages = jax.lax.scan(
partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda),
advantages,
(dones[1:], values[1:], values[:-1], storage.rewards),
reverse=True,
)
return storage.replace(advantages=advantages, returns=advantages + storage.values)
# END GAE
# --- Main training loop ---
global_step = 0
start_time = time.time()
# Reset once to get initial state
print("Resetting the environment...")
next_env_state = env.reset(seed=args.seed)
next_obs = convert_obs_dict_to_array(next_env_state.observations)
next_done = jnp.zeros(args.num_envs, dtype=jnp.bool_)
def step_once(carry, _, env_step_fn):
agent_state, episode_stats, obs, done, key, env_state = carry
action, logprob, value, key = get_action_and_value_noise(agent_state, obs, key)
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
episode_stats, env_state, action
)
storage = Storage(
obs=obs,
actions=action,
logprobs=logprob,
dones=done,
values=value,
rewards=reward,
returns=jnp.zeros_like(reward),
advantages=jnp.zeros_like(reward),
)
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
def rollout(
agent_state, episode_stats, next_obs, next_done, key, env_state, step_once_fn, max_steps
):
(agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan(
step_once_fn,
(agent_state, episode_stats, next_obs, next_done, key, env_state),
(),
max_steps,
)
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
rollout = partial(
rollout,
step_once_fn=partial(step_once, env_step_fn=step_env_wrapped),
max_steps=args.num_steps,
)
print("Starting training...")
iters_bar = tqdm.tqdm(range(1, args.num_iterations + 1))
returns = []
for _ in iters_bar:
iteration_time_start = time.time()
agent_state, episode_stats, next_obs, next_done, storage, key, next_env_state = rollout(
agent_state, episode_stats, next_obs, next_done, key, next_env_state
)
global_step += args.num_steps * args.num_envs
storage = compute_gae(agent_state, next_obs, next_done, storage)
agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo(
agent_state, storage, key
)
avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns))
iters_bar.set_postfix_str(
f"global_step={global_step}, avg_episodic_return={avg_episodic_return}"
)
returns.append(avg_episodic_return)
writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step)
writer.add_scalar(
"charts/avg_episodic_length",
np.mean(jax.device_get(episode_stats.returned_episode_lengths)),
global_step,
)
writer.add_scalar(
"charts/learning_rate",
agent_state.opt_state[1].hyperparams["learning_rate"].item(),
global_step,
)
writer.add_scalar("losses/value_loss", v_loss[-1, -1].item(), global_step)
writer.add_scalar("losses/policy_loss", pg_loss[-1, -1].item(), global_step)
writer.add_scalar("losses/entropy", entropy_loss[-1, -1].item(), global_step)
writer.add_scalar("losses/approx_kl", approx_kl[-1, -1].item(), global_step)
writer.add_scalar("losses/loss", loss[-1, -1].item(), global_step)
# iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}")
writer.add_scalar("charts/SPS", int(global_step / (time.time() - start_time)), global_step)
writer.add_scalar(
"charts/SPS_update",
int(args.num_envs * args.num_steps / (time.time() - iteration_time_start)),
global_step,
)
if args.save_model:
model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model"
save_model(model_path, agent_state, args)
print(f"model saved to {model_path}")
env.close()
writer.close()
print("Saving loss plot...")
simple_plot(
list(range(len(returns))),
returns,
show_window=True,
filename=f"runs/{run_name}/{args.exp_name}_losses.png",
)
def main() -> None:
args = tyro.cli(PPOArgs)
train(args)
if __name__ == "__main__":
main()
ppo_trainer = PPOTrainer(args, env, run_name)
ppo_trainer.train()

350
experiments/train.py.back Normal file
View file

@ -0,0 +1,350 @@
import random
import time
from dataclasses import asdict
from functools import partial
from typing import Callable
import flax
import jax
import jax.numpy as jnp
import numpy as np
import optax
import torch
import tqdm
import tyro
from flax.training.train_state import TrainState
from torch.utils.tensorboard import SummaryWriter
from brittle_star_project.dataclasses import PPOArgs
from brittle_star_project.dataclasses.EpisodeStatistics import EpisodeStatistics
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
from MLPs.mlps import (
GenericDenseLayersWithActivation,
Actor,
OneDenseLayerMLP,
AgentParams,
Storage,
)
from plots import simple_plot
from ppo import PPO
def convert_obs_dict_to_array(obs_dict: dict) -> jnp.ndarray:
return jax.vmap(lambda o: jnp.concatenate([v.flatten() for v in o.values() if v.size > 0]))(
obs_dict
)
def make_env(config_path: str | None, num_envs: int) -> Callable:
def thunk():
if config_path is None:
return BrittleStarJaxEnvWrapper.default(num_envs=num_envs)
return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs)
return thunk
def save_model(model_path: str, agent_state: TrainState, args: PPOArgs):
with open(model_path, "wb") as f:
f.write(
flax.serialization.to_bytes(
[
vars(args),
[
agent_state.params["network_params"],
agent_state.params["actor_params"],
agent_state.params["critic_params"],
],
]
)
)
def train(args: PPOArgs):
args.batch_size = args.num_envs * args.num_steps
args.minibatch_size = args.batch_size // args.num_minibatches
# args.num_iterations = args.total_timesteps // args.batch_size
args.num_iterations = 5
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
print(f"running name: {run_name}")
if args.track:
import wandb
wandb.init(
project=args.wandb_project_name,
entity=args.wandb_entity,
sync_tensorboard=True,
config=vars(args),
name=run_name,
save_code=True,
)
writer = SummaryWriter(f"runs/{run_name}")
writer.add_text(
"hyperparameters",
"|param|value|\n|---|---|\n" + "\n".join(f"|{k}|{v}|" for k, v in vars(args).items()),
)
random.seed(args.seed)
np.random.seed(args.seed)
key = jax.random.PRNGKey(args.seed)
key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(key, 5)
torch.backends.cudnn.deterministic = args.torch_deterministic
device = torch.device("cuda" if torch.cuda.is_available() and args.cuda else "cpu")
print(f"Running on device: {device}")
print("Creating the environment...")
env = make_env(config_path=args.config_path, num_envs=args.num_envs)()
print(f"Environment: {env}")
episode_stats = EpisodeStatistics(
episode_returns=jnp.zeros(args.num_envs, dtype=jnp.float32),
episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32),
returned_episode_returns=jnp.zeros(args.num_envs, jnp.float32),
returned_episode_lengths=jnp.zeros(args.num_envs, dtype=jnp.int32),
)
def step_env_wrapped(episode_stats: EpisodeStatistics, env_state, action):
next_env_state = env.step(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),
)
def linear_schedule(count):
frac = 1.0 - (count // (args.num_minibatches * args.update_epochs)) / args.num_iterations
return args.learning_rate * frac
print("Initializing the models...")
sensor = GenericDenseLayersWithActivation()
feature_extractor = GenericDenseLayersWithActivation()
actor = Actor(action_dim=env.single_action_space.shape[0]) # continuous actions for MJX
critic = OneDenseLayerMLP()
# messager = OneDenseLayerMLP()
sample_obs = jnp.concatenate(
[
v.flatten()
for v in env.single_observation_space.sample(rng=jax.random.PRNGKey(0)).values()
if v.size > 0
]
)
sensor_params = sensor.init(sensor_key, sample_obs)
feature_extractor_params = feature_extractor.init(feature_extractor_key, sample_obs)
actor_params = actor.init(actor_key, sensor.apply(sensor_params, sample_obs))
critic_params = critic.init(
critic_key, feature_extractor.apply(feature_extractor_params, sample_obs)
)
agent_state = TrainState.create(
apply_fn=None,
params=asdict(
AgentParams(sensor_params, actor_params, critic_params, feature_extractor_params)
),
tx=optax.chain(
optax.clip_by_global_norm(args.max_grad_norm),
optax.inject_hyperparams(optax.adam)(
learning_rate=linear_schedule if args.anneal_lr else args.learning_rate, eps=1e-5
),
),
)
sensor.apply = jax.jit(sensor.apply)
feature_extractor.apply = jax.jit(feature_extractor.apply)
actor.apply = jax.jit(actor.apply)
critic.apply = jax.jit(critic.apply)
ppo_instance = PPO(args, sensor, actor, critic, feature_extractor)
@jax.jit
def get_action_and_value_noise(
agent_state: TrainState,
next_obs: jnp.ndarray,
key: jax.random.PRNGKey,
):
hidden = sensor.apply(agent_state.params["sensor_params"], next_obs)
hidden_critic = feature_extractor.apply(
agent_state.params["feature_extractor_params"], next_obs
)
# Continuous actions: sample from a Gaussian parameterized by the actor
mean, log_std = actor.apply(agent_state.params["actor_params"], hidden)
key, subkey = jax.random.split(key)
noise = jax.random.normal(subkey, shape=mean.shape)
std = jnp.exp(log_std)
action = mean + noise * std
logprob = -0.5 * (((action - mean) / std) ** 2 + 2 * log_std + jnp.log(2 * jnp.pi)).sum(-1)
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):
advantages = carry
nextdone, nextvalues, curvalues, reward = inp
nextnonterminal = 1.0 - nextdone
delta = reward + gamma * nextvalues * nextnonterminal - curvalues
advantages = delta + gamma * gae_lambda * nextnonterminal * advantages
return advantages, advantages
@jax.jit
def compute_gae(agent_state, next_obs, next_done, storage):
next_value = critic.apply(
agent_state.params["critic_params"],
sensor.apply(agent_state.params["sensor_params"], next_obs),
).squeeze(-1)
advantages = jnp.zeros((args.num_envs,))
dones = jnp.concatenate([storage.dones, next_done[None, :]], axis=0)
values = jnp.concatenate([storage.values, next_value[None, :]], axis=0)
_, advantages = jax.lax.scan(
partial(compute_gae_once, gamma=args.gamma, gae_lambda=args.gae_lambda),
advantages,
(dones[1:], values[1:], values[:-1], storage.rewards),
reverse=True,
)
return storage.replace(advantages=advantages, returns=advantages + storage.values)
# END GAE
# --- Main training loop ---
global_step = 0
start_time = time.time()
# Reset once to get initial state
print("Resetting the environment...")
next_env_state = env.reset(seed=args.seed)
next_obs = convert_obs_dict_to_array(next_env_state.observations)
next_done = jnp.zeros(args.num_envs, dtype=jnp.bool_)
def step_once(carry, _, env_step_fn):
agent_state, episode_stats, obs, done, key, env_state = carry
action, logprob, value, key = get_action_and_value_noise(agent_state, obs, key)
episode_stats, env_state, (next_obs, reward, next_done) = env_step_fn(
episode_stats, env_state, action
)
storage = Storage(
obs=obs,
actions=action,
logprobs=logprob,
dones=done,
values=value,
rewards=reward,
returns=jnp.zeros_like(reward),
advantages=jnp.zeros_like(reward),
)
return (agent_state, episode_stats, next_obs, next_done, key, env_state), storage
def rollout(
agent_state, episode_stats, next_obs, next_done, key, env_state, step_once_fn, max_steps
):
(agent_state, episode_stats, next_obs, next_done, key, env_state), storage = jax.lax.scan(
step_once_fn,
(agent_state, episode_stats, next_obs, next_done, key, env_state),
(),
max_steps,
)
return agent_state, episode_stats, next_obs, next_done, storage, key, env_state
rollout = partial(
rollout,
step_once_fn=partial(step_once, env_step_fn=step_env_wrapped),
max_steps=args.num_steps,
)
print("Starting training...")
iters_bar = tqdm.tqdm(range(1, args.num_iterations + 1))
returns = []
for _ in iters_bar:
iteration_time_start = time.time()
agent_state, episode_stats, next_obs, next_done, storage, key, next_env_state = rollout(
agent_state, episode_stats, next_obs, next_done, key, next_env_state
)
global_step += args.num_steps * args.num_envs
storage = compute_gae(agent_state, next_obs, next_done, storage)
agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo(
agent_state, storage, key
)
avg_episodic_return = np.mean(jax.device_get(episode_stats.returned_episode_returns))
iters_bar.set_postfix_str(
f"global_step={global_step}, avg_episodic_return={avg_episodic_return}"
)
returns.append(avg_episodic_return)
writer.add_scalar("charts/avg_episodic_return", avg_episodic_return, global_step)
writer.add_scalar(
"charts/avg_episodic_length",
np.mean(jax.device_get(episode_stats.returned_episode_lengths)),
global_step,
)
writer.add_scalar(
"charts/learning_rate",
agent_state.opt_state[1].hyperparams["learning_rate"].item(),
global_step,
)
writer.add_scalar("losses/value_loss", v_loss[-1, -1].item(), global_step)
writer.add_scalar("losses/policy_loss", pg_loss[-1, -1].item(), global_step)
writer.add_scalar("losses/entropy", entropy_loss[-1, -1].item(), global_step)
writer.add_scalar("losses/approx_kl", approx_kl[-1, -1].item(), global_step)
writer.add_scalar("losses/loss", loss[-1, -1].item(), global_step)
# iters_bar.set_postfix_str(f"SPS: {int(global_step / (time.time() - start_time))}")
writer.add_scalar("charts/SPS", int(global_step / (time.time() - start_time)), global_step)
writer.add_scalar(
"charts/SPS_update",
int(args.num_envs * args.num_steps / (time.time() - iteration_time_start)),
global_step,
)
if args.save_model:
model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model"
save_model(model_path, agent_state, args)
print(f"model saved to {model_path}")
env.close()
writer.close()
print("Saving loss plot...")
simple_plot(
list(range(len(returns))),
returns,
show_window=True,
filename=f"runs/{run_name}/{args.exp_name}_losses.png",
)
def main() -> None:
args = tyro.cli(PPOArgs)
train(args)
if __name__ == "__main__":
main()