Spaces:
Sleeping
Sleeping
File size: 1,409 Bytes
bda104d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 | from lightning.pytorch.callbacks import Callback
from lightning.pytorch import Trainer
from lightning.pytorch.core import LightningModule
import torch
class BaseCallback(Callback):
def __init__(self, every_n_steps = 1, every_n_epochs = 1, **kwargs):
super().__init__(**kwargs)
self.every_n_steps = every_n_steps
self.every_n_epochs = every_n_epochs
def _check_step(self, trainer: Trainer, pl_module: LightningModule) -> bool:
if self.every_n_steps is not None:
return trainer.global_step % self.every_n_steps == 0
else:
return False
def _check_epoch(self, trainer: Trainer, pl_module: LightningModule) -> bool:
if self.every_n_epochs is not None:
return trainer.current_epoch % self.every_n_epochs == 0
else:
return False
def _should_run(self, trainer: Trainer, pl_module: LightningModule) -> bool:
return self._check_step(trainer, pl_module) or self._check_epoch(trainer, pl_module)
def _should_run_on_validation(self, trainer: Trainer, pl_module: LightningModule) -> bool:
return self._check_step(trainer, pl_module) or self._check_epoch(trainer, pl_module)
def _should_run_on_test(self, trainer: Trainer, pl_module: LightningModule) -> bool:
return self._check_step(trainer, pl_module) or self._check_epoch(trainer, pl_module)
|