| 40 | |
| 41 | |
| 42 | class EarlyStopping: |
| 43 | def __init__(self, patience=7, verbose=False, delta=0): |
| 44 | self.patience = patience |
| 45 | self.verbose = verbose |
| 46 | self.counter = 0 |
| 47 | self.best_score = None |
| 48 | self.early_stop = False |
| 49 | self.val_loss_min = np.Inf |
| 50 | self.delta = delta |
| 51 | |
| 52 | def __call__(self, val_loss, model, path): |
| 53 | score = -val_loss |
| 54 | if self.best_score is None: |
| 55 | self.best_score = score |
| 56 | self.save_checkpoint(val_loss, model, path) |
| 57 | elif score < self.best_score + self.delta: |
| 58 | self.counter += 1 |
| 59 | print(f'EarlyStopping counter: {self.counter} out of {self.patience}') |
| 60 | if self.counter >= self.patience: |
| 61 | self.early_stop = True |
| 62 | else: |
| 63 | self.best_score = score |
| 64 | self.save_checkpoint(val_loss, model, path) |
| 65 | self.counter = 0 |
| 66 | |
| 67 | def save_checkpoint(self, val_loss, model, path): |
| 68 | if self.verbose: |
| 69 | print(f'Validation loss decreased ({self.val_loss_min:.6f} --> {val_loss:.6f}). Saving model ...') |
| 70 | torch.save(model.state_dict(), path + '/' + 'checkpoint.pth') |
| 71 | self.val_loss_min = val_loss |
| 72 | |
| 73 | |
| 74 | class dotdict(dict): |