Skip to content

Commit aaf6295

Browse files
michal2409nv-kkudrynski
authored andcommitted
[nnUNet/PyT] Add support for multi-node benchmarking
1 parent 9cb7dd0 commit aaf6295

5 files changed

Lines changed: 21 additions & 29 deletions

File tree

PyTorch/Segmentation/nnUNet/main.py

Lines changed: 2 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@
1515
import os
1616

1717
from pytorch_lightning import Trainer, seed_everything
18-
from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint, ModelSummary, RichProgressBar
18+
from pytorch_lightning.callbacks import ModelCheckpoint, ModelSummary, RichProgressBar
1919
from pytorch_lightning.loggers import TensorBoardLogger
2020

2121
from data_loading.data_module import DataModule
@@ -43,7 +43,7 @@
4343
LoggingCallback(
4444
log_dir=args.results,
4545
filnename=filnename,
46-
global_batch_size=batch_size * args.gpus,
46+
global_batch_size=batch_size * args.gpus * args.nodes,
4747
mode=args.exec_mode,
4848
warmup=args.warmup,
4949
dim=args.dim,
@@ -57,14 +57,6 @@
5757
default_hp_metric=False,
5858
version=0,
5959
)
60-
callbacks.append(
61-
EarlyStopping(
62-
monitor="dice",
63-
patience=args.patience,
64-
verbose=True,
65-
mode="max",
66-
)
67-
)
6860
if args.save_ckpt:
6961
callbacks.append(
7062
ModelCheckpoint(

PyTorch/Segmentation/nnUNet/nnunet/nn_unet.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,10 +23,11 @@
2323
from monai.inferers import sliding_window_inference
2424
from monai.networks.nets import DynUNet
2525
from monai.optimizers.lr_scheduler import WarmupCosineSchedule
26+
from pytorch_lightning.utilities import rank_zero_only
2627
from scipy.special import expit, softmax
2728
from skimage.transform import resize
2829
from utils.logger import DLLogger
29-
from utils.utils import get_config_file, print0, rank_zero
30+
from utils.utils import get_config_file, print0
3031

3132
from nnunet.loss import Loss, LossBraTS
3233
from nnunet.metrics import Dice
@@ -82,6 +83,8 @@ def training_step(self, batch, batch_idx):
8283
return loss
8384

8485
def validation_step(self, batch, batch_idx):
86+
if self.current_epoch < self.args.skip_first_n_eval:
87+
return None
8588
img, lbl = batch["image"], batch["label"]
8689
pred = self._forward(img)
8790
loss = self.loss(pred, lbl)
@@ -205,6 +208,11 @@ def round(self, tensor):
205208
return round(torch.mean(tensor).item(), 2)
206209

207210
def validation_epoch_end(self, outputs):
211+
if self.current_epoch < self.args.skip_first_n_eval:
212+
self.log("Dice", 0.001 * self.current_epoch) # To prevent early stopping
213+
self.dice.reset()
214+
return None
215+
208216
dice, loss = self.dice.compute()
209217
self.dice.reset()
210218

@@ -233,7 +241,7 @@ def test_epoch_end(self, outputs):
233241
if self.args.exec_mode == "evaluate":
234242
self.eval_dice, _ = self.dice.compute()
235243

236-
@rank_zero
244+
@rank_zero_only
237245
def on_fit_end(self):
238246
if not self.args.benchmark:
239247
metrics = {}

PyTorch/Segmentation/nnUNet/scripts/benchmark.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
parser.add_argument("--mode", type=str, required=True, choices=["train", "predict"], help="Benchmarking mode")
2222
parser.add_argument("--task", type=str, default="01", help="Task code")
2323
parser.add_argument("--gpus", type=int, default=1, help="Number of GPUs to use")
24+
parser.add_argument("--nodes", type=int, default=1, help="Number of nodes to use")
2425
parser.add_argument("--dim", type=int, required=True, help="Dimension of UNet")
2526
parser.add_argument("--batch_size", type=int, default=2, help="Batch size")
2627
parser.add_argument("--amp", action="store_true", help="Enable automatic mixed precision")
@@ -40,6 +41,7 @@
4041
cmd += f"--exec_mode {args.mode} "
4142
cmd += f"--dim {args.dim} "
4243
cmd += f"--gpus {args.gpus} "
44+
cmd += f"--nodes {args.nodes} "
4345
cmd += f"--train_batches {args.train_batches} "
4446
cmd += f"--test_batches {args.test_batches} "
4547
cmd += f"--warmup {args.warmup} "

PyTorch/Segmentation/nnUNet/utils/logger.py

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -19,30 +19,29 @@
1919
import numpy as np
2020
from dllogger import JSONStreamBackend, StdOutBackend, Verbosity
2121
from pytorch_lightning import Callback
22-
23-
from utils.utils import rank_zero
22+
from pytorch_lightning.utilities import rank_zero_only
2423

2524

2625
class DLLogger:
2726
def __init__(self, log_dir, filename, append=True):
2827
super().__init__()
2928
self._initialize_dllogger(log_dir, filename, append)
3029

31-
@rank_zero
30+
@rank_zero_only
3231
def _initialize_dllogger(self, log_dir, filename, append):
3332
backends = [
3433
JSONStreamBackend(Verbosity.VERBOSE, os.path.join(log_dir, filename), append=append),
3534
StdOutBackend(Verbosity.VERBOSE),
3635
]
3736
logger.init(backends=backends)
3837

39-
@rank_zero
38+
@rank_zero_only
4039
def log_metrics(self, metrics, step=None):
4140
if step is None:
4241
step = ()
4342
logger.log(step=step, data=metrics)
4443

45-
@rank_zero
44+
@rank_zero_only
4645
def flush(self):
4746
logger.flush()
4847

@@ -85,7 +84,7 @@ def _round3(val):
8584

8685
return stats
8786

88-
@rank_zero
87+
@rank_zero_only
8988
def _log(self):
9089
stats = self.process_performance_stats(np.diff(self.timestamps))
9190
self.dllogger.log_metrics(metrics=stats)

PyTorch/Segmentation/nnUNet/utils/utils.py

Lines changed: 2 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -15,23 +15,14 @@
1515
import ctypes
1616
import os
1717
import pickle
18-
from functools import wraps
1918
from subprocess import run
2019

2120
import numpy as np
2221
import torch
22+
from pytorch_lightning.utilities import rank_zero_only
2323

2424

25-
def rank_zero(fn):
26-
@wraps(fn)
27-
def wrapped_fn(*args, **kwargs):
28-
if int(os.getenv("LOCAL_RANK", "0")) == 0:
29-
return fn(*args, **kwargs)
30-
31-
return wrapped_fn
32-
33-
34-
@rank_zero
25+
@rank_zero_only
3526
def print0(text):
3627
print(text)
3728

0 commit comments

Comments
 (0)