1
Fork 0

feat(log): non-interactive logging, incl. progress bar

This commit is contained in:
Tibo De Peuter 2026-04-08 19:54:05 +02:00
parent faf31567e8
commit 46ae13abc3
Signed by: tdpeuter
SSH key fingerprint: SHA256:u/h/LVoqKF1Iz02uOyxe6hcjmoZASCGV2HM0TG9ZMoU
2 changed files with 38 additions and 34 deletions

View file

@ -1,6 +1,5 @@
import datetime 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
@ -11,7 +10,6 @@ import jax
import jax.numpy as jnp import jax.numpy as jnp
import numpy as np import numpy as np
import optax import optax
import tqdm
from flax.training.train_state import TrainState from flax.training.train_state import TrainState
from experiment_logger import get_logger from experiment_logger import get_logger
@ -367,9 +365,9 @@ class PPOTrainer:
} }
self.logger.log(metrics, step=global_step) self.logger.log(metrics, step=global_step)
def _step(self, env_state, next_obs, next_done, is_tty: bool, iteration: int) -> tuple: def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple:
if not is_tty and iteration == 1: if iteration == 1:
self.logger.info(f">>> [HPC] Starting first rollout (JIT): {time.ctime()}") self.logger.log_non_interactive(f"Starting first rollout (JIT): {time.ctime()}")
( (
self.agent_state, self.agent_state,
@ -381,20 +379,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 not is_tty and iteration == 1: if iteration == 1:
self.logger.info(f">>> [HPC] First rollout completed: {time.ctime()}") self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}")
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: if iteration == 1:
self.logger.info(f">>> [HPC] Starting first PPO update (JIT): {time.ctime()}") self.logger.log_non_interactive(f"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 not is_tty and iteration == 1: if iteration == 1:
self.logger.info(f">>> [HPC] First PPO update completed: {time.ctime()}") self.logger.log_non_interactive(f"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()
@ -439,49 +437,41 @@ class PPOTrainer:
""" """
self.logger.info(f"running name: {self.run_name}") self.logger.info(f"running name: {self.run_name}")
is_tty = sys.stdout.isatty()
self.logger.info("[TRAIN]: Resetting environment...") self.logger.info("[TRAIN]: Resetting environment...")
self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}")
if not is_tty:
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 not is_tty: self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}")
self.logger.info(f">>> [HPC] Initial reset completed: {time.ctime()}")
global_step = 0 global_step = 0
start_time = time.time() start_time = time.time()
iter_bar = tqdm.tqdm( iter_bar = self.logger.progress_bar(range(1, self.args.num_iterations + 1))
range(1, self.args.num_iterations + 1),
disable=not is_tty,
)
for iteration in iter_bar: 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, loss_info = self._step(
env_state, next_obs, next_done, is_tty=is_tty, iteration=iteration env_state, next_obs, next_done, iteration=iteration
) )
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))
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))
self.logger.info( self.logger.log_non_interactive(
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}"
) )
if self.args.save_model: if self.args.save_model:
model_path = f"{self.run_dir}/{self.args.exp_name}.cleanrl_model" model_path = f"{self.run_dir}/{self.args.exp_name}.cleanrl_model"

View file

@ -10,6 +10,7 @@ import datetime
import json import json
import logging import logging
import subprocess import subprocess
import sys
import time import time
from pathlib import Path from pathlib import Path
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
@ -86,6 +87,7 @@ class UnifiedLogger:
self.use_wandb = use_wandb self.use_wandb = use_wandb
self.wandb_available = False self.wandb_available = False
self.wandb_run = None self.wandb_run = None
self.is_interactive = sys.stdout.isatty()
# Setup local storage # Setup local storage
self.run_dir = Path(base_dir) / run_name self.run_dir = Path(base_dir) / run_name
@ -150,6 +152,18 @@ class UnifiedLogger:
"""Dynamically update the verbosity of the stdout/text logger.""" """Dynamically update the verbosity of the stdout/text logger."""
self._text_logger.setLevel(level) self._text_logger.setLevel(level)
def log_non_interactive(self, msg: str, *args, **kwargs):
"""Log an info message only if running in a non-interactive environment."""
if not self.is_interactive:
self.info(msg, *args, **kwargs)
def progress_bar(self, iterable=None, *args, **kwargs):
"""Wrapper around tqdm that automatically disables in non-interactive environments."""
import tqdm
kwargs.setdefault("disable", not self.is_interactive)
return tqdm.tqdm(iterable, *args, **kwargs)
def info(self, msg: str, *args, **kwargs): def info(self, msg: str, *args, **kwargs):
"""Log an info message to stdout and disk.""" """Log an info message to stdout and disk."""
self._text_logger.info(msg, *args, **kwargs) self._text_logger.info(msg, *args, **kwargs)