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
+6 -13
View File
@@ -1,8 +1,7 @@
import argparse
import csv
import os
from argparse import Namespace
from typing import Optional, Dict, Union
from typing import Optional, Dict, Union, Any
from warnings import warn
import torch
@@ -100,14 +99,8 @@ class TensorBoardLogger(LightningLoggerBase):
return self._experiment
@rank_zero_only
def log_hyperparams(self, params: argparse.Namespace):
if params is None:
return
# in case converting from namespace
if isinstance(params, Namespace):
params = vars(params)
params = dict(params)
def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None:
params = self._convert_params(params)
if parse_version(torch.__version__) < parse_version("1.3.0"):
warn(
@@ -126,14 +119,14 @@ class TensorBoardLogger(LightningLoggerBase):
self.tags.update(params)
@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():
if isinstance(v, torch.Tensor):
v = v.item()
self.experiment.add_scalar(k, v, step)
@rank_zero_only
def save(self):
def save(self) -> None:
try:
self.experiment.flush()
except AttributeError:
@@ -156,7 +149,7 @@ class TensorBoardLogger(LightningLoggerBase):
writer.writerow({'key': k, 'value': v})
@rank_zero_only
def finalize(self, status: str):
def finalize(self, status: str) -> None:
self.save()
@property