feat: checkpoints
This commit is contained in:
parent
3efee5d746
commit
8b6dcbae7c
2 changed files with 32 additions and 16 deletions
|
|
@ -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
|
||||
|
|
|
|||
21
src/brittle_star_project/trainers/utils.py
Normal file
21
src/brittle_star_project/trainers/utils.py
Normal file
|
|
@ -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
|
||||
Reference in a new issue