From 76af84718a1f9c61903db3e1557ebd5b436e1c97 Mon Sep 17 00:00:00 2001 From: Anthony Bisulco Date: Sun, 10 May 2020 13:15:51 -0400 Subject: [PATCH] Group argument wandb (#1760) * group argument wandb * formatting fix --- pytorch_lightning/loggers/wandb.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/pytorch_lightning/loggers/wandb.py b/pytorch_lightning/loggers/wandb.py index 0d5ff985..3b15677a 100644 --- a/pytorch_lightning/loggers/wandb.py +++ b/pytorch_lightning/loggers/wandb.py @@ -39,6 +39,7 @@ class WandbLogger(LightningLoggerBase): log_model: Save checkpoints in wandb dir to upload on W&B servers. experiment: WandB experiment object entity: The team posting this run (default: your username or your default team) + group: A unique string shared by all runs in a given group Example: >>> from pytorch_lightning.loggers import WandbLogger @@ -64,7 +65,8 @@ class WandbLogger(LightningLoggerBase): tags: Optional[List[str]] = None, log_model: bool = False, experiment=None, - entity=None): + entity=None, + group: Optional[str] = None): super().__init__() self._name = name self._save_dir = save_dir @@ -76,6 +78,7 @@ class WandbLogger(LightningLoggerBase): self._offline = offline self._entity = entity self._log_model = log_model + self._group = group def __getstate__(self): state = self.__dict__.copy() @@ -103,7 +106,8 @@ class WandbLogger(LightningLoggerBase): os.environ['WANDB_MODE'] = 'dryrun' self._experiment = wandb.init( name=self._name, dir=self._save_dir, project=self._project, anonymous=self._anonymous, - reinit=True, id=self._id, resume='allow', tags=self._tags, entity=self._entity) + reinit=True, id=self._id, resume='allow', tags=self._tags, entity=self._entity, + group=self._group) # save checkpoints in wandb dir to upload on W&B servers if self._log_model: self.save_dir = self._experiment.dir