hparams as dict [blocked by 1041] (#1029)

* hparams as dict

* hparams as dict

* fixing

* fixing

* fixing

* fixing

* typing

* typing

* chnagelog

* update set hparams

* use setter

* simplify

* chnagelog

* imports

* pylint

* typing

* Update training_io.py

* Update training_io.py

* Update lightning.py

* Update test_trainer.py

* Update __init__.py

* Update base.py

* Update utils.py

* Update test_trainer.py

* Update training_io.py

* Update test_trainer.py

* Update test_trainer.py

* Update test_trainer.py

* Update test_trainer.py

* Update callback_config.py

* Update callback_config.py

* Update test_trainer.py

Co-authored-by: William Falcon <waf2107@columbia.edu>
This commit is contained in:
Jirka Borovec
2020-03-04 09:33:39 -05:00
committed by GitHub
co-authored by William Falcon
parent 6a39573267
commit e586ed4767
18 changed files with 168 additions and 87 deletions
+3
View File
@@ -24,6 +24,7 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/).
- Split callbacks in multiple files ([#849](https://github.com/PyTorchLightning/pytorch-lightning/pull/849)) - Split callbacks in multiple files ([#849](https://github.com/PyTorchLightning/pytorch-lightning/pull/849))
- Support for user defined callbacks ([#889](https://github.com/PyTorchLightning/pytorch-lightning/pull/889) and [#950](https://github.com/PyTorchLightning/pytorch-lightning/pull/950)) - Support for user defined callbacks ([#889](https://github.com/PyTorchLightning/pytorch-lightning/pull/889) and [#950](https://github.com/PyTorchLightning/pytorch-lightning/pull/950))
- Added support for multiple loggers to be passed to `Trainer` as an iterable (e.g. list, tuple, etc.) ([#903](https://github.com/PyTorchLightning/pytorch-lightning/pull/903)) - Added support for multiple loggers to be passed to `Trainer` as an iterable (e.g. list, tuple, etc.) ([#903](https://github.com/PyTorchLightning/pytorch-lightning/pull/903))
- Added support for logging hparams as dict ([#1029](https://github.com/PyTorchLightning/pytorch-lightning/pull/1029))
### Changed ### Changed
@@ -32,6 +33,8 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/).
- Changed `pytorch_lightning.logging` to `pytorch_lightning.loggers` ([#767](https://github.com/PyTorchLightning/pytorch-lightning/pull/767)) - Changed `pytorch_lightning.logging` to `pytorch_lightning.loggers` ([#767](https://github.com/PyTorchLightning/pytorch-lightning/pull/767))
- Moved the default `tqdm_dict` definition from Trainer to `LightningModule`, so it can be overridden by the user ([#749](https://github.com/PyTorchLightning/pytorch-lightning/pull/749)) - Moved the default `tqdm_dict` definition from Trainer to `LightningModule`, so it can be overridden by the user ([#749](https://github.com/PyTorchLightning/pytorch-lightning/pull/749))
- Moved functionality of `LightningModule.load_from_metrics` into `LightningModule.load_from_checkpoint` ([#995](https://github.com/PyTorchLightning/pytorch-lightning/pull/995)) - Moved functionality of `LightningModule.load_from_metrics` into `LightningModule.load_from_checkpoint` ([#995](https://github.com/PyTorchLightning/pytorch-lightning/pull/995))
- Changed Checkpoint path parameter from `filepath` to `dirpath` ([#1016](https://github.com/PyTorchLightning/pytorch-lightning/pull/1016))
- Freezed models `hparams` as `Namespace` property ([#1029](https://github.com/PyTorchLightning/pytorch-lightning/pull/1029))
### Deprecated ### Deprecated
+1 -1
View File
@@ -62,7 +62,7 @@ The Lightning checkpoint also saves the hparams (hyperparams) passed into the Li
from argparse import Namespace from argparse import Namespace
# usually these come from command line args # usually these come from command line args
args = Namespace(**{'learning_rate':0.001}) args = Namespace(learning_rate=0.001)
# define you module to have hparams as the first arg # define you module to have hparams as the first arg
# this means your checkpoint will have everything that went into making # this means your checkpoint will have everything that went into making
@@ -27,9 +27,11 @@ class ModelCheckpoint(Callback):
# save epoch and val_loss in name # save epoch and val_loss in name
ModelCheckpoint(filepath='{epoch:02d}-{val_loss:.2f}.hdf5') ModelCheckpoint(filepath='{epoch:02d}-{val_loss:.2f}.hdf5')
# saves file like: /my/path/here/sample-mnist_epoch=02_val_loss=0.32.ckpt # saves file like: /my/path/here/sample-mnist_epoch=02_val_loss=0.32.ckpt
# if model already exits, the file will be: /my/path/here/sample-mnist-v0_epoch=02_val_loss=0.32.ckpt # if model already exits, the file will be: /my/path/here/sample-mnist-v0_epoch=02_val_loss=0.32.ckpt
monitor: quantity to monitor. monitor: quantity to monitor.
verbose: verbosity mode, False or True. verbose: verbosity mode, False or True.
save_top_k: if `save_top_k == k`, save_top_k: if `save_top_k == k`,
@@ -135,7 +137,7 @@ class ModelCheckpoint(Callback):
if self.save_function is not None: if self.save_function is not None:
self.save_function(filepath) self.save_function(filepath)
else: else:
raise ValueError(".save_function() not set") raise ValueError("Method `.save_function()` not set")
def check_monitor_top_k(self, current: float) -> bool: def check_monitor_top_k(self, current: float) -> bool:
less_than_k_models = len(self.best_k_models) < self.save_top_k less_than_k_models = len(self.best_k_models) < self.save_top_k
+17 -2
View File
@@ -5,7 +5,7 @@ import os
import warnings import warnings
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from argparse import Namespace from argparse import Namespace
from typing import Optional, Union, Dict, Callable from typing import Any, Callable, Dict, Optional, Union
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -68,6 +68,20 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
#: True if using amp #: True if using amp
self.use_amp = False self.use_amp = False
@property
def hparams(self) -> Namespace:
if not hasattr(self, '_hparams'):
return Namespace()
assert isinstance(self._hparams, dict)
return Namespace(**self._hparams)
@hparams.setter
def hparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
"""Set the model hyper-parameters."""
if isinstance(params, Namespace):
params = vars(params)
self._hparams = params
def print(self, *args, **kwargs): def print(self, *args, **kwargs):
r""" r"""
Prints only from process 0. Use this in any distributed mode to log only once Prints only from process 0. Use this in any distributed mode to log only once
@@ -1201,7 +1215,8 @@ class LightningModule(ABC, GradInformation, ModelIO, ModelHooks):
if cls_takes_hparams: if cls_takes_hparams:
if ckpt_hparams is not None: if ckpt_hparams is not None:
hparams = Namespace(**ckpt_hparams) is_namespace = checkpoint.get('hparams_type') == 'namespace'
hparams = Namespace(**ckpt_hparams) if is_namespace else ckpt_hparams
else: else:
warnings.warn( warnings.warn(
f"Checkpoint does not contain hyperparameters but {cls.__name__}'s __init__ contains" f"Checkpoint does not contain hyperparameters but {cls.__name__}'s __init__ contains"
+2 -4
View File
@@ -36,16 +36,14 @@ class ModelIO(object):
""" """
def load_hparams_from_tags_csv(tags_csv): def load_hparams_from_tags_csv(tags_csv) -> Namespace:
if not os.path.isfile(tags_csv): if not os.path.isfile(tags_csv):
log.warning(f'Missing Tags: {tags_csv}.') log.warning(f'Missing Tags: {tags_csv}.')
return Namespace() return Namespace()
tags = {}
with open(tags_csv) as f: with open(tags_csv) as f:
csv_reader = csv.reader(f, delimiter=',') csv_reader = csv.reader(f, delimiter=',')
for row in list(csv_reader)[1:]: tags = {row[0]: convert(row[1]) for row in list(csv_reader)[1:]}
tags[row[0]] = convert(row[1])
ns = Namespace(**tags) ns = Namespace(**tags)
return ns return ns
+17 -10
View File
@@ -1,5 +1,6 @@
import argparse import argparse
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from argparse import Namespace
from functools import wraps from functools import wraps
from typing import Union, Optional, Dict, Iterable, Any, Callable, List from typing import Union, Optional, Dict, Iterable, Any, Callable, List
@@ -41,6 +42,12 @@ class LightningLoggerBase(ABC):
""" """
pass pass
def _convert_params(self, params: Union[Dict[str, Any], Namespace]) -> Dict[str, Any]:
# in case converting from namespace
if isinstance(params, Namespace):
params = vars(params)
return params
@abstractmethod @abstractmethod
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: argparse.Namespace):
"""Record hyperparameters. """Record hyperparameters.
@@ -50,11 +57,11 @@ class LightningLoggerBase(ABC):
""" """
pass pass
def save(self): def save(self) -> None:
"""Save log data.""" """Save log data."""
pass pass
def finalize(self, status: str): def finalize(self, status: str) -> None:
"""Do any processing that is necessary to finalize an experiment. """Do any processing that is necessary to finalize an experiment.
Args: Args:
@@ -62,7 +69,7 @@ class LightningLoggerBase(ABC):
""" """
pass pass
def close(self): def close(self) -> None:
"""Do any cleanup that is necessary to close an experiment.""" """Do any cleanup that is necessary to close an experiment."""
pass pass
@@ -72,7 +79,7 @@ class LightningLoggerBase(ABC):
return self._rank return self._rank
@rank.setter @rank.setter
def rank(self, value: int): def rank(self, value: int) -> None:
"""Set the process rank.""" """Set the process rank."""
self._rank = value self._rank = value
@@ -107,23 +114,23 @@ class LoggerCollection(LightningLoggerBase):
def experiment(self) -> List[Any]: def experiment(self) -> List[Any]:
return [logger.experiment for logger in self._logger_iterable] return [logger.experiment for logger in self._logger_iterable]
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None): def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None) -> None:
[logger.log_metrics(metrics, step) for logger in self._logger_iterable] [logger.log_metrics(metrics, step) for logger in self._logger_iterable]
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
[logger.log_hyperparams(params) for logger in self._logger_iterable] [logger.log_hyperparams(params) for logger in self._logger_iterable]
def save(self): def save(self) -> None:
[logger.save() for logger in self._logger_iterable] [logger.save() for logger in self._logger_iterable]
def finalize(self, status: str): def finalize(self, status: str) -> None:
[logger.finalize(status) for logger in self._logger_iterable] [logger.finalize(status) for logger in self._logger_iterable]
def close(self): def close(self) -> None:
[logger.close() for logger in self._logger_iterable] [logger.close() for logger in self._logger_iterable]
@LightningLoggerBase.rank.setter @LightningLoggerBase.rank.setter
def rank(self, value: int): def rank(self, value: int) -> None:
self._rank = value self._rank = value
for logger in self._logger_iterable: for logger in self._logger_iterable:
logger.rank = value logger.rank = value
+8 -7
View File
@@ -5,9 +5,9 @@ r"""
CometLogger CometLogger
------------- -------------
""" """
import argparse from argparse import Namespace
from logging import getLogger from logging import getLogger
from typing import Optional, Dict, Union from typing import Optional, Dict, Union, Any
try: try:
from comet_ml import Experiment as CometExperiment from comet_ml import Experiment as CometExperiment
@@ -162,15 +162,16 @@ class CometLogger(LightningLoggerBase):
return self._experiment return self._experiment
@rank_zero_only @rank_zero_only
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
self.experiment.log_parameters(vars(params)) params = self._convert_params(params)
self.experiment.log_parameters(params)
@rank_zero_only @rank_zero_only
def log_metrics( def log_metrics(
self, self,
metrics: Dict[str, Union[torch.Tensor, float]], metrics: Dict[str, Union[torch.Tensor, float]],
step: Optional[int] = None step: Optional[int] = None
): ) -> None:
# Comet.ml expects metrics to be a dictionary of detached tensors on CPU # Comet.ml expects metrics to be a dictionary of detached tensors on CPU
for key, val in metrics.items(): for key, val in metrics.items():
if is_tensor(val): if is_tensor(val):
@@ -182,7 +183,7 @@ class CometLogger(LightningLoggerBase):
self._experiment = None self._experiment = None
@rank_zero_only @rank_zero_only
def finalize(self, status: str): def finalize(self, status: str) -> None:
r""" r"""
When calling self.experiment.end(), that experiment won't log any more data to Comet. That's why, if you need When calling self.experiment.end(), that experiment won't log any more data to Comet. That's why, if you need
to log any more data you need to create an ExistingCometExperiment. For example, to log data when testing your to log any more data you need to create an ExistingCometExperiment. For example, to log data when testing your
@@ -199,7 +200,7 @@ class CometLogger(LightningLoggerBase):
return self.experiment.project_name return self.experiment.project_name
@name.setter @name.setter
def name(self, value: str): def name(self, value: str) -> None:
self.experiment.set_name(value) self.experiment.set_name(value)
@property @property
+7 -6
View File
@@ -23,10 +23,10 @@ Use the logger anywhere in you LightningModule as follows:
self.logger.experiment.whatever_ml_flow_supports(...) self.logger.experiment.whatever_ml_flow_supports(...)
""" """
import argparse from argparse import Namespace
from logging import getLogger from logging import getLogger
from time import time from time import time
from typing import Optional, Dict, Any from typing import Optional, Dict, Any, Union
try: try:
import mlflow import mlflow
@@ -88,12 +88,13 @@ class MLFlowLogger(LightningLoggerBase):
return self._run_id return self._run_id
@rank_zero_only @rank_zero_only
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
for k, v in vars(params).items(): params = self._convert_params(params)
for k, v in params.items():
self.experiment.log_param(self.run_id, k, v) self.experiment.log_param(self.run_id, k, v)
@rank_zero_only @rank_zero_only
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None): def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None) -> None:
timestamp_ms = int(time() * 1000) timestamp_ms = int(time() * 1000)
for k, v in metrics.items(): for k, v in metrics.items():
if isinstance(v, str): if isinstance(v, str):
@@ -105,7 +106,7 @@ class MLFlowLogger(LightningLoggerBase):
pass pass
@rank_zero_only @rank_zero_only
def finalize(self, status: str = 'FINISHED'): def finalize(self, status: str = 'FINISHED') -> None:
if status == 'success': if status == 'success':
status = 'FINISHED' status = 'FINISHED'
self.experiment.set_terminated(self.run_id, status) self.experiment.set_terminated(self.run_id, status)
+12 -11
View File
@@ -6,7 +6,7 @@ Log using `neptune-logger <https://www.neptune.ml>`_
NeptuneLogger NeptuneLogger
-------------- --------------
""" """
import argparse from argparse import Namespace
from logging import getLogger from logging import getLogger
from typing import Optional, List, Dict, Any, Union, Iterable from typing import Optional, List, Dict, Any, Union, Iterable
@@ -164,8 +164,9 @@ class NeptuneLogger(LightningLoggerBase):
return self._experiment return self._experiment
@rank_zero_only @rank_zero_only
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
for key, val in vars(params).items(): params = self._convert_params(params)
for key, val in params.items():
self.experiment.set_property(f'param__{key}', val) self.experiment.set_property(f'param__{key}', val)
@rank_zero_only @rank_zero_only
@@ -173,7 +174,7 @@ class NeptuneLogger(LightningLoggerBase):
self, self,
metrics: Dict[str, Union[torch.Tensor, float]], metrics: Dict[str, Union[torch.Tensor, float]],
step: Optional[int] = None step: Optional[int] = None
): ) -> None:
"""Log metrics (numeric values) in Neptune experiments """Log metrics (numeric values) in Neptune experiments
Args: Args:
@@ -184,7 +185,7 @@ class NeptuneLogger(LightningLoggerBase):
self.log_metric(key, val, step=step) self.log_metric(key, val, step=step)
@rank_zero_only @rank_zero_only
def finalize(self, status: str): def finalize(self, status: str) -> None:
self.experiment.stop() self.experiment.stop()
@property @property
@@ -207,7 +208,7 @@ class NeptuneLogger(LightningLoggerBase):
metric_name: str, metric_name: str,
metric_value: Union[torch.Tensor, float, str], metric_value: Union[torch.Tensor, float, str],
step: Optional[int] = None step: Optional[int] = None
): ) -> None:
"""Log metrics (numeric values) in Neptune experiments """Log metrics (numeric values) in Neptune experiments
Args: Args:
@@ -224,7 +225,7 @@ class NeptuneLogger(LightningLoggerBase):
self.experiment.log_metric(metric_name, x=step, y=metric_value) self.experiment.log_metric(metric_name, x=step, y=metric_value)
@rank_zero_only @rank_zero_only
def log_text(self, log_name: str, text: str, step: Optional[int] = None): def log_text(self, log_name: str, text: str, step: Optional[int] = None) -> None:
"""Log text data in Neptune experiment """Log text data in Neptune experiment
Args: Args:
@@ -235,7 +236,7 @@ class NeptuneLogger(LightningLoggerBase):
self.log_metric(log_name, text, step=step) self.log_metric(log_name, text, step=step)
@rank_zero_only @rank_zero_only
def log_image(self, log_name: str, image: Union[str, Any], step: Optional[int] = None): def log_image(self, log_name: str, image: Union[str, Any], step: Optional[int] = None) -> None:
"""Log image data in Neptune experiment """Log image data in Neptune experiment
Args: Args:
@@ -250,7 +251,7 @@ class NeptuneLogger(LightningLoggerBase):
self.experiment.log_image(log_name, x=step, y=image) self.experiment.log_image(log_name, x=step, y=image)
@rank_zero_only @rank_zero_only
def log_artifact(self, artifact: str, destination: Optional[str] = None): def log_artifact(self, artifact: str, destination: Optional[str] = None) -> None:
"""Save an artifact (file) in Neptune experiment storage. """Save an artifact (file) in Neptune experiment storage.
Args: Args:
@@ -261,7 +262,7 @@ class NeptuneLogger(LightningLoggerBase):
self.experiment.log_artifact(artifact, destination) self.experiment.log_artifact(artifact, destination)
@rank_zero_only @rank_zero_only
def set_property(self, key: str, value: Any): def set_property(self, key: str, value: Any) -> None:
"""Set key-value pair as Neptune experiment property. """Set key-value pair as Neptune experiment property.
Args: Args:
@@ -271,7 +272,7 @@ class NeptuneLogger(LightningLoggerBase):
self.experiment.set_property(key, value) self.experiment.set_property(key, value)
@rank_zero_only @rank_zero_only
def append_tags(self, tags: Union[str, Iterable[str]]): def append_tags(self, tags: Union[str, Iterable[str]]) -> None:
"""appends tags to neptune experiment """appends tags to neptune experiment
Args: Args:
+6 -13
View File
@@ -1,8 +1,7 @@
import argparse
import csv import csv
import os import os
from argparse import Namespace from argparse import Namespace
from typing import Optional, Dict, Union from typing import Optional, Dict, Union, Any
from warnings import warn from warnings import warn
import torch import torch
@@ -100,14 +99,8 @@ class TensorBoardLogger(LightningLoggerBase):
return self._experiment return self._experiment
@rank_zero_only @rank_zero_only
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
if params is None: params = self._convert_params(params)
return
# in case converting from namespace
if isinstance(params, Namespace):
params = vars(params)
params = dict(params)
if parse_version(torch.__version__) < parse_version("1.3.0"): if parse_version(torch.__version__) < parse_version("1.3.0"):
warn( warn(
@@ -126,14 +119,14 @@ class TensorBoardLogger(LightningLoggerBase):
self.tags.update(params) self.tags.update(params)
@rank_zero_only @rank_zero_only
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None): def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None) -> None:
for k, v in metrics.items(): for k, v in metrics.items():
if isinstance(v, torch.Tensor): if isinstance(v, torch.Tensor):
v = v.item() v = v.item()
self.experiment.add_scalar(k, v, step) self.experiment.add_scalar(k, v, step)
@rank_zero_only @rank_zero_only
def save(self): def save(self) -> None:
try: try:
self.experiment.flush() self.experiment.flush()
except AttributeError: except AttributeError:
@@ -156,7 +149,7 @@ class TensorBoardLogger(LightningLoggerBase):
writer.writerow({'key': k, 'value': v}) writer.writerow({'key': k, 'value': v})
@rank_zero_only @rank_zero_only
def finalize(self, status: str): def finalize(self, status: str) -> None:
self.save() self.save()
@property @property
+10 -9
View File
@@ -1,5 +1,5 @@
import argparse from argparse import Namespace
from typing import Optional, Dict, Any from typing import Optional, Dict, Any, Union
try: try:
from test_tube import Experiment from test_tube import Experiment
@@ -92,32 +92,33 @@ class TestTubeLogger(LightningLoggerBase):
return self._experiment return self._experiment
@rank_zero_only @rank_zero_only
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
# TODO: HACK figure out where this is being set to true # TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug self.experiment.debug = self.debug
self.experiment.argparse(params) params = self._convert_params(params)
self.experiment.argparse(Namespace(**params))
@rank_zero_only @rank_zero_only
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None): def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None) -> None:
# TODO: HACK figure out where this is being set to true # TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug self.experiment.debug = self.debug
self.experiment.log(metrics, global_step=step) self.experiment.log(metrics, global_step=step)
@rank_zero_only @rank_zero_only
def save(self): def save(self) -> None:
# TODO: HACK figure out where this is being set to true # TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug self.experiment.debug = self.debug
self.experiment.save() self.experiment.save()
@rank_zero_only @rank_zero_only
def finalize(self, status: str): def finalize(self, status: str) -> None:
# TODO: HACK figure out where this is being set to true # TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug self.experiment.debug = self.debug
self.save() self.save()
self.close() self.close()
@rank_zero_only @rank_zero_only
def close(self): def close(self) -> None:
# TODO: HACK figure out where this is being set to true # TODO: HACK figure out where this is being set to true
self.experiment.debug = self.debug self.experiment.debug = self.debug
if not self.debug: if not self.debug:
@@ -129,7 +130,7 @@ class TestTubeLogger(LightningLoggerBase):
return self._rank return self._rank
@rank.setter @rank.setter
def rank(self, value: int): def rank(self, value: int) -> None:
self._rank = value self._rank = value
if self._experiment is not None: if self._experiment is not None:
self.experiment.rank = value self.experiment.rank = value
+6 -5
View File
@@ -5,9 +5,9 @@ r"""
WandbLogger WandbLogger
------------- -------------
""" """
import argparse
import os import os
from typing import Optional, List, Dict from argparse import Namespace
from typing import Optional, List, Dict, Union, Any
import torch.nn as nn import torch.nn as nn
@@ -91,17 +91,18 @@ class WandbLogger(LightningLoggerBase):
wandb.watch(model, log=log, log_freq=log_freq) wandb.watch(model, log=log, log_freq=log_freq)
@rank_zero_only @rank_zero_only
def log_hyperparams(self, params: argparse.Namespace): def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
params = self._convert_params(params)
self.experiment.config.update(params) self.experiment.config.update(params)
@rank_zero_only @rank_zero_only
def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None): def log_metrics(self, metrics: Dict[str, float], step: Optional[int] = None) -> None:
if step is not None: if step is not None:
metrics['global_step'] = step metrics['global_step'] = step
self.experiment.log(metrics) self.experiment.log(metrics)
@rank_zero_only @rank_zero_only
def finalize(self, status: str = 'success'): def finalize(self, status: str = 'success') -> None:
try: try:
exit_code = 0 if status == 'success' else 1 exit_code = 0 if status == 'success' else 1
wandb.join(exit_code) wandb.join(exit_code)
+6 -4
View File
@@ -4,9 +4,9 @@ import re
import signal import signal
import warnings import warnings
from abc import ABC from abc import ABC
from argparse import Namespace
from subprocess import call from subprocess import call
from typing import Union from typing import Union
from copy import deepcopy
import torch import torch
import torch.distributed as dist import torch.distributed as dist
@@ -238,7 +238,9 @@ class TrainerIOMixin(ABC):
checkpoint['state_dict'] = model.state_dict() checkpoint['state_dict'] = model.state_dict()
if hasattr(model, "hparams"): if hasattr(model, "hparams"):
checkpoint['hparams'] = vars(model.hparams) is_namespace = isinstance(model.hparams, Namespace)
checkpoint['hparams'] = vars(model.hparams) if is_namespace else model.hparams
checkpoint['hparams_type'] = 'namespace' if is_namespace else 'dict'
else: else:
warnings.warn( warnings.warn(
"Did not find hyperparameters at model.hparams. Saving checkpoint without" "Did not find hyperparameters at model.hparams. Saving checkpoint without"
@@ -322,7 +324,7 @@ class TrainerIOMixin(ABC):
# ---------------------------------- # ----------------------------------
# PRIVATE OPS # PRIVATE OPS
# ---------------------------------- # ----------------------------------
def hpc_save(self, folderpath, logger): def hpc_save(self, folderpath: str, logger):
# make sure the checkpoint folder exists # make sure the checkpoint folder exists
os.makedirs(folderpath, exist_ok=True) os.makedirs(folderpath, exist_ok=True)
@@ -333,7 +335,7 @@ class TrainerIOMixin(ABC):
if not os.path.exists(folderpath): if not os.path.exists(folderpath):
os.makedirs(folderpath, exist_ok=True) os.makedirs(folderpath, exist_ok=True)
filepath = '{}/hpc_ckpt_{}.ckpt'.format(folderpath, ckpt_number) filepath = os.path.join(folderpath, f'hpc_ckpt_{ckpt_number}.ckpt')
# give model a chance to do something on hpc_save # give model a chance to do something on hpc_save
model = self.get_model() model = self.get_model()
+2 -2
View File
@@ -24,8 +24,8 @@ def test_wandb_logger(wandb):
logger.log_metrics({'acc': 1.0}, step=3) logger.log_metrics({'acc': 1.0}, step=3)
wandb.init().log.assert_called_once_with({'global_step': 3, 'acc': 1.0}) wandb.init().log.assert_called_once_with({'global_step': 3, 'acc': 1.0})
logger.log_hyperparams('test') logger.log_hyperparams({'test': None})
wandb.init().config.update.assert_called_once_with('test') wandb.init().config.update.assert_called_once_with({'test': None})
logger.watch('model', 'log', 10) logger.watch('model', 'log', 10)
wandb.watch.assert_called_once_with('model', log='log', log_freq=10) wandb.watch.assert_called_once_with('model', log='log', log_freq=10)
+1 -1
View File
@@ -2,7 +2,7 @@
import torch import torch
from .base import TestModelBase from .base import TestModelBase, DictHparamsModel
from .mixins import ( from .mixins import (
LightEmptyTestStep, LightEmptyTestStep,
LightValidationStepMixin, LightValidationStepMixin,
+23 -1
View File
@@ -6,9 +6,9 @@ import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
from torch import optim from torch import optim
from torch.utils.data import DataLoader from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from torchvision import transforms from torchvision import transforms
from torchvision.datasets import MNIST from torchvision.datasets import MNIST
from typing import Dict
try: try:
from test_tube import HyperOptArgumentParser from test_tube import HyperOptArgumentParser
@@ -36,6 +36,28 @@ class TestingMNIST(MNIST):
self.targets = self.targets[:num_samples] self.targets = self.targets[:num_samples]
class DictHparamsModel(LightningModule):
def __init__(self, hparams: Dict):
super(DictHparamsModel, self).__init__()
self.l1 = torch.nn.Linear(hparams.get('in_features'), hparams['out_features'])
def forward(self, x):
return torch.relu(self.l1(x.view(x.size(0), -1)))
def training_step(self, batch, batch_idx):
x, y = batch
y_hat = self.forward(x)
return {'loss': F.cross_entropy(y_hat, y)}
def configure_optimizers(self):
return torch.optim.Adam(self.parameters(), lr=0.02)
def train_dataloader(self):
return DataLoader(MNIST(os.getcwd(), train=True, download=True,
transform=transforms.ToTensor()), batch_size=32)
class TestModelBase(LightningModule): class TestModelBase(LightningModule):
""" """
Base LightningModule for testing. Implements only the required Base LightningModule for testing. Implements only the required
+14
View File
@@ -168,6 +168,20 @@ def load_model(exp, root_weights_dir, module_class=LightningTemplateModel, path_
return trained_model return trained_model
def load_model_from_checkpoint(root_weights_dir, module_class=LightningTemplateModel):
# load trained model
checkpoints = [x for x in os.listdir(root_weights_dir) if '.ckpt' in x]
weights_dir = os.path.join(root_weights_dir, checkpoints[0])
trained_model = module_class.load_from_checkpoint(
checkpoint_path=weights_dir,
)
assert trained_model is not None, 'loading model failed'
return trained_model
def run_prediction(dataloader, trained_model, dp=False, min_acc=0.45): def run_prediction(dataloader, trained_model, dp=False, min_acc=0.45):
# run prediction on 1 batch # run prediction on 1 batch
for batch in dataloader: for batch in dataloader:
+30 -10
View File
@@ -3,30 +3,28 @@ import math
import os import os
import pytest import pytest
import torch import torch
import argparse from argparse import ArgumentParser, Namespace
import tests.models.utils as tutils import tests.models.utils as tutils
from unittest import mock from unittest import mock
from pytorch_lightning import Trainer from pytorch_lightning import Trainer, LightningModule
from pytorch_lightning.callbacks import ( from pytorch_lightning.callbacks import (
EarlyStopping, EarlyStopping,
ModelCheckpoint, ModelCheckpoint,
) )
from tests.models import ( from tests.models import (
TestModelBase, TestModelBase,
DictHparamsModel,
LightningTestModel, LightningTestModel,
LightEmptyTestStep, LightEmptyTestStep,
LightValidationStepMixin, LightValidationStepMixin,
LightValidationMultipleDataloadersMixin, LightValidationMultipleDataloadersMixin,
LightTrainDataloader, LightTrainDataloader,
LightTestDataloader, LightTestDataloader,
LightValidationMixin,
LightTestMixin
) )
from pytorch_lightning.core.lightning import load_hparams_from_tags_csv from pytorch_lightning.core.lightning import load_hparams_from_tags_csv
from pytorch_lightning.trainer.logging import TrainerLoggingMixin from pytorch_lightning.trainer.logging import TrainerLoggingMixin
from pytorch_lightning.utilities.debugging import MisconfigurationException from pytorch_lightning.utilities.debugging import MisconfigurationException
from pytorch_lightning import Callback
def test_no_val_module(tmpdir): def test_no_val_module(tmpdir):
@@ -128,7 +126,7 @@ def test_gradient_accumulation_scheduling(tmpdir):
assert Trainer(accumulate_grad_batches={1: 2.5, 3: 5}) assert Trainer(accumulate_grad_batches={1: 2.5, 3: 5})
# test optimizer call freq matches scheduler # test optimizer call freq matches scheduler
def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx, second_order_closure=None): def _optimizer_step(self, epoch, batch_idx, optimizer, optimizer_idx, second_order_closure=None):
# only test the first 12 batches in epoch # only test the first 12 batches in epoch
if batch_idx < 12: if batch_idx < 12:
if epoch == 0: if epoch == 0:
@@ -179,7 +177,7 @@ def test_gradient_accumulation_scheduling(tmpdir):
default_save_path=tmpdir) default_save_path=tmpdir)
# for the test # for the test
trainer.optimizer_step = optimizer_step trainer.optimizer_step = _optimizer_step
model.prev_called_batch_idx = 0 model.prev_called_batch_idx = 0
trainer.fit(model) trainer.fit(model)
@@ -188,7 +186,6 @@ def test_gradient_accumulation_scheduling(tmpdir):
def test_loading_meta_tags(tmpdir): def test_loading_meta_tags(tmpdir):
tutils.reset_seed() tutils.reset_seed()
from argparse import Namespace
hparams = tutils.get_hparams() hparams = tutils.get_hparams()
# save tags # save tags
@@ -604,8 +601,9 @@ def test_testpass_overrides(tmpdir):
model = LightningTestModel(hparams) model = LightningTestModel(hparams)
Trainer().test(model) Trainer().test(model)
@mock.patch('argparse.ArgumentParser.parse_args', @mock.patch('argparse.ArgumentParser.parse_args',
return_value=argparse.Namespace(**Trainer.default_attributes())) return_value=Namespace(**Trainer.default_attributes()))
def test_default_args(tmpdir): def test_default_args(tmpdir):
"""Tests default argument parser for Trainer""" """Tests default argument parser for Trainer"""
tutils.reset_seed() tutils.reset_seed()
@@ -613,7 +611,7 @@ def test_default_args(tmpdir):
# logger file to get meta # logger file to get meta
logger = tutils.get_test_tube_logger(tmpdir, False) logger = tutils.get_test_tube_logger(tmpdir, False)
parser = argparse.ArgumentParser(add_help=False) parser = ArgumentParser(add_help=False)
args = parser.parse_args() args = parser.parse_args()
args.logger = logger args.logger = logger
@@ -622,3 +620,25 @@ def test_default_args(tmpdir):
assert isinstance(trainer, Trainer) assert isinstance(trainer, Trainer)
assert trainer.max_epochs == 5 assert trainer.max_epochs == 5
def test_hparams_save_load(tmpdir):
model = DictHparamsModel({'in_features': 28 * 28, 'out_features': 10})
# logger file to get meta
trainer_options = dict(
default_save_path=tmpdir,
max_epochs=2,
)
# fit model
trainer = Trainer(**trainer_options)
result = trainer.fit(model)
assert result == 1
# try to load the model now
pretrained_model = tutils.load_model_from_checkpoint(
trainer.checkpoint_callback.dirpath,
module_class=DictHparamsModel
)