1
Fork 0

feat(experiments): renamed _train_backup.py to _train_backup.py.back to remove duplicated code warning in pycharm

This commit is contained in:
Robin Meersman 2026-04-06 13:58:08 +02:00
parent b97fc30d5d
commit d5b911fa32

View file

@ -1,435 +0,0 @@
import datetime
import random
import yaml
import subprocess
import sys
import time
from dataclasses import asdict
from functools import partial
from typing import Callable
import flax
import jax
import jax.numpy as jnp
import matplotlib.pyplot as plt
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(env_config_path: str | None, num_envs: int) -> Callable:
def thunk():
if env_config_path is None:
return BrittleStarJaxEnvWrapper.default(num_envs=num_envs)
return BrittleStarJaxEnvWrapper.from_config(env_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
# Try to get git short hash
try:
git_hash = (
subprocess.check_output(["git", "rev-parse", "--short", "HEAD"]).decode("ascii").strip()
)
except Exception:
git_hash = "none"
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
# args.num_iterations = args.total_timesteps // args.batch_size
args.num_iterations = 5
print(f"running name: {run_name}")
if args.run_dir is None:
args.run_dir = f"runs/{run_name}"
import os
os.makedirs(args.run_dir, exist_ok=True)
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(args.run_dir)
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(env_config_path=args.env_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...")
if not sys.stdout.isatty():
print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True)
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_)
if not sys.stdout.isatty():
print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True)
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),
disable=not sys.stdout.isatty(),
)
returns = []
is_tty = sys.stdout.isatty()
for iteration in iters_bar:
iteration_time_start = time.time()
if not is_tty and iteration == 1:
print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True)
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
)
if not is_tty and iteration == 1:
print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True)
global_step += args.num_steps * args.num_envs
storage = compute_gae(agent_state, next_obs, next_done, storage)
if not is_tty and iteration == 1:
print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True)
agent_state, loss, pg_loss, v_loss, entropy_loss, approx_kl, key = ppo_instance.update_ppo(
agent_state, storage, key
)
if not is_tty and iteration == 1:
print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True)
losses.append(jnp.mean(loss))
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}"
)
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 not is_tty:
sps = int(global_step / (time.time() - start_time))
remaining_steps = args.total_timesteps - global_step
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
eta_str = str(datetime.timedelta(seconds=eta_seconds))
print(
f"Iteration {iteration}/{args.num_iterations} | "
f"Step {global_step}/{args.total_timesteps} | "
f"SPS {sps} | "
f"Return {avg_episodic_return:.4f} | "
f"ETA {eta_str}",
flush=True,
)
if args.save_model:
model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model"
with open(model_path, "wb") as f:
f.write(
flax.serialization.to_bytes(
[
vars(args),
[
agent_state.params["sensor_params"],
agent_state.params["actor_params"],
agent_state.params["critic_params"],
agent_state.params["feature_extractor_params"],
],
]
)
)
print(f"model saved to {model_path}")
env.close()
writer.close()
print("Saving loss plot...")
plt.plot(losses)
plt.title("PPO Loss, mean over minibatches")
plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png")
plt.close()
def main() -> None:
temp_args = tyro.cli(PPOArgs)
if temp_args.env_config_path is not None:
with open(temp_args.env_config_path, "r") as f:
config = yaml.safe_load(f)
if config:
# parse PPOArgs with defaults from yaml.
for key, value in config.items():
if hasattr(temp_args, key):
setattr(temp_args, key, value)
# Re-parse CLI to ensure they OVERRIDE the yaml
args = tyro.cli(PPOArgs, default=temp_args)
else:
args = temp_args
train(args)
if __name__ == "__main__":
main()