mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-10 12:21:57 +08:00
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:
co-authored by
William Falcon
parent
6a39573267
commit
e586ed4767
@@ -1,5 +1,6 @@
|
||||
import argparse
|
||||
from abc import ABC, abstractmethod
|
||||
from argparse import Namespace
|
||||
from functools import wraps
|
||||
from typing import Union, Optional, Dict, Iterable, Any, Callable, List
|
||||
|
||||
@@ -41,6 +42,12 @@ class LightningLoggerBase(ABC):
|
||||
"""
|
||||
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
|
||||
def log_hyperparams(self, params: argparse.Namespace):
|
||||
"""Record hyperparameters.
|
||||
@@ -50,11 +57,11 @@ class LightningLoggerBase(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def save(self):
|
||||
def save(self) -> None:
|
||||
"""Save log data."""
|
||||
pass
|
||||
|
||||
def finalize(self, status: str):
|
||||
def finalize(self, status: str) -> None:
|
||||
"""Do any processing that is necessary to finalize an experiment.
|
||||
|
||||
Args:
|
||||
@@ -62,7 +69,7 @@ class LightningLoggerBase(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
def close(self):
|
||||
def close(self) -> None:
|
||||
"""Do any cleanup that is necessary to close an experiment."""
|
||||
pass
|
||||
|
||||
@@ -72,7 +79,7 @@ class LightningLoggerBase(ABC):
|
||||
return self._rank
|
||||
|
||||
@rank.setter
|
||||
def rank(self, value: int):
|
||||
def rank(self, value: int) -> None:
|
||||
"""Set the process rank."""
|
||||
self._rank = value
|
||||
|
||||
@@ -107,23 +114,23 @@ class LoggerCollection(LightningLoggerBase):
|
||||
def experiment(self) -> List[Any]:
|
||||
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]
|
||||
|
||||
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]
|
||||
|
||||
def save(self):
|
||||
def save(self) -> None:
|
||||
[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]
|
||||
|
||||
def close(self):
|
||||
def close(self) -> None:
|
||||
[logger.close() for logger in self._logger_iterable]
|
||||
|
||||
@LightningLoggerBase.rank.setter
|
||||
def rank(self, value: int):
|
||||
def rank(self, value: int) -> None:
|
||||
self._rank = value
|
||||
for logger in self._logger_iterable:
|
||||
logger.rank = value
|
||||
|
||||
Reference in New Issue
Block a user