feat(PPOTrainer.py): added logging messages
This commit is contained in:
parent
84f6fbd76e
commit
b97fc30d5d
5 changed files with 135 additions and 27 deletions
|
|
@ -1,5 +1,5 @@
|
||||||
# Minimal config to verify HPC setup is functional.
|
# Minimal config to verify HPC setup is functional.
|
||||||
# Run with: python src/train.py --config-path configs/hpc/smoke_test.yaml
|
# Run with: python experiments/train.py --config-path configs/hpc/smoke_test.yaml
|
||||||
exp_name: "hpc_smoke_test"
|
exp_name: "hpc_smoke_test"
|
||||||
seed: 0
|
seed: 0
|
||||||
track: false # Test WandB integration
|
track: false # Test WandB integration
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,6 @@
|
||||||
|
import datetime
|
||||||
import random
|
import random
|
||||||
|
import sys
|
||||||
import time
|
import time
|
||||||
from dataclasses import asdict, dataclass
|
from dataclasses import asdict, dataclass
|
||||||
from functools import partial
|
from functools import partial
|
||||||
|
|
@ -200,11 +202,12 @@ class LossInfo:
|
||||||
|
|
||||||
|
|
||||||
class PPOTrainer:
|
class PPOTrainer:
|
||||||
def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_name: str):
|
def __init__(self, args: PPOArgs, env: BrittleStarJaxEnvWrapper, run_dir: str, run_name: str):
|
||||||
self.args = args
|
self.args = args
|
||||||
self.env = env
|
self.env = env
|
||||||
|
self.run_dir = run_dir
|
||||||
self.run_name = run_name
|
self.run_name = run_name
|
||||||
self.writer = SummaryWriter(f"runs/{self.run_name}")
|
self.writer = SummaryWriter(self.run_dir)
|
||||||
|
|
||||||
self.key = jax.random.PRNGKey(args.seed)
|
self.key = jax.random.PRNGKey(args.seed)
|
||||||
|
|
||||||
|
|
@ -244,11 +247,17 @@ class PPOTrainer:
|
||||||
|
|
||||||
self._init_random()
|
self._init_random()
|
||||||
|
|
||||||
def _init_random(self):
|
def _init_random(self, log: bool = True):
|
||||||
|
if log:
|
||||||
|
print(f"[RANDOM]: Setting random seed to {self.args.seed}")
|
||||||
|
|
||||||
random.seed(self.args.seed)
|
random.seed(self.args.seed)
|
||||||
np.random.seed(self.args.seed)
|
np.random.seed(self.args.seed)
|
||||||
|
|
||||||
def _init_agent(self):
|
def _init_agent(self, log: bool = True):
|
||||||
|
if log:
|
||||||
|
print("[AGENT]: Initializing agent...")
|
||||||
|
|
||||||
sensor = GenericDenseLayersWithActivation()
|
sensor = GenericDenseLayersWithActivation()
|
||||||
feature_extractor = GenericDenseLayersWithActivation()
|
feature_extractor = GenericDenseLayersWithActivation()
|
||||||
actor = Actor(
|
actor = Actor(
|
||||||
|
|
@ -258,7 +267,10 @@ class PPOTrainer:
|
||||||
# messenger = OneDenseLayerMLP()
|
# messenger = OneDenseLayerMLP()
|
||||||
return sensor, feature_extractor, actor, critic
|
return sensor, feature_extractor, actor, critic
|
||||||
|
|
||||||
def _init_agent_state(self) -> TrainState:
|
def _init_agent_state(self, log: bool = True) -> TrainState:
|
||||||
|
if log:
|
||||||
|
print("[AGENT STATE]: Initializing agent state...")
|
||||||
|
|
||||||
self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(
|
self.key, sensor_key, actor_key, critic_key, feature_extractor_key = jax.random.split(
|
||||||
self.key, 5
|
self.key, 5
|
||||||
)
|
)
|
||||||
|
|
@ -301,7 +313,10 @@ class PPOTrainer:
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_episode_stats(self) -> EpisodeStatistics:
|
def _init_episode_stats(self, log: bool = True) -> EpisodeStatistics:
|
||||||
|
if log:
|
||||||
|
print("[EPISODE STATS]: Initializing episode stats...")
|
||||||
|
|
||||||
return EpisodeStatistics(
|
return EpisodeStatistics(
|
||||||
episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32),
|
episode_returns=jnp.zeros(self.args.num_envs, dtype=jnp.float32),
|
||||||
episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32),
|
episode_lengths=jnp.zeros(self.args.num_envs, dtype=jnp.int32),
|
||||||
|
|
@ -335,6 +350,7 @@ class PPOTrainer:
|
||||||
iteration_time_start,
|
iteration_time_start,
|
||||||
loss_info,
|
loss_info,
|
||||||
):
|
):
|
||||||
|
|
||||||
self.writer.add_scalar(
|
self.writer.add_scalar(
|
||||||
"charts/avg_episodic_return", loss_info.avg_episodic_return, global_step
|
"charts/avg_episodic_return", loss_info.avg_episodic_return, global_step
|
||||||
)
|
)
|
||||||
|
|
@ -353,7 +369,6 @@ class PPOTrainer:
|
||||||
self.writer.add_scalar("losses/entropy", loss_info.entropy_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/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)
|
||||||
|
|
||||||
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
|
||||||
)
|
)
|
||||||
|
|
@ -363,7 +378,10 @@ class PPOTrainer:
|
||||||
global_step,
|
global_step,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _step(self, env_state, next_obs, next_done) -> tuple:
|
def _step(self, env_state, next_obs, next_done, is_tty: bool, iteration: int) -> tuple:
|
||||||
|
if not is_tty and iteration == 1:
|
||||||
|
print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True)
|
||||||
|
|
||||||
(
|
(
|
||||||
self.agent_state,
|
self.agent_state,
|
||||||
self.episode_stats,
|
self.episode_stats,
|
||||||
|
|
@ -374,12 +392,21 @@ class PPOTrainer:
|
||||||
next_env_state,
|
next_env_state,
|
||||||
) = self._rollout(env_state, next_obs, next_done)
|
) = self._rollout(env_state, next_obs, next_done)
|
||||||
|
|
||||||
|
if not is_tty and iteration == 1:
|
||||||
|
print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True)
|
||||||
|
|
||||||
storage = self._compute_gae(storage, next_obs, next_done)
|
storage = self._compute_gae(storage, next_obs, next_done)
|
||||||
|
|
||||||
|
if not is_tty and iteration == 1:
|
||||||
|
print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True)
|
||||||
|
|
||||||
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)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if not is_tty and iteration == 1:
|
||||||
|
print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True)
|
||||||
|
|
||||||
avg_episodic_return = float(
|
avg_episodic_return = float(
|
||||||
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))
|
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns))
|
||||||
)
|
)
|
||||||
|
|
@ -402,7 +429,10 @@ class PPOTrainer:
|
||||||
self.env.close()
|
self.env.close()
|
||||||
self.writer.close()
|
self.writer.close()
|
||||||
|
|
||||||
def _save_model(self, model_path: str):
|
def _save_model(self, model_path: str, log: bool = True):
|
||||||
|
if log:
|
||||||
|
print(f"[SAVE]: Saving the model to: {model_path}...")
|
||||||
|
|
||||||
with open(model_path, "wb") as f:
|
with open(model_path, "wb") as f:
|
||||||
f.write(
|
f.write(
|
||||||
flax.serialization.to_bytes(
|
flax.serialization.to_bytes(
|
||||||
|
|
@ -418,21 +448,38 @@ class PPOTrainer:
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
def train(self):
|
def train(self, log: bool = True):
|
||||||
"""
|
"""
|
||||||
Train the PPO agent for a specified number of iterations
|
Train the PPO agent for a specified number of iterations
|
||||||
(passed through PPOArgs in constructor).
|
(passed through PPOArgs in constructor).
|
||||||
Closes the environment at the end of training.
|
Closes the environment at the end of training.
|
||||||
"""
|
"""
|
||||||
|
if log:
|
||||||
|
print(f"running name: {self.run_name}")
|
||||||
|
|
||||||
|
is_tty = sys.stdout.isatty()
|
||||||
|
if log:
|
||||||
|
print("[TRAIN]: Resetting environment...")
|
||||||
|
|
||||||
|
if not is_tty:
|
||||||
|
print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True)
|
||||||
|
|
||||||
env_state = self.env.reset(seed=self.args.seed)
|
env_state = self.env.reset(seed=self.args.seed)
|
||||||
next_obs = convert_obs_dict_to_array(env_state.observations)
|
next_obs = convert_obs_dict_to_array(env_state.observations)
|
||||||
next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_)
|
next_done = jnp.zeros(self.args.num_envs, dtype=jnp.bool_)
|
||||||
|
|
||||||
|
if log and not is_tty:
|
||||||
|
print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True)
|
||||||
|
|
||||||
global_step = 0
|
global_step = 0
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
|
|
||||||
if self.args.track:
|
if self.args.track:
|
||||||
import wandb
|
import wandb
|
||||||
|
|
||||||
|
if log:
|
||||||
|
print("[TRAIN]: Initializing Weights and Biases...")
|
||||||
|
|
||||||
wandb.init(
|
wandb.init(
|
||||||
project=self.args.wandb_project_name,
|
project=self.args.wandb_project_name,
|
||||||
entity=self.args.wandb_entity,
|
entity=self.args.wandb_entity,
|
||||||
|
|
@ -442,21 +489,47 @@ class PPOTrainer:
|
||||||
save_code=True,
|
save_code=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if log:
|
||||||
|
print("[TRAIN]: Adding hyperparameters to TensorBoard...")
|
||||||
|
|
||||||
self.writer.add_text(
|
self.writer.add_text(
|
||||||
"hyperparameters",
|
"hyperparameters",
|
||||||
"|param|value|\n|---|---|\n"
|
"|param|value|\n|---|---|\n"
|
||||||
+ "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()),
|
+ "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()),
|
||||||
)
|
)
|
||||||
|
|
||||||
for _ in tqdm.tqdm(range(self.args.num_iterations)):
|
iter_bar = tqdm.tqdm(
|
||||||
|
range(1, self.args.num_iterations + 1),
|
||||||
|
disable=not sys.stdout.isatty(),
|
||||||
|
)
|
||||||
|
for iteration in iter_bar:
|
||||||
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)
|
||||||
|
|
||||||
|
if not is_tty and iteration == 1:
|
||||||
|
print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True)
|
||||||
|
|
||||||
global_step += self.args.num_steps * self.args.num_envs
|
global_step += self.args.num_steps * self.args.num_envs
|
||||||
self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info)
|
self._log(global_step, self.episode_stats, start_time, iteration_time_start, loss_info)
|
||||||
|
|
||||||
|
if not is_tty:
|
||||||
|
sps = int(global_step / (time.time() - start_time))
|
||||||
|
remaining_steps = self.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}/{self.args.num_iterations} | "
|
||||||
|
f"Step {global_step}/{self.args.total_timesteps} | "
|
||||||
|
f"SPS {sps} | "
|
||||||
|
f"Return {loss_info.avg_episodic_return:.4f} | "
|
||||||
|
f"ETA {eta_str}",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
|
||||||
if self.args.save_model:
|
if self.args.save_model:
|
||||||
model_path = f"runs/{self.run_name}/{self.args.exp_name}.cleanrl_model"
|
model_path = f"{self.run_dir}/{self.args.exp_name}.cleanrl_model"
|
||||||
self._save_model(model_path=model_path)
|
self._save_model(model_path=model_path)
|
||||||
|
|
||||||
self._close()
|
self._close()
|
||||||
|
|
|
||||||
|
|
@ -81,7 +81,6 @@ def train(args: PPOArgs):
|
||||||
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
|
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 = args.total_timesteps // args.batch_size
|
||||||
args.num_iterations = 5
|
args.num_iterations = 5
|
||||||
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
|
|
||||||
print(f"running name: {run_name}")
|
print(f"running name: {run_name}")
|
||||||
|
|
||||||
if args.run_dir is None:
|
if args.run_dir is None:
|
||||||
|
|
@ -385,8 +384,6 @@ def train(args: PPOArgs):
|
||||||
)
|
)
|
||||||
|
|
||||||
if args.save_model:
|
if args.save_model:
|
||||||
model_path = f"runs/{run_name}/{args.exp_name}.cleanrl_model"
|
|
||||||
save_model(model_path, agent_state, args)
|
|
||||||
model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model"
|
model_path = f"{args.run_dir}/{args.exp_name}.cleanrl_model"
|
||||||
with open(model_path, "wb") as f:
|
with open(model_path, "wb") as f:
|
||||||
f.write(
|
f.write(
|
||||||
|
|
@ -408,12 +405,6 @@ def train(args: PPOArgs):
|
||||||
writer.close()
|
writer.close()
|
||||||
|
|
||||||
print("Saving loss plot...")
|
print("Saving loss plot...")
|
||||||
simple_plot(
|
|
||||||
list(range(len(returns))),
|
|
||||||
returns,
|
|
||||||
show_window=True,
|
|
||||||
filename=f"runs/{run_name}/{args.exp_name}_losses.png",
|
|
||||||
)
|
|
||||||
plt.plot(losses)
|
plt.plot(losses)
|
||||||
plt.title("PPO Loss, mean over minibatches")
|
plt.title("PPO Loss, mean over minibatches")
|
||||||
plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png")
|
plt.savefig(f"{args.run_dir}/{args.exp_name}_losses.png")
|
||||||
|
|
|
||||||
|
|
@ -1,7 +1,10 @@
|
||||||
|
import subprocess
|
||||||
import time
|
import time
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import tyro
|
import tyro
|
||||||
|
import yaml
|
||||||
|
import os
|
||||||
|
|
||||||
from brittle_star_project.dataclasses import PPOArgs
|
from brittle_star_project.dataclasses import PPOArgs
|
||||||
from PPOTrainer import PPOTrainer
|
from PPOTrainer import PPOTrainer
|
||||||
|
|
@ -14,16 +17,53 @@ def make_env(config_path: str | None, num_envs: int) -> BrittleStarJaxEnvWrapper
|
||||||
return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs)
|
return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs)
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args() -> PPOArgs:
|
||||||
|
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)
|
||||||
|
|
||||||
|
# Reparse CLI to ensure they OVERRIDE the yaml
|
||||||
|
args = tyro.cli(PPOArgs, default=temp_args)
|
||||||
|
else:
|
||||||
|
args = temp_args
|
||||||
|
return args
|
||||||
|
|
||||||
|
|
||||||
|
def get_git_hash() -> str:
|
||||||
|
try:
|
||||||
|
return (
|
||||||
|
subprocess.check_output(["git", "rev-parse", "--short", "HEAD"]).decode("ascii").strip()
|
||||||
|
)
|
||||||
|
except subprocess.CalledProcessError | UnicodeDecodeError:
|
||||||
|
return "none"
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
args = tyro.cli(PPOArgs)
|
args = parse_args()
|
||||||
|
|
||||||
args.batch_size = args.num_envs * args.num_steps
|
args.batch_size = args.num_envs * args.num_steps
|
||||||
args.minibatch_size = args.batch_size // args.num_minibatches
|
args.minibatch_size = args.batch_size // args.num_minibatches
|
||||||
args.num_iterations = args.total_timesteps // args.batch_size
|
args.num_iterations = args.total_timesteps // args.batch_size
|
||||||
run_name = f"{args.exp_name}__seed_{args.seed}__{int(time.time())}"
|
|
||||||
env = make_env(args.config_path, args.num_envs)
|
git_hash = get_git_hash()
|
||||||
|
run_name = f"{args.exp_name}__seed_{args.seed}__{git_hash}__{int(time.time())}"
|
||||||
|
if args.run_dir is None:
|
||||||
|
run_dir = f"runs/{run_name}"
|
||||||
|
else:
|
||||||
|
run_dir = args.run_dir
|
||||||
|
|
||||||
|
os.makedirs(run_dir, exist_ok=True)
|
||||||
|
|
||||||
|
env = make_env(args.env_config_path, args.num_envs)
|
||||||
|
|
||||||
torch.backends.cudnn.deterministic = args.torch_deterministic
|
torch.backends.cudnn.deterministic = args.torch_deterministic
|
||||||
|
|
||||||
ppo_trainer = PPOTrainer(args, env, run_name)
|
ppo_trainer = PPOTrainer(args, env, run_dir, run_name)
|
||||||
ppo_trainer.train()
|
ppo_trainer.train()
|
||||||
|
|
|
||||||
6
uv.lock
generated
6
uv.lock
generated
|
|
@ -35,6 +35,9 @@ dependencies = [
|
||||||
]
|
]
|
||||||
|
|
||||||
[package.optional-dependencies]
|
[package.optional-dependencies]
|
||||||
|
analysis = [
|
||||||
|
{ name = "tensorboard" },
|
||||||
|
]
|
||||||
cuda = [
|
cuda = [
|
||||||
{ name = "jax", extra = ["cuda13"] },
|
{ name = "jax", extra = ["cuda13"] },
|
||||||
]
|
]
|
||||||
|
|
@ -64,12 +67,13 @@ requires-dist = [
|
||||||
{ name = "protobuf", specifier = ">=5.0.0" },
|
{ name = "protobuf", specifier = ">=5.0.0" },
|
||||||
{ name = "pyopengl", specifier = ">=3.1.10" },
|
{ name = "pyopengl", specifier = ">=3.1.10" },
|
||||||
{ name = "pyopengl-accelerate", specifier = ">=3.1.10" },
|
{ name = "pyopengl-accelerate", specifier = ">=3.1.10" },
|
||||||
|
{ name = "tensorboard", marker = "extra == 'analysis'" },
|
||||||
{ name = "torch", specifier = ">=2.4.0" },
|
{ name = "torch", specifier = ">=2.4.0" },
|
||||||
{ name = "tyro", specifier = ">=1.0.10" },
|
{ name = "tyro", specifier = ">=1.0.10" },
|
||||||
{ name = "wandb", specifier = "==0.24.2" },
|
{ name = "wandb", specifier = "==0.24.2" },
|
||||||
{ name = "warp-lang" },
|
{ name = "warp-lang" },
|
||||||
]
|
]
|
||||||
provides-extras = ["cuda"]
|
provides-extras = ["cuda", "analysis"]
|
||||||
|
|
||||||
[package.metadata.requires-dev]
|
[package.metadata.requires-dev]
|
||||||
dev = [
|
dev = [
|
||||||
|
|
|
||||||
Reference in a new issue