From 34237cfcaf5c53a15e62a3219d72bf33b667d214 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adrian=20W=C3=A4lchli?= Date: Mon, 25 May 2020 22:01:29 +0200 Subject: [PATCH] handle unknown args passed to Trainer.from_argparse_args (#1932) * filter valid args * error on unknown manual args * added test * changelog * update docs and doctest * simplify * doctest * doctest * doctest * better test with mock check for init call * fstring * extend test * skip test on 3.6 not working Co-authored-by: William Falcon --- CHANGELOG.md | 3 +++ pytorch_lightning/trainer/trainer.py | 22 +++++++++++++++++----- tests/trainer/test_trainer_cli.py | 26 ++++++++++++++++++++++++++ 3 files changed, 46 insertions(+), 5 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index f87b746a..b6280a1b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -38,8 +38,11 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). - Fixed user warning when apex was used together with learning rate schedulers ([#1873](https://github.com/PyTorchLightning/pytorch-lightning/pull/1873)) +- Fixed an issue with `Trainer.from_argparse_args` when passing in unknown Trainer args ([#1932](https://github.com/PyTorchLightning/pytorch-lightning/pull/1932)) + - Fix bug related to logger not being reset correctly for model after tuner algorithms ([#1933](https://github.com/PyTorchLightning/pytorch-lightning/pull/1933)) + ## [0.7.6] - 2020-05-16 ### Added diff --git a/pytorch_lightning/trainer/trainer.py b/pytorch_lightning/trainer/trainer.py index eefbfe1a..c373bf41 100644 --- a/pytorch_lightning/trainer/trainer.py +++ b/pytorch_lightning/trainer/trainer.py @@ -725,20 +725,32 @@ class Trainer( @classmethod def from_argparse_args(cls, args: Union[Namespace, ArgumentParser], **kwargs) -> 'Trainer': - """create an instance from CLI arguments + """ + Create an instance from CLI arguments. + + Args: + args: The parser or namespace to take arguments from. Only known arguments will be + parsed and passed to the :class:`Trainer`. + **kwargs: Additional keyword arguments that may override ones in the parser or namespace. + These must be valid Trainer arguments. Example: >>> parser = ArgumentParser(add_help=False) >>> parser = Trainer.add_argparse_args(parser) + >>> parser.add_argument('--my_custom_arg', default='something') # doctest: +SKIP >>> args = Trainer.parse_argparser(parser.parse_args("")) - >>> trainer = Trainer.from_argparse_args(args) + >>> trainer = Trainer.from_argparse_args(args, logger=False) """ if isinstance(args, ArgumentParser): - args = Trainer.parse_argparser(args) + args = cls.parse_argparser(args) params = vars(args) - params.update(**kwargs) - return cls(**params) + # we only want to pass in valid Trainer args, the rest may be user specific + valid_kwargs = inspect.signature(cls.__init__).parameters + trainer_kwargs = dict((name, params[name]) for name in valid_kwargs if name in params) + trainer_kwargs.update(**kwargs) + + return cls(**trainer_kwargs) @property def num_gpus(self) -> int: diff --git a/tests/trainer/test_trainer_cli.py b/tests/trainer/test_trainer_cli.py index 922acf93..c66d6149 100644 --- a/tests/trainer/test_trainer_cli.py +++ b/tests/trainer/test_trainer_cli.py @@ -1,5 +1,6 @@ import inspect import pickle +import sys from argparse import ArgumentParser, Namespace from unittest import mock @@ -110,3 +111,28 @@ def test_argparse_args_parsing(cli_args, expected): for k, v in expected.items(): assert getattr(args, k) == v assert Trainer.from_argparse_args(args) + + +@pytest.mark.skipif( + sys.version_info < (3, 7), + reason="signature inspection while mocking is not working in Python < 3.7 despite autospec" +) +@pytest.mark.parametrize(['cli_args', 'extra_args'], [ + pytest.param({}, {}), + pytest.param({'logger': False}, {}), + pytest.param({'logger': False}, {'logger': True}), + pytest.param({'logger': False}, {'checkpoint_callback': True}), +]) +def test_init_from_argparse_args(cli_args, extra_args): + unknown_args = dict(unknown_arg=0) + + # unkown args in the argparser/namespace should be ignored + with mock.patch('pytorch_lightning.Trainer.__init__', autospec=True, return_value=None) as init: + trainer = Trainer.from_argparse_args(Namespace(**cli_args, **unknown_args), **extra_args) + expected = dict(cli_args) + expected.update(extra_args) # extra args should override any cli arg + init.assert_called_with(trainer, **expected) + + # passing in unknown manual args should throw an error + with pytest.raises(TypeError, match=r"__init__\(\) got an unexpected keyword argument 'unknown_arg'"): + Trainer.from_argparse_args(Namespace(**cli_args), **extra_args, **unknown_args)