mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-11 12:31:23 +08:00
add version_ prefix to log_dir (#706)
* add version_ prefix to log_dir * add version_ prefix
This commit is contained in:
@@ -63,7 +63,7 @@ class TensorBoardLogger(LightningLoggerBase):
|
||||
|
||||
root_dir = os.path.join(self.save_dir, self.name)
|
||||
os.makedirs(root_dir, exist_ok=True)
|
||||
log_dir = os.path.join(root_dir, str(self.version))
|
||||
log_dir = os.path.join(root_dir, "version_" + str(self.version))
|
||||
self._experiment = SummaryWriter(log_dir=log_dir, **self.kwargs)
|
||||
return self._experiment
|
||||
|
||||
@@ -131,9 +131,11 @@ class TensorBoardLogger(LightningLoggerBase):
|
||||
|
||||
def _get_next_version(self):
|
||||
root_dir = os.path.join(self.save_dir, self.name)
|
||||
existing_versions = [
|
||||
int(d) for d in os.listdir(root_dir) if os.path.isdir(os.path.join(root_dir, d)) and d.isdigit()
|
||||
]
|
||||
existing_versions = []
|
||||
for d in os.listdir(root_dir):
|
||||
if os.path.isdir(os.path.join(root_dir, d)) and d.startswith("version_"):
|
||||
existing_versions.append(int(d.split("_")[1]))
|
||||
|
||||
if len(existing_versions) == 0:
|
||||
return 0
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user