From 8a2f8b14bcc4d5c4fdfbaeecdc56732ee3ce6a46 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Mon, 6 Apr 2026 14:07:46 +0200 Subject: [PATCH] fix(PPOTrainer.py, PPOArgs.py): added hyperparam config path to cli args, fixed missing arguments error in _step call --- experiments/PPOTrainer.py | 6 ++++-- experiments/train.py | 12 +++++++++--- src/brittle_star_project/dataclasses/PPOArgs.py | 3 +++ 3 files changed, 16 insertions(+), 5 deletions(-) diff --git a/experiments/PPOTrainer.py b/experiments/PPOTrainer.py index cfadb54..d947fb7 100644 --- a/experiments/PPOTrainer.py +++ b/experiments/PPOTrainer.py @@ -500,12 +500,14 @@ class PPOTrainer: iter_bar = tqdm.tqdm( range(1, self.args.num_iterations + 1), - disable=not sys.stdout.isatty(), + disable=not is_tty, ) for iteration in iter_bar: 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, is_tty=is_tty, iteration=iteration + ) if not is_tty and iteration == 1: print(f">>> [HPC] First rollout completed: {time.ctime()}", flush=True) diff --git a/experiments/train.py b/experiments/train.py index 596dad9..0a65493 100644 --- a/experiments/train.py +++ b/experiments/train.py @@ -17,11 +17,14 @@ def make_env(config_path: str | None, num_envs: int) -> BrittleStarJaxEnvWrapper return BrittleStarJaxEnvWrapper.from_config(config_path, num_envs=num_envs) -def parse_args() -> PPOArgs: +def parse_args(log: bool = True) -> 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: + if temp_args.hyperparameter_config_path is not None: + if log: + print(f"Loading hyperparameter config from {temp_args.hyperparameter_config_path}") + + with open(temp_args.hyperparameter_config_path, "r") as f: config = yaml.safe_load(f) if config: # parse PPOArgs with defaults from yaml. @@ -32,6 +35,9 @@ def parse_args() -> PPOArgs: # Reparse CLI to ensure they OVERRIDE the yaml args = tyro.cli(PPOArgs, default=temp_args) else: + if log: + print("No hyperparameter config provided, using default config") + args = temp_args return args diff --git a/src/brittle_star_project/dataclasses/PPOArgs.py b/src/brittle_star_project/dataclasses/PPOArgs.py index f6b20bf..3f04c16 100644 --- a/src/brittle_star_project/dataclasses/PPOArgs.py +++ b/src/brittle_star_project/dataclasses/PPOArgs.py @@ -13,6 +13,9 @@ class PPOArgs: # path to environment config file, if None, use default config env_config_path: str | None = None + # path to hyperparameter config file (yaml), if None, use default config + hyperparameter_config_path: str | None = None + # the name of this experiment exp_name: str = "brittle_star_ppo"