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,
|
Storage,
|
||||||
)
|
)
|
||||||
from brittle_star_project.ppo import PPO
|
from brittle_star_project.ppo import PPO
|
||||||
|
from brittle_star_project.trainers.utils import serialize_training_state
|
||||||
|
|
||||||
# TODO: move to config
|
# TODO: move to config
|
||||||
_ALLOWED_OBS_KEYS = {
|
_ALLOWED_OBS_KEYS = {
|
||||||
|
|
@ -573,23 +574,13 @@ class PPOTrainer:
|
||||||
|
|
||||||
def _save_model(self, model_path: str):
|
def _save_model(self, model_path: str):
|
||||||
self.logger.info("[SAVE]: Saving the final model...")
|
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
|
def _save_checkpoint(self, iteration: int):
|
||||||
|
self.logger.info(f"[SAVE]: Saving checkpoint at iteration {iteration}...")
|
||||||
config_dict = {
|
config_dict, params = serialize_training_state(self.cfg, self.agent_state)
|
||||||
"experiment": _asdict(self.experiment),
|
self.logger.save_checkpoint(params=params, step=iteration, metadata=config_dict)
|
||||||
"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 train(self):
|
def train(self):
|
||||||
"""
|
"""
|
||||||
|
|
@ -647,6 +638,10 @@ class PPOTrainer:
|
||||||
f"ETA {eta_str}"
|
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):
|
if getattr(self.cfg.experiment, "debug_sanity", False):
|
||||||
self.logger.info("\n[SANITY CHECK] Successfully completed 1 epoch")
|
self.logger.info("\n[SANITY CHECK] Successfully completed 1 epoch")
|
||||||
break
|
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