MCPcopy Create free account
hub / github.com/Anoise/WTFlib / EarlyStopping

Class EarlyStopping

LDPS_Graph/utils/tools.py:42–71  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

40
41
42class 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
74class dotdict(dict):

Callers 1

trainMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected