diff --git a/CHANGELOG.md b/CHANGELOG.md index 91ac33b8..6a1f31e3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -60,8 +60,10 @@ The format is based on [Keep a Changelog](http://keepachangelog.com/en/1.0.0/). - Fixed all warnings and errors in the docs build process ([#1191](https://github.com/PyTorchLightning/pytorch-lightning/pull/1191)) - Fixed an issue where `val_percent_check=0` would not disable validation ([#1251](https://github.com/PyTorchLightning/pytorch-lightning/pull/1251)) - Fixed average of incomplete `TensorRunningMean` ([#1309](https://github.com/PyTorchLightning/pytorch-lightning/pull/1309)) +- Fixed `WandbLogger.watch` with `wandb.init()` ([#1311](https://github.com/PyTorchLightning/pytorch-lightning/pull/1311)) - Fixed an issue with early stopping that would prevent it from monitoring training metrics when validation is disabled / not implemented ([#1235](https://github.com/PyTorchLightning/pytorch-lightning/pull/1235)). - Fixed a bug that would cause `trainer.test()` to run on the validation set when overloading `validation_epoch_end ` and `test_end` ([#1353](https://github.com/PyTorchLightning/pytorch-lightning/pull/1353)). +- Fixed `WandbLogger.watch` ([#1311](https://github.com/PyTorchLightning/pytorch-lightning/pull/1311)) ## [0.7.1] - 2020-03-07 diff --git a/pytorch_lightning/loggers/wandb.py b/pytorch_lightning/loggers/wandb.py index 5890873c..39bfcc3c 100644 --- a/pytorch_lightning/loggers/wandb.py +++ b/pytorch_lightning/loggers/wandb.py @@ -94,7 +94,7 @@ class WandbLogger(LightningLoggerBase): return self._experiment def watch(self, model: nn.Module, log: str = 'gradients', log_freq: int = 100): - wandb.watch(model, log=log, log_freq=log_freq) + self.experiment.watch(model, log=log, log_freq=log_freq) @rank_zero_only def log_hyperparams(self, params: Union[Dict[str, Any], Namespace]) -> None: diff --git a/tests/loggers/test_wandb.py b/tests/loggers/test_wandb.py index 774f722c..bd430ca2 100644 --- a/tests/loggers/test_wandb.py +++ b/tests/loggers/test_wandb.py @@ -28,7 +28,7 @@ def test_wandb_logger(wandb): wandb.init().config.update.assert_called_once_with({'test': None}) logger.watch('model', 'log', 10) - wandb.watch.assert_called_once_with('model', log='log', log_freq=10) + wandb.init().watch.assert_called_once_with('model', log='log', log_freq=10) assert logger.name == wandb.init().project_name() assert logger.version == wandb.init().id