mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
@@ -57,6 +57,7 @@ class EarlyStopping(Callback):
|
||||
self.min_delta = min_delta
|
||||
self.wait = 0
|
||||
self.stopped_epoch = 0
|
||||
self.mode = mode
|
||||
|
||||
mode_dict = {
|
||||
'min': torch.lt,
|
||||
@@ -67,9 +68,8 @@ class EarlyStopping(Callback):
|
||||
if mode not in mode_dict:
|
||||
if self.verbose > 0:
|
||||
log.info(f'EarlyStopping mode {mode} is unknown, fallback to auto mode.')
|
||||
mode = 'auto'
|
||||
self.mode = 'auto'
|
||||
|
||||
self.monitor_op = mode_dict[mode]
|
||||
self.min_delta *= 1 if self.monitor_op == torch.gt else -1
|
||||
|
||||
def _validate_condition_metric(self, logs):
|
||||
@@ -94,6 +94,15 @@ class EarlyStopping(Callback):
|
||||
|
||||
return True
|
||||
|
||||
@property
|
||||
def monitor_op(self):
|
||||
mode_dict = {
|
||||
'min': torch.lt,
|
||||
'max': torch.gt,
|
||||
'auto': torch.gt if 'acc' in self.monitor else torch.lt
|
||||
}
|
||||
return mode_dict[self.mode]
|
||||
|
||||
def on_train_start(self, trainer, pl_module):
|
||||
# Allow instances to be re-used
|
||||
self.wait = 0
|
||||
|
||||
@@ -240,6 +240,15 @@ def test_early_stopping_no_val_step(tmpdir):
|
||||
assert trainer.current_epoch < trainer.max_epochs
|
||||
|
||||
|
||||
def test_pickling(tmpdir):
|
||||
import pickle
|
||||
early_stopping = EarlyStopping()
|
||||
ckpt = ModelCheckpoint(tmpdir)
|
||||
|
||||
pickle.dumps(ckpt)
|
||||
pickle.dumps(early_stopping)
|
||||
|
||||
|
||||
def test_model_checkpoint_with_non_string_input(tmpdir):
|
||||
""" Test that None in checkpoint callback is valid and that chkp_path is
|
||||
set correctly """
|
||||
|
||||
@@ -88,6 +88,9 @@ def test_loggers_pickle(tmpdir, monkeypatch, logger_class):
|
||||
logger_args = _get_logger_args(logger_class, tmpdir)
|
||||
logger = logger_class(**logger_args)
|
||||
|
||||
# test pickling loggers
|
||||
pickle.dumps(logger)
|
||||
|
||||
trainer = Trainer(
|
||||
max_epochs=1,
|
||||
logger=logger
|
||||
|
||||
@@ -26,6 +26,13 @@ from tests.base import (
|
||||
)
|
||||
|
||||
|
||||
def test_model_pickle(tmpdir):
|
||||
import pickle
|
||||
|
||||
model = TestModelBase(tutils.get_default_hparams())
|
||||
pickle.dumps(model)
|
||||
|
||||
|
||||
def test_hparams_save_load(tmpdir):
|
||||
model = DictHparamsModel({'in_features': 28 * 28, 'out_features': 10, 'failed_key': lambda x: x})
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import inspect
|
||||
from argparse import ArgumentParser, Namespace
|
||||
from unittest import mock
|
||||
import pickle
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -42,14 +43,14 @@ def test_add_argparse_args_redefined(cli_args):
|
||||
|
||||
args = parser.parse_args(cli_args)
|
||||
|
||||
# make sure we can pickle args
|
||||
pickle.dumps(args)
|
||||
|
||||
# Check few deprecated args are not in namespace:
|
||||
for depr_name in ('gradient_clip', 'nb_gpu_nodes', 'max_nb_epochs'):
|
||||
assert depr_name not in args
|
||||
|
||||
trainer = Trainer.from_argparse_args(args=args)
|
||||
|
||||
# make sure trainer can be pickled
|
||||
import pickle
|
||||
pickle.dumps(trainer)
|
||||
|
||||
assert isinstance(trainer, Trainer)
|
||||
|
||||
Reference in New Issue
Block a user