refactor(log): PPOTrainer
This commit is contained in:
parent
795ae48520
commit
faf31567e8
1 changed files with 61 additions and 107 deletions
|
|
@ -13,7 +13,8 @@ import numpy as np
|
||||||
import optax
|
import optax
|
||||||
import tqdm
|
import tqdm
|
||||||
from flax.training.train_state import TrainState
|
from flax.training.train_state import TrainState
|
||||||
from torch.utils.tensorboard import SummaryWriter
|
|
||||||
|
from experiment_logger import get_logger
|
||||||
|
|
||||||
from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs
|
from brittle_star_project.dataclasses import EpisodeStatistics, PPOArgs
|
||||||
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
|
from brittle_star_project.environment.BrittleStarJaxEnvWrapper import BrittleStarJaxEnvWrapper
|
||||||
|
|
@ -207,7 +208,7 @@ class PPOTrainer:
|
||||||
self.env = env
|
self.env = env
|
||||||
self.run_dir = run_dir
|
self.run_dir = run_dir
|
||||||
self.run_name = run_name
|
self.run_name = run_name
|
||||||
self.writer = SummaryWriter(self.run_dir)
|
self.logger = get_logger()
|
||||||
|
|
||||||
self.key = jax.random.PRNGKey(args.seed)
|
self.key = jax.random.PRNGKey(args.seed)
|
||||||
|
|
||||||
|
|
@ -247,16 +248,14 @@ class PPOTrainer:
|
||||||
|
|
||||||
self._init_random()
|
self._init_random()
|
||||||
|
|
||||||
def _init_random(self, log: bool = True):
|
def _init_random(self):
|
||||||
if log:
|
self.logger.info(f"[RANDOM]: Setting random seed to {self.args.seed}")
|
||||||
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, log: bool = True):
|
def _init_agent(self):
|
||||||
if log:
|
self.logger.info("[AGENT]: Initializing agent...")
|
||||||
print("[AGENT]: Initializing agent...")
|
|
||||||
|
|
||||||
sensor = GenericDenseLayersWithActivation()
|
sensor = GenericDenseLayersWithActivation()
|
||||||
feature_extractor = GenericDenseLayersWithActivation()
|
feature_extractor = GenericDenseLayersWithActivation()
|
||||||
|
|
@ -267,9 +266,8 @@ class PPOTrainer:
|
||||||
# messenger = OneDenseLayerMLP()
|
# messenger = OneDenseLayerMLP()
|
||||||
return sensor, feature_extractor, actor, critic
|
return sensor, feature_extractor, actor, critic
|
||||||
|
|
||||||
def _init_agent_state(self, log: bool = True) -> TrainState:
|
def _init_agent_state(self) -> TrainState:
|
||||||
if log:
|
self.logger.info("[AGENT STATE]: Initializing agent state...")
|
||||||
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
|
||||||
|
|
@ -313,9 +311,8 @@ class PPOTrainer:
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_episode_stats(self, log: bool = True) -> EpisodeStatistics:
|
def _init_episode_stats(self) -> EpisodeStatistics:
|
||||||
if log:
|
self.logger.info("[EPISODE STATS]: Initializing episode stats...")
|
||||||
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),
|
||||||
|
|
@ -350,39 +347,29 @@ class PPOTrainer:
|
||||||
iteration_time_start,
|
iteration_time_start,
|
||||||
loss_info,
|
loss_info,
|
||||||
):
|
):
|
||||||
|
metrics = {
|
||||||
|
"charts/avg_episodic_return": loss_info.avg_episodic_return,
|
||||||
|
"charts/avg_episodic_length": np.mean(
|
||||||
|
jax.device_get(episode_stats.returned_episode_lengths)
|
||||||
|
),
|
||||||
|
"charts/learning_rate": self.agent_state.opt_state[1]
|
||||||
|
.hyperparams["learning_rate"]
|
||||||
|
.item(),
|
||||||
|
"losses/value_loss": loss_info.v_loss[-1, -1].item(),
|
||||||
|
"losses/policy_loss": loss_info.pg_loss[-1, -1].item(),
|
||||||
|
"losses/entropy": loss_info.entropy_loss[-1, -1].item(),
|
||||||
|
"losses/approx_kl": loss_info.approx_kl[-1, -1].item(),
|
||||||
|
"losses/loss": loss_info.loss[-1, -1].item(),
|
||||||
|
"charts/SPS": int(global_step / (time.time() - start_time)),
|
||||||
|
"charts/SPS_update": int(
|
||||||
|
self.args.num_envs * self.args.num_steps / (time.time() - iteration_time_start)
|
||||||
|
),
|
||||||
|
}
|
||||||
|
self.logger.log(metrics, step=global_step)
|
||||||
|
|
||||||
self.writer.add_scalar(
|
def _step(self, env_state, next_obs, next_done, is_tty: bool, iteration: int) -> tuple:
|
||||||
"charts/avg_episodic_return", loss_info.avg_episodic_return, global_step
|
if not is_tty and iteration == 1:
|
||||||
)
|
self.logger.info(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}")
|
||||||
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)
|
|
||||||
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, is_tty: bool, iteration: int, log: bool = True
|
|
||||||
) -> tuple:
|
|
||||||
if log and not is_tty and iteration == 1:
|
|
||||||
print(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}", flush=True)
|
|
||||||
|
|
||||||
(
|
(
|
||||||
self.agent_state,
|
self.agent_state,
|
||||||
|
|
@ -394,20 +381,20 @@ 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 log and not is_tty and iteration == 1:
|
if not is_tty and iteration == 1:
|
||||||
print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True)
|
self.logger.info(f">>> [HPC] First rollout completed: {time.ctime()}")
|
||||||
|
|
||||||
storage = self._compute_gae(storage, next_obs, next_done)
|
storage = self._compute_gae(storage, next_obs, next_done)
|
||||||
|
|
||||||
if log and not is_tty and iteration == 1:
|
if not is_tty and iteration == 1:
|
||||||
print(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}", flush=True)
|
self.logger.info(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}")
|
||||||
|
|
||||||
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 log and not is_tty and iteration == 1:
|
if not is_tty and iteration == 1:
|
||||||
print(f">>> [HPC] First PPO update completed: {time.ctime()}", flush=True)
|
self.logger.info(f">>> [HPC] First PPO update completed: {time.ctime()}")
|
||||||
|
|
||||||
avg_episodic_return = float(
|
avg_episodic_return = float(
|
||||||
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item()
|
jnp.mean(jax.device_get(self.episode_stats.returned_episode_returns)).item()
|
||||||
|
|
@ -429,77 +416,45 @@ class PPOTrainer:
|
||||||
|
|
||||||
def _close(self):
|
def _close(self):
|
||||||
self.env.close()
|
self.env.close()
|
||||||
self.writer.close()
|
|
||||||
|
|
||||||
def _save_model(self, model_path: str, log: bool = True):
|
def _save_model(self, model_path: str):
|
||||||
if log:
|
self.logger.info("[SAVE]: Saving the final model...")
|
||||||
print(f"[SAVE]: Saving the model to: {model_path}...")
|
|
||||||
|
|
||||||
with open(model_path, "wb") as f:
|
params = [
|
||||||
f.write(
|
vars(self.args),
|
||||||
flax.serialization.to_bytes(
|
[
|
||||||
[
|
self.agent_state.params["sensor_params"],
|
||||||
vars(self.args),
|
self.agent_state.params["actor_params"],
|
||||||
[
|
self.agent_state.params["critic_params"],
|
||||||
self.agent_state.params["sensor_params"],
|
self.agent_state.params["feature_extractor_params"],
|
||||||
self.agent_state.params["actor_params"],
|
],
|
||||||
self.agent_state.params["critic_params"],
|
]
|
||||||
self.agent_state.params["feature_extractor_params"],
|
self.logger.save_final_model(params=params)
|
||||||
],
|
|
||||||
]
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
def train(self, log: bool = True):
|
def train(self):
|
||||||
"""
|
"""
|
||||||
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:
|
self.logger.info(f"running name: {self.run_name}")
|
||||||
print(f"running name: {self.run_name}")
|
|
||||||
|
|
||||||
is_tty = sys.stdout.isatty()
|
is_tty = sys.stdout.isatty()
|
||||||
if log:
|
self.logger.info("[TRAIN]: Resetting environment...")
|
||||||
print("[TRAIN]: Resetting environment...")
|
|
||||||
|
|
||||||
if not is_tty:
|
if not is_tty:
|
||||||
print(f">>> [HPC] Initial reset started: {time.ctime()}", flush=True)
|
self.logger.info(f">>> [HPC] Initial reset started: {time.ctime()}")
|
||||||
|
|
||||||
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:
|
if not is_tty:
|
||||||
print(f">>> [HPC] Initial reset completed: {time.ctime()}", flush=True)
|
self.logger.info(f">>> [HPC] Initial reset completed: {time.ctime()}")
|
||||||
|
|
||||||
global_step = 0
|
global_step = 0
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
|
|
||||||
if self.args.track:
|
|
||||||
import wandb
|
|
||||||
|
|
||||||
if log:
|
|
||||||
print("[TRAIN]: Initializing Weights and Biases...")
|
|
||||||
|
|
||||||
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,
|
|
||||||
)
|
|
||||||
|
|
||||||
if log:
|
|
||||||
print("[TRAIN]: Adding hyperparameters to TensorBoard...")
|
|
||||||
|
|
||||||
self.writer.add_text(
|
|
||||||
"hyperparameters",
|
|
||||||
"|param|value|\n|---|---|\n"
|
|
||||||
+ "\n".join(f"|{k}|{v}|" for k, v in vars(self.args).items()),
|
|
||||||
)
|
|
||||||
|
|
||||||
iter_bar = tqdm.tqdm(
|
iter_bar = tqdm.tqdm(
|
||||||
range(1, self.args.num_iterations + 1),
|
range(1, self.args.num_iterations + 1),
|
||||||
disable=not is_tty,
|
disable=not is_tty,
|
||||||
|
|
@ -514,19 +469,18 @@ class PPOTrainer:
|
||||||
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 log and not is_tty:
|
if not is_tty:
|
||||||
sps = int(global_step / (time.time() - start_time))
|
sps = int(global_step / (time.time() - start_time))
|
||||||
remaining_steps = self.args.total_timesteps - global_step
|
remaining_steps = self.args.total_timesteps - global_step
|
||||||
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
|
eta_seconds = int(remaining_steps / sps) if sps > 0 else 0
|
||||||
eta_str = str(datetime.timedelta(seconds=eta_seconds))
|
eta_str = str(datetime.timedelta(seconds=eta_seconds))
|
||||||
|
|
||||||
print(
|
self.logger.info(
|
||||||
f"Iteration {iteration}/{self.args.num_iterations} | "
|
f"Iteration {iteration}/{self.args.num_iterations} | "
|
||||||
f"Step {global_step}/{self.args.total_timesteps} | "
|
f"Step {global_step}/{self.args.total_timesteps} | "
|
||||||
f"SPS {sps} | "
|
f"SPS {sps} | "
|
||||||
f"Return {loss_info.avg_episodic_return:.4f} | "
|
f"Return {loss_info.avg_episodic_return:.4f} | "
|
||||||
f"ETA {eta_str}",
|
f"ETA {eta_str}"
|
||||||
flush=True,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.args.save_model:
|
if self.args.save_model:
|
||||||
|
|
|
||||||
Reference in a new issue