From 8b6dcbae7c899a02919f288341a224f2813f6d2a Mon Sep 17 00:00:00 2001 From: Tibo De Peuter Date: Thu, 16 Apr 2026 12:27:56 +0200 Subject: [PATCH] feat: checkpoints --- .../trainers/PPOTrainer.py | 27 ++++++++----------- src/brittle_star_project/trainers/utils.py | 21 +++++++++++++++ 2 files changed, 32 insertions(+), 16 deletions(-) create mode 100644 src/brittle_star_project/trainers/utils.py diff --git a/src/brittle_star_project/trainers/PPOTrainer.py b/src/brittle_star_project/trainers/PPOTrainer.py index e5ecca6..ea77f12 100644 --- a/src/brittle_star_project/trainers/PPOTrainer.py +++ b/src/brittle_star_project/trainers/PPOTrainer.py @@ -24,6 +24,7 @@ from brittle_star_project.MLPs.mlps import ( Storage, ) from brittle_star_project.ppo import PPO +from brittle_star_project.trainers.utils import serialize_training_state # TODO: move to config _ALLOWED_OBS_KEYS = { @@ -573,23 +574,13 @@ class PPOTrainer: def _save_model(self, model_path: str): self.logger.info("[SAVE]: Saving the final model...") + config_dict, params = serialize_training_state(self.cfg, self.agent_state) + self.logger.save_final_model(params=params, metadata=config_dict) - from dataclasses import asdict as _asdict - - config_dict = { - "experiment": _asdict(self.experiment), - "ppo": _asdict(self.ppo), - } - params = [ - config_dict, - [ - self.agent_state.params["sensor_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 _save_checkpoint(self, iteration: int): + self.logger.info(f"[SAVE]: Saving checkpoint at iteration {iteration}...") + config_dict, params = serialize_training_state(self.cfg, self.agent_state) + self.logger.save_checkpoint(params=params, step=iteration, metadata=config_dict) def train(self): """ @@ -647,6 +638,10 @@ class PPOTrainer: f"ETA {eta_str}" ) + if self.logging_cfg.save_model and self.logging_cfg.checkpoint_frequency > 0: + if iteration % self.logging_cfg.checkpoint_frequency == 0: + self._save_checkpoint(iteration) + if getattr(self.cfg.experiment, "debug_sanity", False): self.logger.info("\n[SANITY CHECK] Successfully completed 1 epoch") break diff --git a/src/brittle_star_project/trainers/utils.py b/src/brittle_star_project/trainers/utils.py new file mode 100644 index 0000000..e271366 --- /dev/null +++ b/src/brittle_star_project/trainers/utils.py @@ -0,0 +1,21 @@ +from flax.training.train_state import TrainState +from brittle_star_project.configs.main_config import BrittleStarConfig + + +def serialize_training_state(cfg: BrittleStarConfig, agent_state: TrainState): + from dataclasses import asdict as _asdict + + config_dict = { + "experiment": _asdict(cfg.experiment), + "ppo": _asdict(cfg.ppo), + } + params = [ + config_dict, + [ + agent_state.params["sensor_params"], + agent_state.params["actor_params"], + agent_state.params["critic_params"], + agent_state.params["feature_extractor_params"], + ], + ] + return config_dict, params