jone/Music_Source_Separation
3
1import logging2import os3from typing import NoReturn4 5import pytorch_lightning as pl6import torch7import torch.nn as nn8from pytorch_lightning.utilities import rank_zero_only9 10 11class SaveCheckpointsCallback(pl.Callback):12 def __init__(13 self,14 model: nn.Module,15 checkpoints_dir: str,16 save_step_frequency: int,17 ):18 r"""Callback to save checkpoints every #save_step_frequency steps.19 20 Args:21 model: nn.Module22 checkpoints_dir: str, directory to save checkpoints23 save_step_frequency: int24 """25 self.model = model26 self.checkpoints_dir = checkpoints_dir27 self.save_step_frequency = save_step_frequency28 os.makedirs(self.checkpoints_dir, exist_ok=True)29 30 @rank_zero_only31 def on_batch_end(self, trainer: pl.Trainer, _) -> NoReturn:32 r"""Save checkpoint."""33 global_step = trainer.global_step34 35 if global_step % self.save_step_frequency == 0:36 37 checkpoint_path = os.path.join(38 self.checkpoints_dir, "step={}.pth".format(global_step)39 )40 41 checkpoint = {'step': global_step, 'model': self.model.state_dict()}42 43 torch.save(checkpoint, checkpoint_path)44 logging.info("Save checkpoint to {}".format(checkpoint_path))45 