From 31d7745e624f141d5080090198f61b6df879a472 Mon Sep 17 00:00:00 2001 From: sophmrtn Date: Thu, 17 Jun 2021 19:25:29 +0100 Subject: [PATCH 01/62] Adds learning rate scheduling options --- src/rectangle/utils/train.py | 36 +++++++++++++++++++++++++++++++++++- 1 file changed, 35 insertions(+), 1 deletion(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index a6802b8..cde40fb 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -15,7 +15,7 @@ class Trainer(nn.Module): def __init__(self, model, nb_epochs=200, outdir='./logs', loss=DiceLoss(), metric=DiceLoss(), opt='adam', print_interval=1, val_interval=5, device='cuda', - early_stop=5, ensemble=None): + early_stop=5, lr_schedule=None, ensemble=None): super().__init__() @@ -53,6 +53,13 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', self.opt = opt + if lr_schedule: + if lr_schedule not in ['lambda', 'exponential', 'reduce_on_plateau']: + raise ValueError('Available learning rate schedules are LambdaLR, ExponentialLR or ReduceLROnPlateau.') + elif self.ensemble: + lr_schedule = [lr_schedule for model in self.model_ensemble] + + self.lr_schedule = lr_schedule def train(self, train_data, val_data=None, oname=None, train_pre=None, train_post=None, train_batch=128, train_shuffle=True, @@ -93,6 +100,17 @@ def train(self, train_data, val_data=None, oname=None, train = train_list[i] val = val_list[i] opt_ = self.opt[i] + if self.lr_schedule: + lr_schedule_ = self.lr_schedule[i] + if lr_schedule_ == 'lambda': + lr_schedule_ = torch.optim.lr_scheduler.LambdaLR(opt_, lambda epoch: 0.95 ** epoch) + elif lr_schedule_ == 'exponential': + lr_schedule_ = torch.optim.lr_scheduler.ExponentialLR(opt_, 0.95) + else: + lr_schedule_ = torch.optim.lr_scheduler.ReduceLROnPlateau(opt_) + else: + lr_schedule_ = None + print('Beginning training of model #{}'.format(i)) for epoch in range(self.nb_epochs): if self.early_stop: @@ -113,6 +131,8 @@ def train(self, train_data, val_data=None, oname=None, loss_ = self.loss(pred, label) loss_.backward() opt_.step() + if lr_schedule_ and lr_schedule_ != 'reduce_on_plateau': + lr_schedule_.step() loss_epoch.append(loss_.item()) loss_log_ensemble[i,epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: @@ -131,6 +151,8 @@ def train(self, train_data, val_data=None, oname=None, for aug in val_post: pred = aug(pred) dice_metric = self.metric(pred, label) + if lr_schedule_ == 'reduce_on_plateau': + lr_schedule_.step(dice_metric) dice_epoch.append(1 - dice_metric.item()) dice_log_ensemble[i,int(epoch//self.val_interval)] = np.nanmean(dice_epoch) if epoch >= self.val_interval: @@ -155,6 +177,14 @@ def train(self, train_data, val_data=None, oname=None, early_ = 0 dice_max = 0 model = self.model + if self.lr_schedule: + if self.lr_schedule == 'lambda': + self.lr_schedule = torch.optim.lr_scheduler.LambdaLR(self.opt, lambda epoch: 0.95 ** epoch) + elif self.lr_schedule == 'exponential': + self.lr_schedule = torch.optim.lr_scheduler.ExponentialLR(self.opt, 0.95) + else: + self.lr_schedule = torch.optim.lr_scheduler.ReduceLROnPlateau(self.opt) + for epoch in range(self.nb_epochs): if self.early_stop: if early_ == self.early_stop: @@ -174,6 +204,8 @@ def train(self, train_data, val_data=None, oname=None, loss_ = self.loss(pred, label) loss_.backward() self.opt.step() + if self.lr_schedule and self.lr_schedule != 'reduce_on_plateau': + self.lr_schedule.step() loss_epoch.append(loss_.item()) loss_log[epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: @@ -192,6 +224,8 @@ def train(self, train_data, val_data=None, oname=None, for aug in val_post: pred = aug(pred) dice_metric = self.metric(pred, label) + if self.lr_schedule == 'reduce_on_plateau': + self.lr_schedule.step(dice_metric) # monitors validation loss dice_epoch.append(1 - dice_metric.item()) dice_log[int(epoch//self.val_interval)] = np.nanmean(dice_epoch) if epoch % self.print_interval == 0: From 43c8f903182c038f53bff1dde31bfd197df1fe3e Mon Sep 17 00:00:00 2001 From: Jiongqi Date: Sun, 20 Jun 2021 22:33:20 +0800 Subject: [PATCH 02/62] add code for tensorboard for both models --- src/rectangle/utils/train.py | 62 ++++++++++++++++++++++++++++++++++++ 1 file changed, 62 insertions(+) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index cde40fb..1cfdb67 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -9,6 +9,8 @@ from torch.utils.data import DataLoader, random_split, ConcatDataset import numpy as np from scipy.ndimage import laplace +from torch.utils.tensorboard import SummaryWriter +from torchvision.utils import make_grid class Trainer(nn.Module): @@ -29,6 +31,7 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', self.ensemble = ensemble self.outdir = outdir self.device = device + self.writer = SummaryWriter() if self.ensemble == 0: self.ensemble = None @@ -136,6 +139,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log_ensemble[i,epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: + self.writer.add_scalar('train/dice_loss_ensemble', loss_, epoch) + self.writer.add_scalar('train/dice_coefficient_ensemble', 1-loss_, epoch) print('Epoch #{}: Mean Dice Loss: {}'.format(epoch, loss_log_ensemble[i,epoch])) if epoch % self.val_interval == 0: dice_epoch = [] @@ -155,6 +160,29 @@ def train(self, train_data, val_data=None, oname=None, lr_schedule_.step(dice_metric) dice_epoch.append(1 - dice_metric.item()) dice_log_ensemble[i,int(epoch//self.val_interval)] = np.nanmean(dice_epoch) + + self.writer.add_scalar('val/dice_loss_ensemble', dice_metric, epoch) + self.writer.add_scalar('val/dice_coefficient_ensemble', 1-dice_metric, epoch) + + ## show some (e.g.,10) example images in tensorboard + ex_num = 10 + ex_label = label[:ex_num] + ex_pred = pred[:ex_num] + + ex_labels = torch.empty(0).to(self.device) + ex_pres = torch.empty(0).to(self.device) + for i in range(ex_num): + ex_labels = torch.cat([ex_labels,ex_label[i]], dim=1) + ex_preds = torch.cat([ex_labels,ex_label[i]], dim=1) + ex_images = torch.cat([ex_labels,ex_preds], dim=0) + image_grid = (make_grid(ex_images, nrow=ex_num)[0]+0.5)/ex_num + self.writer.add_images( + "val/example_images_ensemble", + image_grid, + self.trainer.global_step, + dataformats="HW", + ) + if epoch >= self.val_interval: if dice_log_ensemble[i,int(epoch//self.val_interval)] > dice_max: early_ = 0 @@ -170,6 +198,7 @@ def train(self, train_data, val_data=None, oname=None, print('Mean Validation Dice: {}'.format(dice_log_ensemble[i,int(epoch//self.val_interval)])) print('Finished training of model #{}'.format(i)) else: + print('non-ensemble') loss_log = np.empty(self.nb_epochs) loss_log[:] = np.nan dice_log = np.empty(int(self.nb_epochs//self.val_interval)) @@ -209,6 +238,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log[epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: + self.writer.add_scalar('train/dice_loss', loss_, epoch) + self.writer.add_scalar('train/dice_coefficient', 1-loss_, epoch) print('Epoch #{}: Mean Dice Loss: {}'.format(epoch, loss_log[epoch])) if epoch % self.val_interval == 0: dice_epoch = [] @@ -228,6 +259,29 @@ def train(self, train_data, val_data=None, oname=None, self.lr_schedule.step(dice_metric) # monitors validation loss dice_epoch.append(1 - dice_metric.item()) dice_log[int(epoch//self.val_interval)] = np.nanmean(dice_epoch) + + self.writer.add_scalar('val/dice_loss', dice_metric, epoch) + self.writer.add_scalar('val/dice_coefficient', 1-dice_metric, epoch) + + ## show some (e.g.,10) example images in tensorboard + ex_num = 10 + ex_label = label[:ex_num] + ex_pred = pred[:ex_num] + + ex_labels = torch.empty(0).to(self.device) + ex_pres = torch.empty(0).to(self.device) + for i in range(ex_num): + ex_labels = torch.cat([ex_labels,ex_label[i]], dim=1) + ex_preds = torch.cat([ex_labels,ex_label[i]], dim=1) + ex_images = torch.cat([ex_labels,ex_preds], dim=0) + image_grid = (make_grid(ex_images, nrow=ex_num)[0]+0.5)/ex_num + self.writer.add_images( + "val/example_images", + image_grid, + self.trainer.global_step, + dataformats="HW", + ) + if epoch % self.print_interval == 0: print('Mean Validation Dice: {}'.format(dice_log[int(epoch//self.val_interval)])) if epoch >= self.val_interval: @@ -517,6 +571,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log_ensemble[i,epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: + self.writer.add_scalar('class_train/dice_loss_ensemble', loss_, epoch) + self.writer.add_scalar('class_train/dice_coefficient_ensemble', 1-loss_, epoch) print('Epoch #{}: Mean acc Loss: {}'.format(epoch, loss_log_ensemble[i,epoch])) if epoch % self.val_interval == 0: acc_epoch = [] @@ -534,6 +590,8 @@ def train(self, train_data, val_data=None, oname=None, acc_metric = self.metric(pred, label) acc_epoch.append(acc_metric) acc_log_ensemble[i,int(epoch//self.val_interval)] = np.nanmean(acc_epoch) + self.writer.add_scalar('class_val/dice_loss_ensemble', acc_metric, epoch) + self.writer.add_scalar('class_val/dice_coefficient_ensemble', 1-acc_metric, epoch) if epoch >= self.val_interval: if acc_log_ensemble[i,int(epoch//self.val_interval)] > acc_max: early_ = 0 @@ -578,6 +636,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log[epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: + self.writer.add_scalar('class_train/dice_loss', loss_, epoch) + self.writer.add_scalar('class_train/dice_coefficient', 1-loss_, epoch) print('Epoch #{}: Mean acc Loss: {}'.format(epoch, loss_log[epoch])) if epoch % self.val_interval == 0: acc_epoch = [] @@ -595,6 +655,8 @@ def train(self, train_data, val_data=None, oname=None, acc_metric = self.metric(pred, label) acc_epoch.append(acc_metric) acc_log[int(epoch//self.val_interval)] = np.nanmean(acc_epoch) + self.writer.add_scalar('class_val/dice_loss', acc_metric, epoch) + self.writer.add_scalar('class_val/dice_coefficient', 1-acc_metric, epoch) if epoch % self.print_interval == 0: print('Mean Validation acc: {}'.format(acc_log[int(epoch//self.val_interval)])) if epoch >= self.val_interval: From d742df941e1628ba0feed1d95e0e4f739e8c9ae2 Mon Sep 17 00:00:00 2001 From: Jiongqi Date: Sun, 20 Jun 2021 23:43:41 +0800 Subject: [PATCH 03/62] fix bugs in code relevant to tensorboard --- src/rectangle/utils/train.py | 36 ++++++++++++++---------------------- 1 file changed, 14 insertions(+), 22 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index 1cfdb67..c97effb 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -17,7 +17,7 @@ class Trainer(nn.Module): def __init__(self, model, nb_epochs=200, outdir='./logs', loss=DiceLoss(), metric=DiceLoss(), opt='adam', print_interval=1, val_interval=5, device='cuda', - early_stop=5, lr_schedule=None, ensemble=None): + early_stop=5, lr_schedule=None, ensemble=None): super().__init__() @@ -166,20 +166,16 @@ def train(self, train_data, val_data=None, oname=None, ## show some (e.g.,10) example images in tensorboard ex_num = 10 - ex_label = label[:ex_num] - ex_pred = pred[:ex_num] - - ex_labels = torch.empty(0).to(self.device) - ex_pres = torch.empty(0).to(self.device) - for i in range(ex_num): - ex_labels = torch.cat([ex_labels,ex_label[i]], dim=1) - ex_preds = torch.cat([ex_labels,ex_label[i]], dim=1) - ex_images = torch.cat([ex_labels,ex_preds], dim=0) - image_grid = (make_grid(ex_images, nrow=ex_num)[0]+0.5)/ex_num + ex_label = label[:ex_num,0] + ex_pred = pred[:ex_num,0] + ex_image = torch.cat([ex_label,ex_pred], dim=2) + + ex_images = ex_image.reshape(-1,ex_image.shape[2]) + image_grid = (make_grid(ex_images, nrow=ex_num)[0]+0.5)/ex_num self.writer.add_images( "val/example_images_ensemble", image_grid, - self.trainer.global_step, + epoch, dataformats="HW", ) @@ -265,20 +261,16 @@ def train(self, train_data, val_data=None, oname=None, ## show some (e.g.,10) example images in tensorboard ex_num = 10 - ex_label = label[:ex_num] - ex_pred = pred[:ex_num] - - ex_labels = torch.empty(0).to(self.device) - ex_pres = torch.empty(0).to(self.device) - for i in range(ex_num): - ex_labels = torch.cat([ex_labels,ex_label[i]], dim=1) - ex_preds = torch.cat([ex_labels,ex_label[i]], dim=1) - ex_images = torch.cat([ex_labels,ex_preds], dim=0) + ex_label = label[:ex_num,0] + ex_pred = pred[:ex_num,0] + ex_image = torch.cat([ex_label,ex_pred], dim=2) + + ex_images = ex_image.reshape(-1,ex_image.shape[2]) image_grid = (make_grid(ex_images, nrow=ex_num)[0]+0.5)/ex_num self.writer.add_images( "val/example_images", image_grid, - self.trainer.global_step, + epoch, dataformats="HW", ) From 671be7f6d3c8547bf40f65ee0e379063741fb038 Mon Sep 17 00:00:00 2001 From: Jiongqi Date: Sun, 20 Jun 2021 23:52:54 +0800 Subject: [PATCH 04/62] remove unnecessary comment in train.py while adding tensorboard code --- src/rectangle/utils/train.py | 1 - 1 file changed, 1 deletion(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index c97effb..0a92d08 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -194,7 +194,6 @@ def train(self, train_data, val_data=None, oname=None, print('Mean Validation Dice: {}'.format(dice_log_ensemble[i,int(epoch//self.val_interval)])) print('Finished training of model #{}'.format(i)) else: - print('non-ensemble') loss_log = np.empty(self.nb_epochs) loss_log[:] = np.nan dice_log = np.empty(int(self.nb_epochs//self.val_interval)) From 64cdd6bfba44ee0e3c90b6fd5ecbc67774f286c6 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Sun, 20 Jun 2021 19:36:04 +0100 Subject: [PATCH 05/62] Send precision and recall values to cpu --- src/rectangle/utils/train.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index cde40fb..f34710f 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -323,9 +323,9 @@ def test(self, test_data, oname=None, for aug in test_post: pred = aug(pred) dice_metric = self.metric(pred, label) - dice_log.append(1-dice_metric.item()) - prec_log.append(precision(pred, label)) - rec_log.append(recall(pred, label)) + dice_log.append(1-dice_metric.item().detach().cpu().numpy()) + prec_log.append(precision(pred, label).detach().cpu().numpy()) + rec_log.append(recall(pred, label).detach().cpu().numpy()) input_img = input.detach().cpu().numpy() pred_img = pred.detach().cpu().numpy() From 93b60da5ba60e595213f3483de30627e3b2a546a Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Sun, 20 Jun 2021 19:36:20 +0100 Subject: [PATCH 06/62] Split up CLI stuff --- test.py | 135 ++++++++++++++++++++++++++++++++++++++++++++ train.py | 44 --------------- train_classifier.py | 133 +++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 268 insertions(+), 44 deletions(-) create mode 100644 test.py create mode 100644 train_classifier.py diff --git a/test.py b/test.py new file mode 100644 index 0000000..8be99b5 --- /dev/null +++ b/test.py @@ -0,0 +1,135 @@ +import os +import argparse + +parser = argparse.ArgumentParser(prog='test', + description="Test RectAngle model. See list of available arguments for more info.") + +parser.add_argument('--test', + '--te', + metavar='test', + type=str, + action='store', + default=None, + help='Path to test data.') + +parser.add_argument('--ensemble', + '--en', + metavar='ensemble', + type=str, + action='store', + default=None, + help='Number of ensembled models.') + +parser.add_argument('--weights', + '--w', + metavar='weights', + type=str, + nargs='*', + action='store', + default=None, + help='Path to saved model weights.') + +parser.add_argument('--gate', + '--g', + metavar='gate', + type=str, + action='store', + default=None, + help='(Optional) Attention gating.') + +parser.add_argument('--odir', + '--o', + metavar='odir', + type=str, + action='store', + default='./', + help='Path to output folder.') + +parser.add_argument('--depth', + '--d', + metavar='depth', + type=str, + action='store', + default='5', + help='Depth of U-Net architecture used.') + +parser.add_argument('--classifier', + '--c', + metavar='classifier', + type=bool, + action='store', + default=True, + help='Use of classifier for pre-screening. If selected will train without and then perform test without + with.') + +parser.add_argument('--classweights', + '--cw', + metavar='classweights', + type=str, + action='store', + default=True, + help='Path to trained weights for classifier.') + +parser.add_argument('--threshold', + '--th', + metavar='threshold', + type=str, + action='store', + default='0.5', + help='Activation threshold for classifier.') + +parser.add_argument('--seed', + '--s', + metavar='seed', + type=str, + action='store', + default=None, + help='Random seed for training.') + +args = parser.parse_args() + +## convert arguments to useable form +if args.ensemble: + ensemble = int(args.ensemble) +else: + ensemble = None + +## run training +import rectangle as rect +import h5py +import torch +import random +import numpy as np + +# set seeds for repeatable results +if args.seed: + seed = int(args.seed) + torch.manual_seed(seed) + random.seed(seed) + np.random.seed(seed) + +model = [rect.model.networks.UNet(n_layers=int(args.depth), device=device, + gate=args.gate) for e in int(args.ensemble)] + +for n, m in enumerate(model): + m.load_state_dict(torch.load(args.weights[n])) + +class_model = rect.model.networks.MakeDenseNet(freeze_weights=False).to(device) + +class_model.load_state_dict(torch.load(args.classweights)) + +if torch.cuda.is_available(): + device = torch.device('cuda') + torch.backends.cudnn.benchmark = True +else: + device = torch.device('cpu') + +f_test = h5py.File(args.test, 'r') +if args.classifier: + train_data = rect.utils.io.PreScreenLoader(class_model.eval(), f_test, label=args.label, threshold=float(args.thresh)) +else: + test_data = rect.utils.io.H5DataLoader(f_test, label='vote') + +trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir) + +trainer.test(test_data, test_pre=[rect.utils.transforms.z_score()], + test_post=[rect.utils.transforms.Binary(), rect.utils.transforms.KeepLargestComponent()]) diff --git a/train.py b/train.py index 973a6d6..5ca5ee4 100644 --- a/train.py +++ b/train.py @@ -1,5 +1,3 @@ -## CLI for running training - import os import argparse @@ -22,14 +20,6 @@ default=None, help='Path to validation data.') -parser.add_argument('--test', - '--te', - metavar='test', - type=str, - action='store', - default=None, - help='Path to test data.') - parser.add_argument('--label', '--l', metavar='label', @@ -86,14 +76,6 @@ default='32', help='Batch size. Note images are large (~400x~300).') -parser.add_argument('--classifier', - '--c', - metavar='classifier', - type=bool, - action='store', - default=True, - help='Use of classifier for pre-screening. If selected will train without and then perform test without + with.') - parser.add_argument('--seed', '--s', metavar='seed', @@ -154,29 +136,3 @@ else: trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(), rect.utils.transforms.SpeckleNoise()], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) - -if args.test: - trainer.test(test_data, test_pre=[rect.utils.transforms.z_score()], - test_post=[rect.utils.transforms.Binary(), rect.utils.transforms.KeepLargestComponent()]) - -if args.classifier: - class_train_data = rect.utils.io.ClassifyDataLoader(f_train) - if args.val: - class_val_data = rect.utils.io.ClassifyDataLoader(f_val) - else: - class_val_data = None - - class_model = rect.model.networks.MakeDenseNet(freeze_weights=False).to(device) - class_trainer = rect.utils.train.ClassTrainer(class_model, outdir=os.path.join(args.odir, 'classlogs'), - ensemble=None, early_stop=1000) - - class_trainer.train(class_train_data, class_val_data, train_batch=int(args.batch)) - - threshRange = np.linspace(0, 0.6, 20) - - if args.test: - for i, thresh in enumerate(threshRange): - test_screen_data = rect.utils.io.PreScreenLoader(class_model.eval(), f_test, label='vote', threshold = thresh) - trainer.test(test_screen_data, test_pre=[rect.utils.transforms.z_score()], - test_post=[rect.utils.transforms.Binary(), rect.utils.transforms.KeepLargestComponent()], oname='class_thresh_{}'.format(i)) - diff --git a/train_classifier.py b/train_classifier.py new file mode 100644 index 0000000..f20ed32 --- /dev/null +++ b/train_classifier.py @@ -0,0 +1,133 @@ +import os +import argparse + +parser = argparse.ArgumentParser(prog='train', + description="Train RectAngle model. See list of available arguments for more info.") + +parser.add_argument('--train', + '--tr', + metavar='train', + type=str, + action='store', + default='./miccai_us_data/train.h5', + help='Path to training data. Note that for ensemble this should include train + val pre-split.') + +parser.add_argument('--val', + '--v', + metavar='val', + type=str, + action='store', + default=None, + help='Path to validation data.') + +parser.add_argument('--test', + '--te', + metavar='test', + type=str, + action='store', + default=None, + help='Path to test data.') + +parser.add_argument('--ensemble', + '--en', + metavar='ensemble', + type=str, + action='store', + default=None, + help='Number of ensembled models.') + +parser.add_argument('--freeze', + '--f', + metavar='freeze', + type=bool, + action='store', + default=False, + help='Freeze CNN weights (pre-trained on ImageNet).') + +parser.add_argument('--odir', + '--o', + metavar='odir', + type=str, + action='store', + default='./', + help='Path to output folder.') + +parser.add_argument('--epochs', + '--ep', + metavar='epochs', + type=str, + action='store', + default='200', + help='Max number of training epochs per model.') + +parser.add_argument('--batch', + '--b', + metavar='batch', + type=str, + action='store', + default='32', + help='Batch size. Note images are large (~400x~300).') + +parser.add_argument('--seed', + '--s', + metavar='seed', + type=str, + action='store', + default=None, + help='Random seed for training.') + + +args = parser.parse_args() + +## convert arguments to useable form +if args.ensemble: + ensemble = int(args.ensemble) +else: + ensemble = None + +## run training +import rectangle as rect +import h5py +import torch +import random +import numpy as np + +# set seeds for repeatable results +if args.seed: + seed = int(args.seed) + torch.manual_seed(seed) + random.seed(seed) + np.random.seed(seed) + +f_train = h5py.File(args.train, 'r') +train_data = rect.utils.io.H5DataLoader(f_train, label=args.label) + +if torch.cuda.is_available(): + device = torch.device('cuda') + torch.backends.cudnn.benchmark = True +else: + device = torch.device('cpu') + +if args.val: + f_val = h5py.File(args.val, 'r') + val_data = rect.utils.io.H5DataLoader(f_val, label='vote') + +if args.test: + f_test = h5py.File(args.test, 'r') + test_data = rect.utils.io.H5DataLoader(f_test, label='vote') + +class_train_data = rect.utils.io.ClassifyDataLoader(f_train) +if args.val: + class_val_data = rect.utils.io.ClassifyDataLoader(f_val) +else: + class_val_data = None +if args.test: + class_test_data = rect.utils.io.ClassifyDataLoader(f_test) +else: + class_test_data = None + +class_model = rect.model.networks.MakeDenseNet(freeze_weights=args.freeze).to(device) +class_trainer = rect.utils.train.ClassTrainer(class_model, outdir=os.path.join(args.odir), + ensemble=ensemble, nb_epochs=int(args.epochs)) + +class_trainer.train(class_train_data, class_val_data, train_batch=int(args.batch)) From a6304e26b7841a9c04df4b9afa6d09ef67235ae4 Mon Sep 17 00:00:00 2001 From: Iani Date: Sun, 20 Jun 2021 23:35:57 +0100 Subject: [PATCH 07/62] Modified ClassifyDataLoader, but saved as v2 for proof-checking --- src/rectangle/utils/io.py | 50 ++++++++++++++++++++++++++++++++++++++- 1 file changed, 49 insertions(+), 1 deletion(-) diff --git a/src/rectangle/utils/io.py b/src/rectangle/utils/io.py index 8f8dfe3..66f6236 100644 --- a/src/rectangle/utils/io.py +++ b/src/rectangle/utils/io.py @@ -164,6 +164,55 @@ def __getitem__(self, index): label = torch.tensor([0.0]) return(image, label) +class ClassifyDataLoader_v2(torch.utils.data.Dataset): + + def __init__(self, file, keys=None): + """ Dataloader for hdf5 files, with labels converted to classifier labels + Input arguments: + file : h5py File object + Loaded using h5py.File(path : string) + keys : list, default = None + Keys from h5py file to use. Useful for train-val-test split. + If None, keys generated from entire file. + """ + + super().__init__() + + self.file = file + if not keys: + keys = list(file.keys()) + + self.split_keys = [key.split('_') for key in keys] + start_subj = int(self.split_keys[0][1]) + last_subj = int(self.split_keys[-1][1]) + self.num_subjects = (last_subj - start_subj)+ 1 #Add 1 to account for 0 idx python + self.subjects = [key[1] for key in self.split_keys if key[0] == 'frame'] + #self.subjects = np.linspace(start_subj, last_subj, + # self.num_subjects+1, dtype=int) + + def __len__(self): + return self.num_subjects + + def __getitem__(self, index): + + subj_ix = self.subjects[index] + image_key = 'frame_' + subj_ix + image = torch.unsqueeze(torch.tensor(self.file[image_key][()].astype('float32')), dim=0) + + label_batch = torch.cat([torch.unsqueeze(torch.tensor( + self.file[f'label_{subj_ix}_0{label_ix}' ] + [()].astype('float32')), dim=0) for label_ix in range(3)]) + + label_vote = torch.sum(label_batch, dim=(1,2)) + sum_vote = torch.sum(label_vote != 0) + + #print(sum_vote) + if sum_vote >= 2: + label = torch.tensor([1.0]) + else: + label = torch.tensor([0.0]) + + return(image, label) class TestPlotLoader(torch.utils.data.Dataset): def __init__(self, file, keys=None, label='vote'): @@ -231,7 +280,6 @@ def __getitem__(self, index): label = torch.unsqueeze(torch.mean(label_batch, dim=0), dim=0) return(image, label) - class PreScreenLoader(torch.utils.data.Dataset): def __init__(self, model, file, keys=None, label='random', threshold=0.5): """ Dataloader for hdf5 files. From 8fe2026aa3ec10ddaa07bb4c39faa239c6427186 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 10:51:58 +0100 Subject: [PATCH 08/62] Remove unused args in cli tools --- train.py | 4 ---- train_classifier.py | 7 ++----- 2 files changed, 2 insertions(+), 9 deletions(-) diff --git a/train.py b/train.py index 5ca5ee4..2b1a1d9 100644 --- a/train.py +++ b/train.py @@ -120,10 +120,6 @@ f_val = h5py.File(args.val, 'r') val_data = rect.utils.io.H5DataLoader(f_val, label='vote') -if args.test: - f_test = h5py.File(args.test, 'r') - test_data = rect.utils.io.H5DataLoader(f_test, label='vote') - model = rect.model.networks.UNet(n_layers=int(args.depth), device=device, gate=args.gate) diff --git a/train_classifier.py b/train_classifier.py index f20ed32..7e8d0cf 100644 --- a/train_classifier.py +++ b/train_classifier.py @@ -99,22 +99,19 @@ random.seed(seed) np.random.seed(seed) -f_train = h5py.File(args.train, 'r') -train_data = rect.utils.io.H5DataLoader(f_train, label=args.label) - if torch.cuda.is_available(): device = torch.device('cuda') torch.backends.cudnn.benchmark = True else: device = torch.device('cpu') +f_train = h5py.File(args.train, 'r') + if args.val: f_val = h5py.File(args.val, 'r') - val_data = rect.utils.io.H5DataLoader(f_val, label='vote') if args.test: f_test = h5py.File(args.test, 'r') - test_data = rect.utils.io.H5DataLoader(f_test, label='vote') class_train_data = rect.utils.io.ClassifyDataLoader(f_train) if args.val: From 182e3ef65a2b25a6841842986c2ae5a377a6a3d3 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 11:16:26 +0100 Subject: [PATCH 09/62] change class loader to Iani new class --- train_classifier.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/train_classifier.py b/train_classifier.py index 7e8d0cf..469924e 100644 --- a/train_classifier.py +++ b/train_classifier.py @@ -113,13 +113,13 @@ if args.test: f_test = h5py.File(args.test, 'r') -class_train_data = rect.utils.io.ClassifyDataLoader(f_train) +class_train_data = rect.utils.io.ClassifyDataLoader_v2(f_train) if args.val: - class_val_data = rect.utils.io.ClassifyDataLoader(f_val) + class_val_data = rect.utils.io.ClassifyDataLoader_v2(f_val) else: class_val_data = None if args.test: - class_test_data = rect.utils.io.ClassifyDataLoader(f_test) + class_test_data = rect.utils.io.ClassifyDataLoader_v2(f_test) else: class_test_data = None From 2b09a3aa817865cadabc8048380d318c2f11d5bb Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 11:46:51 +0100 Subject: [PATCH 10/62] Add lr_schedule to CLI train --- train.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/train.py b/train.py index 2b1a1d9..773d07e 100644 --- a/train.py +++ b/train.py @@ -44,6 +44,14 @@ default=None, help='(Optional) Attention gating.') +parser.add_argument('--lr_schedule', + '--lrs', + metavar='lr_schedule', + type=str, + action='store', + default=None, + help="Method for scheduling of learning rate. {None, 'lambda', 'exponential', 'reduce_on_plateau'}") + parser.add_argument('--odir', '--o', metavar='odir', @@ -124,7 +132,7 @@ gate=args.gate) trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, - nb_epochs=int(args.epochs)) + nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule) if args.val: trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(), rect.utils.transforms.SpeckleNoise()], From 2db8bfee73d1039cf74a63b05a5513f78a234988 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 11:48:00 +0100 Subject: [PATCH 11/62] Set tensorboard logs to outdir/runs --- src/rectangle/utils/train.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index ad9a75a..b1e63a0 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -31,7 +31,7 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', self.ensemble = ensemble self.outdir = outdir self.device = device - self.writer = SummaryWriter() + self.writer = SummaryWriter(path.join(outdir,'runs')) if self.ensemble == 0: self.ensemble = None From 226cf260d6cd7f9d0d19f2f8d56bcb98b9bf7212 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 11:48:39 +0100 Subject: [PATCH 12/62] Define log_dir explicitly in summarywriter --- src/rectangle/utils/train.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index b1e63a0..f4df8de 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -31,7 +31,7 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', self.ensemble = ensemble self.outdir = outdir self.device = device - self.writer = SummaryWriter(path.join(outdir,'runs')) + self.writer = SummaryWriter(log_dir=path.join(outdir,'runs')) if self.ensemble == 0: self.ensemble = None From bd6bd2ed9eb3a2218a55b842797019d5f9d119ee Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 11:51:03 +0100 Subject: [PATCH 13/62] Add summarywriter init to ClassTrainer --- src/rectangle/utils/train.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index f4df8de..18f8404 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -499,6 +499,8 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', self.opt = opt + self.writer = SummaryWriter(log_dir=path.join(outdir,'runs')) + def train(self, train_data, val_data=None, oname=None, train_pre=None, train_post=None, train_batch=128, train_shuffle=True, From e9bda1d1e0c94faee58e2bb4f6b9f301bfb731e2 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 14:52:01 +0100 Subject: [PATCH 14/62] add device def to trainer in CLI --- train.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/train.py b/train.py index 773d07e..68d2fb2 100644 --- a/train.py +++ b/train.py @@ -131,7 +131,7 @@ model = rect.model.networks.UNet(n_layers=int(args.depth), device=device, gate=args.gate) -trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, +trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, device=device, nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule) if args.val: From 3154147dd9b48bb8dc7fa161c79e618f218a54dc Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 21 Jun 2021 15:14:48 +0100 Subject: [PATCH 15/62] Add device def to other CLI scripts --- test.py | 2 +- train_classifier.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/test.py b/test.py index 8be99b5..ae8b7a9 100644 --- a/test.py +++ b/test.py @@ -129,7 +129,7 @@ else: test_data = rect.utils.io.H5DataLoader(f_test, label='vote') -trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir) +trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, device=device) trainer.test(test_data, test_pre=[rect.utils.transforms.z_score()], test_post=[rect.utils.transforms.Binary(), rect.utils.transforms.KeepLargestComponent()]) diff --git a/train_classifier.py b/train_classifier.py index 469924e..076a95d 100644 --- a/train_classifier.py +++ b/train_classifier.py @@ -125,6 +125,6 @@ class_model = rect.model.networks.MakeDenseNet(freeze_weights=args.freeze).to(device) class_trainer = rect.utils.train.ClassTrainer(class_model, outdir=os.path.join(args.odir), - ensemble=ensemble, nb_epochs=int(args.epochs)) + ensemble=ensemble, nb_epochs=int(args.epochs), device=device) class_trainer.train(class_train_data, class_val_data, train_batch=int(args.batch)) From a356cc68c5d0ade744e7b574ac8709b0c132fb15 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 22 Jun 2021 11:09:26 +0100 Subject: [PATCH 16/62] Add tensorboard compatibility with ensemble --- src/rectangle/utils/train.py | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index 18f8404..a7a2b75 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -31,7 +31,6 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', self.ensemble = ensemble self.outdir = outdir self.device = device - self.writer = SummaryWriter(log_dir=path.join(outdir,'runs')) if self.ensemble == 0: self.ensemble = None @@ -47,6 +46,10 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', for i in range(self.ensemble): self.model_ensemble.append(deepcopy(model)) + self.writer = [SummaryWriter(log_dir=path.join(outdir,'runs/model_{}'.format(i))) for i in range(self.ensemble)] + else: + self.writer = SummaryWriter(log_dir=path.join(outdir,'runs')) + if opt == 'adam': if self.ensemble: opt = [Adam(model.parameters()) for model in self.model_ensemble] @@ -103,6 +106,7 @@ def train(self, train_data, val_data=None, oname=None, train = train_list[i] val = val_list[i] opt_ = self.opt[i] + writer_ = self.writer[i] if self.lr_schedule: lr_schedule_ = self.lr_schedule[i] if lr_schedule_ == 'lambda': @@ -139,8 +143,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log_ensemble[i,epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: - self.writer.add_scalar('train/dice_loss_ensemble', loss_, epoch) - self.writer.add_scalar('train/dice_coefficient_ensemble', 1-loss_, epoch) + writer_.add_scalar('train/dice_loss_ensemble', loss_, epoch) + writer_.add_scalar('train/dice_coefficient_ensemble', 1-loss_, epoch) print('Epoch #{}: Mean Dice Loss: {}'.format(epoch, loss_log_ensemble[i,epoch])) if epoch % self.val_interval == 0: dice_epoch = [] @@ -161,8 +165,8 @@ def train(self, train_data, val_data=None, oname=None, dice_epoch.append(1 - dice_metric.item()) dice_log_ensemble[i,int(epoch//self.val_interval)] = np.nanmean(dice_epoch) - self.writer.add_scalar('val/dice_loss_ensemble', dice_metric, epoch) - self.writer.add_scalar('val/dice_coefficient_ensemble', 1-dice_metric, epoch) + writer_.add_scalar('val/dice_loss_ensemble', dice_metric, epoch) + writer_.add_scalar('val/dice_coefficient_ensemble', 1-dice_metric, epoch) ## show some (e.g.,10) example images in tensorboard ex_num = 10 @@ -172,7 +176,7 @@ def train(self, train_data, val_data=None, oname=None, ex_images = ex_image.reshape(-1,ex_image.shape[2]) image_grid = (make_grid(ex_images, nrow=ex_num)[0]+0.5)/ex_num - self.writer.add_images( + writer_.add_images( "val/example_images_ensemble", image_grid, epoch, From a8e72c82f923a8cfc92cb80dff5360b314363eb7 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 22 Jun 2021 12:52:51 +0100 Subject: [PATCH 17/62] Set lr_schedule to every epoch --- src/rectangle/utils/train.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index a7a2b75..2536800 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -16,8 +16,8 @@ class Trainer(nn.Module): def __init__(self, model, nb_epochs=200, outdir='./logs', loss=DiceLoss(), metric=DiceLoss(), opt='adam', - print_interval=1, val_interval=5, device='cuda', - early_stop=5, lr_schedule=None, ensemble=None): + print_interval=1, val_interval=1, device='cuda', + early_stop=10, lr_schedule=None, ensemble=None): super().__init__() @@ -138,9 +138,9 @@ def train(self, train_data, val_data=None, oname=None, loss_ = self.loss(pred, label) loss_.backward() opt_.step() - if lr_schedule_ and lr_schedule_ != 'reduce_on_plateau': - lr_schedule_.step() loss_epoch.append(loss_.item()) + if lr_schedule_ and lr_schedule_ != 'reduce_on_plateau': + lr_schedule_.step() loss_log_ensemble[i,epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: writer_.add_scalar('train/dice_loss_ensemble', loss_, epoch) @@ -160,10 +160,10 @@ def train(self, train_data, val_data=None, oname=None, for aug in val_post: pred = aug(pred) dice_metric = self.metric(pred, label) - if lr_schedule_ == 'reduce_on_plateau': - lr_schedule_.step(dice_metric) dice_epoch.append(1 - dice_metric.item()) dice_log_ensemble[i,int(epoch//self.val_interval)] = np.nanmean(dice_epoch) + if lr_schedule_ == 'reduce_on_plateau': + lr_schedule_.step(1-np.nanmean(dice_epoch)) writer_.add_scalar('val/dice_loss_ensemble', dice_metric, epoch) writer_.add_scalar('val/dice_coefficient_ensemble', 1-dice_metric, epoch) From c1ecd531fb48ee645f8a7fbb3d240086bc609fd9 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 09:41:56 +0100 Subject: [PATCH 18/62] Add modification to handle affine and flip --- src/rectangle/utils/train.py | 63 ++++++++++++++++++++++++++++++------ 1 file changed, 54 insertions(+), 9 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index 2536800..8049449 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -130,7 +130,12 @@ def train(self, train_data, val_data=None, oname=None, opt_.zero_grad() if train_pre: for aug in train_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if train_post: for aug in train_post: @@ -154,7 +159,12 @@ def train(self, train_data, val_data=None, oname=None, input, label = input.to(self.device), label.to(self.device) if val_pre: for aug in val_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if val_post: for aug in val_post: @@ -224,7 +234,12 @@ def train(self, train_data, val_data=None, oname=None, self.opt.zero_grad() if train_pre: for aug in train_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if train_post: for aug in train_post: @@ -248,7 +263,12 @@ def train(self, train_data, val_data=None, oname=None, input, label = input.to(self.device), label.to(self.device) if val_pre: for aug in val_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if val_post: for aug in val_post: @@ -361,7 +381,12 @@ def test(self, test_data, oname=None, input, label = input.to(self.device), label.to(self.device) if test_pre: for aug in test_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) if self.ensemble: pred = [model(input) for model in self.model_ensemble] pred = torch.cat(pred, dim=0) @@ -557,7 +582,12 @@ def train(self, train_data, val_data=None, oname=None, opt_.zero_grad() if train_pre: for aug in train_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if train_post: for aug in train_post: @@ -579,7 +609,12 @@ def train(self, train_data, val_data=None, oname=None, input, label = input.to(self.device), label.to(self.device) if val_pre: for aug in val_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if val_post: for aug in val_post: @@ -622,7 +657,12 @@ def train(self, train_data, val_data=None, oname=None, self.opt.zero_grad() if train_pre: for aug in train_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if train_post: for aug in train_post: @@ -644,7 +684,12 @@ def train(self, train_data, val_data=None, oname=None, input, label = input.to(self.device), label.to(self.device) if val_pre: for aug in val_pre: - input = aug(input) + if aug.__class__.__name__ == 'Flip' or 'Affine': + input = torch.stack([input, label]) + input = aug(input) + input, label = torch.chunk(input, 2) + else: + input = aug(input) pred = model(input) if val_post: for aug in val_post: From f366fd9d6e24c7e7accc02195eb93a9e2c32718c Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 10:06:52 +0100 Subject: [PATCH 19/62] Drop empty dim in input,label after affine --- src/rectangle/utils/train.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index 8049449..ef7b399 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -134,6 +134,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -163,6 +164,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -238,6 +240,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -267,6 +270,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -385,6 +389,7 @@ def test(self, test_data, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) if self.ensemble: @@ -586,6 +591,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -613,6 +619,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -661,6 +668,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -688,6 +696,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.stack([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + input, label = input[0], label[0] else: input = aug(input) pred = model(input) From 612a5a01a64eb9048d6043740e05d1525d794a6b Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 10:30:44 +0100 Subject: [PATCH 20/62] Re-try for fix of affine forward --- src/rectangle/utils/train.py | 26 +++++++++----------------- 1 file changed, 9 insertions(+), 17 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index ef7b399..3c436c1 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -131,10 +131,9 @@ def train(self, train_data, val_data=None, oname=None, if train_pre: for aug in train_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -161,10 +160,9 @@ def train(self, train_data, val_data=None, oname=None, if val_pre: for aug in val_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -237,10 +235,9 @@ def train(self, train_data, val_data=None, oname=None, if train_pre: for aug in train_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -267,7 +264,7 @@ def train(self, train_data, val_data=None, oname=None, if val_pre: for aug in val_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) input, label = input[0], label[0] @@ -386,10 +383,9 @@ def test(self, test_data, oname=None, if test_pre: for aug in test_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) if self.ensemble: @@ -588,10 +584,9 @@ def train(self, train_data, val_data=None, oname=None, if train_pre: for aug in train_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -616,10 +611,9 @@ def train(self, train_data, val_data=None, oname=None, if val_pre: for aug in val_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -665,10 +659,9 @@ def train(self, train_data, val_data=None, oname=None, if train_pre: for aug in train_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) @@ -693,10 +686,9 @@ def train(self, train_data, val_data=None, oname=None, if val_pre: for aug in val_pre: if aug.__class__.__name__ == 'Flip' or 'Affine': - input = torch.stack([input, label]) + input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) From 54819632992b67dd08b8f7b454d48ad0290154b0 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 11:14:27 +0100 Subject: [PATCH 21/62] Remove indexing line (prev commit) --- src/rectangle/utils/train.py | 1 - 1 file changed, 1 deletion(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index 3c436c1..f0e0e6c 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -267,7 +267,6 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) - input, label = input[0], label[0] else: input = aug(input) pred = model(input) From a9bfb9a8715d1149275e4f2ff156d4c40e15f8c8 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 12:41:34 +0100 Subject: [PATCH 22/62] Binarise label after applying warps --- src/rectangle/utils/train.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index f0e0e6c..b94d922 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -1,6 +1,7 @@ import torch from torch import nn from rectangle.utils.metrics import DiceLoss, Precision, Recall, Accuracy +from rectangle.utils.transforms import Binary from torch.optim import Adam from copy import deepcopy from os import path, makedirs @@ -31,6 +32,7 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', self.ensemble = ensemble self.outdir = outdir self.device = device + self.bin = Binary() if self.ensemble == 0: self.ensemble = None @@ -134,6 +136,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) @@ -163,6 +166,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) @@ -238,6 +242,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) @@ -267,6 +272,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) @@ -385,6 +391,7 @@ def test(self, test_data, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) if self.ensemble: @@ -586,6 +593,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) @@ -613,6 +621,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) @@ -661,6 +670,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) @@ -688,6 +698,7 @@ def train(self, train_data, val_data=None, oname=None, input = torch.cat([input, label]) input = aug(input) input, label = torch.chunk(input, 2) + label = self.bin(label) else: input = aug(input) pred = model(input) From cb1ba532a8a4f2eeb843d7c34969ace72dabd6fb Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 12:48:31 +0100 Subject: [PATCH 23/62] Fix binarise in dice loss --- src/rectangle/utils/metrics.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/rectangle/utils/metrics.py b/src/rectangle/utils/metrics.py index 6734015..4163d52 100644 --- a/src/rectangle/utils/metrics.py +++ b/src/rectangle/utils/metrics.py @@ -29,7 +29,7 @@ def forward(self, inputs, targets): # Seems to perform very well without binary - soft dice? if not self.soft: - inputs = BinaryDice(inputs, self.threshold) + inputs = self.BinaryDice(inputs, self.threshold) inputs = inputs.view(-1).float() targets = targets.view(-1).float() From 3326aa8b1fead8fb297d3025f01e6fea647a90b9 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 13:25:40 +0100 Subject: [PATCH 24/62] Add qsub scripts --- scripts/baseline_mean.qsub.sh | 23 +++++++++++++++++++++++ scripts/baseline_random.qsub.sh | 23 +++++++++++++++++++++++ scripts/baseline_vote.qsub.sh | 23 +++++++++++++++++++++++ 3 files changed, 69 insertions(+) create mode 100644 scripts/baseline_mean.qsub.sh create mode 100644 scripts/baseline_random.qsub.sh create mode 100644 scripts/baseline_vote.qsub.sh diff --git a/scripts/baseline_mean.qsub.sh b/scripts/baseline_mean.qsub.sh new file mode 100644 index 0000000..318a674 --- /dev/null +++ b/scripts/baseline_mean.qsub.sh @@ -0,0 +1,23 @@ +#$ -S /bin/bash +#$ -l tmem=32G +#$ -l h_vmem=32G +#$ -l h_rt=40:00:00 + +#$ -l gpu=true +#$ -N baseline + +#$ -cwd + +#module purge +#module load default/python/3.8.5 +source /share/apps/source_files/python/python-3.8.5.source +source rectenv/bin/activate + +python ./RectAngle/train.py --train ./miccai_us_data/train.h5 \ +--val ./miccai_us_data/val.h5 \ +--ensemble 5 \ +--lr_schedule exponential \ +--label mean \ +--odir ./baseline_data/mean \ +--epochs 50 \ +--seed 0 diff --git a/scripts/baseline_random.qsub.sh b/scripts/baseline_random.qsub.sh new file mode 100644 index 0000000..9e8179f --- /dev/null +++ b/scripts/baseline_random.qsub.sh @@ -0,0 +1,23 @@ +#$ -S /bin/bash +#$ -l tmem=32G +#$ -l h_vmem=32G +#$ -l h_rt=40:00:00 + +#$ -l gpu=true +#$ -N baseline + +#$ -cwd + +#module purge +#module load default/python/3.8.5 +source /share/apps/source_files/python/python-3.8.5.source +source rectenv/bin/activate + +python ./RectAngle/train.py --train ./miccai_us_data/train.h5 \ +--val ./miccai_us_data/val.h5 \ +--ensemble 5 \ +--lr_schedule exponential \ +--label random \ +--odir ./baseline_data/random \ +--epochs 50 \ +--seed 0 diff --git a/scripts/baseline_vote.qsub.sh b/scripts/baseline_vote.qsub.sh new file mode 100644 index 0000000..6b11ca8 --- /dev/null +++ b/scripts/baseline_vote.qsub.sh @@ -0,0 +1,23 @@ +#$ -S /bin/bash +#$ -l tmem=32G +#$ -l h_vmem=32G +#$ -l h_rt=40:00:00 + +#$ -l gpu=true +#$ -N baseline + +#$ -cwd + +#module purge +#module load default/python/3.8.5 +source /share/apps/source_files/python/python-3.8.5.source +source rectenv/bin/activate + +python ./RectAngle/train.py --train ./miccai_us_data/train.h5 \ +--val ./miccai_us_data/val.h5 \ +--ensemble 5 \ +--label vote \ +--lr_schedule exponential \ +--odir ./baseline_data/vote \ +--epochs 50 \ +--seed 0 From 81821f38e476d27bd65912133f30d0a77d0cc3d4 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 24 Jun 2021 15:39:34 +0100 Subject: [PATCH 25/62] Add TODO of Iani suggestion --- src/rectangle/utils/transforms.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/rectangle/utils/transforms.py b/src/rectangle/utils/transforms.py index dbe1ebe..69ec712 100644 --- a/src/rectangle/utils/transforms.py +++ b/src/rectangle/utils/transforms.py @@ -336,6 +336,7 @@ def __init__(self, threshold=0.5): super().__init__() self.threshold = threshold + # TODO: check for torch auto binarising def __call__(self, image): return (image > self.threshold).int() From c4f0bf2842450be2cc35c7c652a8e1351cdada6c Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 16:54:58 +0100 Subject: [PATCH 26/62] Added print("cuda available" to check GPU is used --- train.py | 1 + 1 file changed, 1 insertion(+) diff --git a/train.py b/train.py index 68d2fb2..88b4026 100644 --- a/train.py +++ b/train.py @@ -121,6 +121,7 @@ if torch.cuda.is_available(): device = torch.device('cuda') torch.backends.cudnn.benchmark = True + print("Cuda available!") else: device = torch.device('cpu') From cd262238187dfe097bb57c7d6e14c33a2209b795 Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 17:06:25 +0100 Subject: [PATCH 27/62] Added print statements to check code runs on CANDI --- train.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/train.py b/train.py index 88b4026..225a292 100644 --- a/train.py +++ b/train.py @@ -108,6 +108,7 @@ import random import numpy as np +print("Code running") # set seeds for repeatable results if args.seed: seed = int(args.seed) @@ -124,6 +125,7 @@ print("Cuda available!") else: device = torch.device('cpu') + print("Using CPU!") if args.val: f_val = h5py.File(args.val, 'r') From 6b538c952f29f7ef2dcbe671dee235a547e9992c Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 17:10:21 +0100 Subject: [PATCH 28/62] Added cuda_visible_devices to use the GPU --- train.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/train.py b/train.py index 225a292..35ab743 100644 --- a/train.py +++ b/train.py @@ -109,6 +109,8 @@ import numpy as np print("Code running") +os.environ["CUDA_VISIBLE_DEVICES"]="0" + # set seeds for repeatable results if args.seed: seed = int(args.seed) From 84389070243b515771fbc5c51e1756631b71a65f Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 18:39:33 +0100 Subject: [PATCH 29/62] Removed os.environ line to avoid bugs with GPU --- train.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/train.py b/train.py index 35ab743..67ba00e 100644 --- a/train.py +++ b/train.py @@ -109,7 +109,7 @@ import numpy as np print("Code running") -os.environ["CUDA_VISIBLE_DEVICES"]="0" +#os.environ["CUDA_VISIBLE_DEVICES"]="0" # set seeds for repeatable results if args.seed: From c4e39a01215ae9561819332041718918273059c0 Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 18:53:01 +0100 Subject: [PATCH 30/62] Changed augmentation to scale only --- src/rectangle/utils/transforms.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/rectangle/utils/transforms.py b/src/rectangle/utils/transforms.py index 69ec712..55262fa 100644 --- a/src/rectangle/utils/transforms.py +++ b/src/rectangle/utils/transforms.py @@ -81,8 +81,10 @@ def __init__(self, prob=0.3,\ def __call__(self, image): rand_ = random.uniform(0,1) if rand_ < self.prob: - RandAffine_ = RandomAffine(degrees=self.degrees, translate=(self.translate,self.translate), - scale=self.scale, shear=self.shear) + #RandAffine_ = RandomAffine(degrees=self.degrees, translate=(self.translate,self.translate), + # scale=self.scale, shear=self.shear) + RandAffine_ = RandomAffine(scale=self.scale) + image = RandAffine_(image) return image From fbb4ff4ea2e194d5ada2c9573dbab74242a065ff Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 18:53:43 +0100 Subject: [PATCH 31/62] Removed speckle noise in augmentation --- train.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/train.py b/train.py index 67ba00e..4293397 100644 --- a/train.py +++ b/train.py @@ -140,8 +140,10 @@ nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule) if args.val: - trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(), rect.utils.transforms.SpeckleNoise()], + trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(), val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) else: - trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(), rect.utils.transforms.SpeckleNoise()], + trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine()], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) + + \ No newline at end of file From 347c50b34bee5092bae388a97564933138a06321 Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 18:56:45 +0100 Subject: [PATCH 32/62] Removed bug in typed code --- train.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/train.py b/train.py index 4293397..7b5b63b 100644 --- a/train.py +++ b/train.py @@ -140,10 +140,9 @@ nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule) if args.val: - trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(), + trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine()], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) else: trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine()], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) - \ No newline at end of file From d4a3eb8c899b2f3b92cbbd7087d7174026ff5955 Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 19:16:04 +0100 Subject: [PATCH 33/62] Added data augmentation to classification training --- train_classifier.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/train_classifier.py b/train_classifier.py index 076a95d..bc9d670 100644 --- a/train_classifier.py +++ b/train_classifier.py @@ -127,4 +127,5 @@ class_trainer = rect.utils.train.ClassTrainer(class_model, outdir=os.path.join(args.odir), ensemble=ensemble, nb_epochs=int(args.epochs), device=device) -class_trainer.train(class_train_data, class_val_data, train_batch=int(args.batch)) +class_trainer.train(class_train_data, class_val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine()], + val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) From f8c4e2e91e64c0d0ecc457523e2b502bff1c29a6 Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 20:34:09 +0100 Subject: [PATCH 34/62] Added to main to access latest code changes --- prescreening_strategy2.py | 156 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 156 insertions(+) create mode 100644 prescreening_strategy2.py diff --git a/prescreening_strategy2.py b/prescreening_strategy2.py new file mode 100644 index 0000000..99719bc --- /dev/null +++ b/prescreening_strategy2.py @@ -0,0 +1,156 @@ +from numpy.lib.arraysetops import unique +import rectangle as rect +import h5py +import torch +import random +import numpy as np +import os +from rectangle.model.networks import DenseNet as DenseNet + +from torch.utils.data import DataLoader +import tensorflow as tf + +def standardise(image): + + batch_ = image.shape[0] + for batch_iter_ in range(batch_): + image[batch_iter_,...] = (image[batch_iter_,...] - \ + torch.mean(image[batch_iter_,...]) / \ + torch.std(image[batch_iter_,...])) + + return image + +def dice_score2(y_pred, y_true, eps=1e-8): + ''' + y_pred, y_true -> [N, C=1, D, H, W] + ''' + #y_pred[y_pred < 0.5] = 0. + #y_pred[y_pred > 0] = 1. + + #Calculate the number of incorrectly labelled pixels + + numerator = torch.sum(y_true*y_pred, dim=(2,3)) * 2 + denominator = torch.sum(y_true, dim=(2,3)) + torch.sum(y_pred, dim=(2,3)) + eps + return torch.mean(numerator / denominator) + +def dice_fp(y_pred, y_true, pos_frames, neg_frames): + """ A function that computes dice score on positive frames, + and FP pixels on negative frames, based off Yipeng's metrics + """ + dice_ = dice_score2(y_pred[pos_frames, :, :], y_true[pos_frames, :, :]) + fp = torch.sum(y_pred[neg_frames, :, :], dim = [1,2,3]) + + return dice_, fp + + +use_cuda = torch.cuda.is_available() + +### Loading ensemble segmentation network ### +num_ensemble = 5 +path_str = '/Users/iani/Documents/Segmentation_project/ensemble/' +latest_model = ['13.pth', '4.pth', '30.pth', '28.pth', '28.pth'] #Checked manually +model_paths = [os.path.join(path_str, 'model_'+ str(idx), latest_model[idx]) +for idx in range(num_ensemble)] + +depth = 5 +device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') +seg_models = [rect.model.networks.UNet(n_layers=depth, device=device, + gate=None) for e in range(int(num_ensemble))] + +for n, m in enumerate(seg_models): + m.load_state_dict(torch.load(model_paths[n], map_location= device)) + +### Loading classifier network ### +class_model = torch.load("/Users/iani/Documents/Segmentation_project/classification_model", map_location = device) + +### Inference ### + +test_file = h5py.File('/Users/iani/Documents/Reg2Seg/dataset/test.h5', 'r') +test_DS = rect.utils.io.H5DataLoader(test_file) +test_DL = DataLoader(test_DS, batch_size = 8, shuffle = False) + +segmentation_threshold = 0.5 +classification_threshold = 0.5 + + +all_dice_screen = [] +all_dice_noscreen = [] + +all_fp_screen = [] +all_fp_noscreen = [] + +with torch.no_grad(): + + for jj, (images_test, labels_test) in enumerate(test_DL): + + if use_cuda: + images_test, labels_test = images_test.cuda(), labels_test.cuda() + + #Obtain positive and negative frames + positive_frames = [(1 in label) for label in labels_test] + negative_frames = [not(1 in label) for label in labels_test] + + #False positives negative frames + + #Dice score : positive frames + + #Obtain prediction for classifier + class_preds = class_model(images_test) + + #Normalise images for segmentation network + norm_images_test = standardise(images_test) + + #Obtain predictions for each ensemble model and combine them + combined_predictions = torch.zeros_like(labels_test, dtype = float) + majority = len(seg_models) - 1 + + for model_ in seg_models: + #Obtain predictions + model_.eval() + seg_predictions = torch.tensor(model_(norm_images_test) > 0.5, dtype = float) + combined_predictions += seg_predictions + + #All segmentation results - only on positive frames + combined_predictions = (combined_predictions >= majority) #Majority vote + dice_noscreen, fp_noscreen = dice_fp(combined_predictions, labels_test, positive_frames, negative_frames) + all_dice_noscreen.append(dice_noscreen) + all_fp_noscreen.append(fp_noscreen) + + + #dice_noscreen = dice_score(combined_predictions, labels_test) + #all_dice_noscreen.append(dice_noscreen) + + #Pre-screened results only + prostate_idx = np.where(class_preds == 1)[0] + #dice_screened = dice_score(combined_predictions[prostate_idx, :,:], labels_test[prostate_idx, :,:]) + + positive_frames_screened = [positive_frames[i] for i in prostate_idx] + negative_frames_screened = [negative_frames[i] for i in prostate_idx] + + dice_screen, fp_screen = dice_fp(combined_predictions[prostate_idx, :,:], labels_test[prostate_idx, :,:], positive_frames_screened, negative_frames_screened) + all_dice_screen.append(dice_screen) + all_fp_screen.append(fp_screen) + + print(f"Dice scores: Not-screened : {dice_noscreen} | Screened : {dice_screen}") + print(f"FP scores: Not-screened : {fp_noscreen} | Screened : {fp_screen}") + + +#Obtaining plots of the histogram + +#Obtain all unique FP scores for screen, no screen method +unique_fp_screen = [np.unique(fp_vals) for fp_vals in all_fp_screen if len(fp_vals) > 0] +unique_fp_screen = np.concatenate(unique_fp_screen, axis = 0) + +#Obtain all unique FP scores for screen, no screen method +unique_fp_noscreen = [np.unique(fp_vals) for fp_vals in all_fp_noscreen if len(fp_vals) > 0] +unique_fp_noscreen = np.concatenate(unique_fp_noscreen, axis = 0) + +from matplotlib import pyplot as plt +plt.hist(unique_fp_noscreen, label = "noscreen") +plt.hist(unique_fp_screen, label = "screen") +plt.xlabel("Number of FP pixels per negative segmented frame") +plt.legend() +plt.show() + +print('Chicken') + From b1cadf5342ecad77da81d34d8ba13ad7a851b6e3 Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 21:52:21 +0100 Subject: [PATCH 35/62] Revert randomafine back to include all --- src/rectangle/utils/transforms.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/src/rectangle/utils/transforms.py b/src/rectangle/utils/transforms.py index 55262fa..69ec712 100644 --- a/src/rectangle/utils/transforms.py +++ b/src/rectangle/utils/transforms.py @@ -81,10 +81,8 @@ def __init__(self, prob=0.3,\ def __call__(self, image): rand_ = random.uniform(0,1) if rand_ < self.prob: - #RandAffine_ = RandomAffine(degrees=self.degrees, translate=(self.translate,self.translate), - # scale=self.scale, shear=self.shear) - RandAffine_ = RandomAffine(scale=self.scale) - + RandAffine_ = RandomAffine(degrees=self.degrees, translate=(self.translate,self.translate), + scale=self.scale, shear=self.shear) image = RandAffine_(image) return image From 16a9903755903fa4c250b894c5e740327d87cf97 Mon Sep 17 00:00:00 2001 From: Iani Date: Thu, 24 Jun 2021 21:53:42 +0100 Subject: [PATCH 36/62] Change affine in train to scale and rotation only --- train.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/train.py b/train.py index 7b5b63b..7825110 100644 --- a/train.py +++ b/train.py @@ -143,6 +143,6 @@ trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine()], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) else: - trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine()], + trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(prob = 0.3, scale = (0.9,1.1), degrees = 5, shear = 0, translate = 0)], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) From 4456ffe04d9231e706fbbebde21762084877f128 Mon Sep 17 00:00:00 2001 From: Jiongqi Date: Fri, 25 Jun 2021 15:17:35 +0800 Subject: [PATCH 37/62] add labelling method: combination+manually adjust percentage --- src/rectangle/utils/io.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/src/rectangle/utils/io.py b/src/rectangle/utils/io.py index 66f6236..98028e4 100644 --- a/src/rectangle/utils/io.py +++ b/src/rectangle/utils/io.py @@ -84,6 +84,8 @@ def __init__(self, file, keys=None, label='random'): self.subjects = np.linspace(start_subj, last_subj, self.num_subjects+1, dtype=int) self.label = label + if label.split('_')[0] == 'combine': + self.label_loop = label def __len__(self): return self.num_subjects @@ -94,6 +96,14 @@ def __getitem__(self, index): image = torch.unsqueeze(torch.tensor( self.file['frame_%05d' % (subj_ix, )][()].astype('float32')), dim=0) + + if self.label_loop.split('_')[0] == 'combine': + label_percent = int(self.label_loop.split('_')[1]) + if index < int(self.num_subjects*label_percent/100): + self.label = 'vote' + else: + self.label = 'random' + if self.label == 'random': label = torch.unsqueeze(torch.tensor( self.file['label_%05d_%02d' % (subj_ix, From 4467b28caee729195b33d811b1de1998139cb214 Mon Sep 17 00:00:00 2001 From: Iani Date: Fri, 25 Jun 2021 09:40:12 +0100 Subject: [PATCH 38/62] Changed AffineTransform to rotation of 5 degrees only --- train.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/train.py b/train.py index 7825110..b076ce6 100644 --- a/train.py +++ b/train.py @@ -1,6 +1,8 @@ import os import argparse +from rectangle.utils.transforms import Affine + parser = argparse.ArgumentParser(prog='train', description="Train RectAngle model. See list of available arguments for more info.") @@ -139,10 +141,13 @@ trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, device=device, nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule) +#Manually setting Affine Transforms +AffineTransform = rect.utils.transforms.Affine(prob = 0.3, scale = None, degrees = 5, shear = None, translate = None) + if args.val: - trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine()], + trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), AffineTransform], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) else: - trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), rect.utils.transforms.Affine(prob = 0.3, scale = (0.9,1.1), degrees = 5, shear = 0, translate = 0)], + trainer.train(train_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), AffineTransform], val_pre=[rect.utils.transforms.z_score()], train_batch=int(args.batch)) From 62d3174726aa2b9692fafd759f087c84fafbb1d1 Mon Sep 17 00:00:00 2001 From: Jiongqi Date: Fri, 25 Jun 2021 17:08:05 +0800 Subject: [PATCH 39/62] remove labelling type combination manually --- src/rectangle/utils/io.py | 11 +---------- 1 file changed, 1 insertion(+), 10 deletions(-) diff --git a/src/rectangle/utils/io.py b/src/rectangle/utils/io.py index 98028e4..f028cbd 100644 --- a/src/rectangle/utils/io.py +++ b/src/rectangle/utils/io.py @@ -84,8 +84,6 @@ def __init__(self, file, keys=None, label='random'): self.subjects = np.linspace(start_subj, last_subj, self.num_subjects+1, dtype=int) self.label = label - if label.split('_')[0] == 'combine': - self.label_loop = label def __len__(self): return self.num_subjects @@ -96,14 +94,7 @@ def __getitem__(self, index): image = torch.unsqueeze(torch.tensor( self.file['frame_%05d' % (subj_ix, )][()].astype('float32')), dim=0) - - if self.label_loop.split('_')[0] == 'combine': - label_percent = int(self.label_loop.split('_')[1]) - if index < int(self.num_subjects*label_percent/100): - self.label = 'vote' - else: - self.label = 'random' - + if self.label == 'random': label = torch.unsqueeze(torch.tensor( self.file['label_%05d_%02d' % (subj_ix, From 5a098947f165ce96e97d46b6fc125ba59c9e9098 Mon Sep 17 00:00:00 2001 From: Iani Date: Fri, 25 Jun 2021 10:32:10 +0100 Subject: [PATCH 40/62] Changed values from None to 0 and 1 for scaling --- train.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/train.py b/train.py index b076ce6..724645c 100644 --- a/train.py +++ b/train.py @@ -142,7 +142,7 @@ nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule) #Manually setting Affine Transforms -AffineTransform = rect.utils.transforms.Affine(prob = 0.3, scale = None, degrees = 5, shear = None, translate = None) +AffineTransform = rect.utils.transforms.Affine(prob = 0.3, scale = (1,1), degrees = 5, shear = 0, translate = 0) if args.val: trainer.train(train_data, val_data, train_pre=[rect.utils.transforms.z_score(), rect.utils.transforms.Flip(), AffineTransform], From ba80f0141e822d9effed24c0e88832cdbaf02cdf Mon Sep 17 00:00:00 2001 From: Jiongqi Date: Sat, 26 Jun 2021 20:52:43 +0800 Subject: [PATCH 41/62] fix the z_score bracket bug --- src/rectangle/utils/transforms.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/rectangle/utils/transforms.py b/src/rectangle/utils/transforms.py index 69ec712..05a84a1 100644 --- a/src/rectangle/utils/transforms.py +++ b/src/rectangle/utils/transforms.py @@ -32,8 +32,8 @@ def __call__(self, image): batch_ = image.shape[0] for batch_iter_ in range(batch_): image[batch_iter_,...] = (image[batch_iter_,...] - \ - torch.mean(image[batch_iter_,...]) / \ - torch.std(image[batch_iter_,...])) + torch.mean(image[batch_iter_,...]))/ \ + torch.std(image[batch_iter_,...]) return image From 11ad5cf0707c2054c2dbed79d43e1595e0b8bbc6 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Sat, 26 Jun 2021 14:35:01 +0100 Subject: [PATCH 42/62] Add CLI arg for early stopping --- train.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/train.py b/train.py index 724645c..c6077df 100644 --- a/train.py +++ b/train.py @@ -94,6 +94,14 @@ default=None, help='Random seed for training.') +parser.add_argument('--earlystop', + '--e', + metavar='earlystop', + type=str, + action='store', + default='10', + help='Number of val steps with no improvement before stopping training early.') + args = parser.parse_args() @@ -139,7 +147,8 @@ gate=args.gate) trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, device=device, - nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule) + nb_epochs=int(args.epochs), lr_schedule=args.lr_schedule, + early_stop=int(args.earlystop)) #Manually setting Affine Transforms AffineTransform = rect.utils.transforms.Affine(prob = 0.3, scale = (1,1), degrees = 5, shear = 0, translate = 0) From 4449249698882c023d2a363a6e24a8ffbf7578ca Mon Sep 17 00:00:00 2001 From: sophmrtn Date: Sat, 26 Jun 2021 20:27:02 +0100 Subject: [PATCH 43/62] fixes lr bug --- src/rectangle/utils/train.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index b94d922..5ab36c8 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -54,9 +54,9 @@ def __init__(self, model, nb_epochs=200, outdir='./logs', if opt == 'adam': if self.ensemble: - opt = [Adam(model.parameters()) for model in self.model_ensemble] + opt = [Adam(model.parameters(), lr=0.0001) for model in self.model_ensemble] else: - opt = Adam(model.parameters()) + opt = Adam(model.parameters(), lr=0.0001) # opt = Adam(model.parameters()) self.opt = opt @@ -252,9 +252,9 @@ def train(self, train_data, val_data=None, oname=None, loss_ = self.loss(pred, label) loss_.backward() self.opt.step() - if self.lr_schedule and self.lr_schedule != 'reduce_on_plateau': - self.lr_schedule.step() loss_epoch.append(loss_.item()) + if self.lr_schedule and self.lr_schedule != 'reduce_on_plateau': + self.lr_schedule.step() loss_log[epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: self.writer.add_scalar('train/dice_loss', loss_, epoch) @@ -280,11 +280,11 @@ def train(self, train_data, val_data=None, oname=None, for aug in val_post: pred = aug(pred) dice_metric = self.metric(pred, label) - if self.lr_schedule == 'reduce_on_plateau': - self.lr_schedule.step(dice_metric) # monitors validation loss + dice_epoch.append(1 - dice_metric.item()) dice_log[int(epoch//self.val_interval)] = np.nanmean(dice_epoch) - + if self.lr_schedule == 'reduce_on_plateau': + self.lr_schedule.step(1-np.nanmean(dice_epoch)) # monitors validation loss self.writer.add_scalar('val/dice_loss', dice_metric, epoch) self.writer.add_scalar('val/dice_coefficient', 1-dice_metric, epoch) From 2af797d5fb15d3d8035f3481cb8668fcdba62cc3 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Mon, 28 Jun 2021 10:51:16 +0100 Subject: [PATCH 44/62] Fix tensorboard metric outputs --- src/rectangle/utils/train.py | 32 ++++++++++++++++---------------- 1 file changed, 16 insertions(+), 16 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index b94d922..bf78c2b 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -151,8 +151,8 @@ def train(self, train_data, val_data=None, oname=None, lr_schedule_.step() loss_log_ensemble[i,epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: - writer_.add_scalar('train/dice_loss_ensemble', loss_, epoch) - writer_.add_scalar('train/dice_coefficient_ensemble', 1-loss_, epoch) + writer_.add_scalar('train/dice_loss_ensemble', np.nanmean(dice_epoch), epoch) + writer_.add_scalar('train/dice_coefficient_ensemble', 1-np.nanmean(dice_epoch), epoch) print('Epoch #{}: Mean Dice Loss: {}'.format(epoch, loss_log_ensemble[i,epoch])) if epoch % self.val_interval == 0: dice_epoch = [] @@ -179,8 +179,8 @@ def train(self, train_data, val_data=None, oname=None, if lr_schedule_ == 'reduce_on_plateau': lr_schedule_.step(1-np.nanmean(dice_epoch)) - writer_.add_scalar('val/dice_loss_ensemble', dice_metric, epoch) - writer_.add_scalar('val/dice_coefficient_ensemble', 1-dice_metric, epoch) + writer_.add_scalar('val/dice_loss_ensemble', np.nanmean(dice_epoch), epoch) + writer_.add_scalar('val/dice_coefficient_ensemble', 1-np.nanmean(dice_epoch), epoch) ## show some (e.g.,10) example images in tensorboard ex_num = 10 @@ -257,8 +257,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log[epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: - self.writer.add_scalar('train/dice_loss', loss_, epoch) - self.writer.add_scalar('train/dice_coefficient', 1-loss_, epoch) + self.writer.add_scalar('train/dice_loss', np.nanmean(loss_epoch), epoch) + self.writer.add_scalar('train/dice_coefficient', 1-np.nanmean(loss_epoch), epoch) print('Epoch #{}: Mean Dice Loss: {}'.format(epoch, loss_log[epoch])) if epoch % self.val_interval == 0: dice_epoch = [] @@ -285,8 +285,8 @@ def train(self, train_data, val_data=None, oname=None, dice_epoch.append(1 - dice_metric.item()) dice_log[int(epoch//self.val_interval)] = np.nanmean(dice_epoch) - self.writer.add_scalar('val/dice_loss', dice_metric, epoch) - self.writer.add_scalar('val/dice_coefficient', 1-dice_metric, epoch) + self.writer.add_scalar('val/dice_loss', np.nanmean(dice_epoch), epoch) + self.writer.add_scalar('val/dice_coefficient', 1-np.nanmean(dice_epoch), epoch) ## show some (e.g.,10) example images in tensorboard ex_num = 10 @@ -606,8 +606,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log_ensemble[i,epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: - self.writer.add_scalar('class_train/dice_loss_ensemble', loss_, epoch) - self.writer.add_scalar('class_train/dice_coefficient_ensemble', 1-loss_, epoch) + self.writer.add_scalar('class_train/dice_loss_ensemble', np.nanmean(loss_epoch), epoch) + self.writer.add_scalar('class_train/dice_coefficient_ensemble', 1-np.nanmean(loss_epoch), epoch) print('Epoch #{}: Mean acc Loss: {}'.format(epoch, loss_log_ensemble[i,epoch])) if epoch % self.val_interval == 0: acc_epoch = [] @@ -631,8 +631,8 @@ def train(self, train_data, val_data=None, oname=None, acc_metric = self.metric(pred, label) acc_epoch.append(acc_metric) acc_log_ensemble[i,int(epoch//self.val_interval)] = np.nanmean(acc_epoch) - self.writer.add_scalar('class_val/dice_loss_ensemble', acc_metric, epoch) - self.writer.add_scalar('class_val/dice_coefficient_ensemble', 1-acc_metric, epoch) + self.writer.add_scalar('class_val/dice_loss_ensemble', np.nanmean(acc_epoch), epoch) + self.writer.add_scalar('class_val/dice_coefficient_ensemble', 1-np.nanmean(acc_epoch), epoch) if epoch >= self.val_interval: if acc_log_ensemble[i,int(epoch//self.val_interval)] > acc_max: early_ = 0 @@ -683,8 +683,8 @@ def train(self, train_data, val_data=None, oname=None, loss_epoch.append(loss_.item()) loss_log[epoch] = np.nanmean(loss_epoch) if epoch % self.print_interval == 0: - self.writer.add_scalar('class_train/dice_loss', loss_, epoch) - self.writer.add_scalar('class_train/dice_coefficient', 1-loss_, epoch) + self.writer.add_scalar('class_train/dice_loss', np.nanmean(loss_epoch), epoch) + self.writer.add_scalar('class_train/dice_coefficient', 1-np.nanmean(loss_epoch), epoch) print('Epoch #{}: Mean acc Loss: {}'.format(epoch, loss_log[epoch])) if epoch % self.val_interval == 0: acc_epoch = [] @@ -708,8 +708,8 @@ def train(self, train_data, val_data=None, oname=None, acc_metric = self.metric(pred, label) acc_epoch.append(acc_metric) acc_log[int(epoch//self.val_interval)] = np.nanmean(acc_epoch) - self.writer.add_scalar('class_val/dice_loss', acc_metric, epoch) - self.writer.add_scalar('class_val/dice_coefficient', 1-acc_metric, epoch) + self.writer.add_scalar('class_val/dice_loss', np.nanmean(acc_epoch), epoch) + self.writer.add_scalar('class_val/dice_coefficient', 1-np.nanmean(acc_epoch), epoch) if epoch % self.print_interval == 0: print('Mean Validation acc: {}'.format(acc_log[int(epoch//self.val_interval)])) if epoch >= self.val_interval: From 947b48a8a3a1509e7844cabe783c01ad7f749865 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 14:22:13 +0100 Subject: [PATCH 45/62] Make classifier optional in test.py --- test.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/test.py b/test.py index ae8b7a9..82d46ef 100644 --- a/test.py +++ b/test.py @@ -113,9 +113,11 @@ for n, m in enumerate(model): m.load_state_dict(torch.load(args.weights[n])) -class_model = rect.model.networks.MakeDenseNet(freeze_weights=False).to(device) -class_model.load_state_dict(torch.load(args.classweights)) +if args.classifier: + class_model = rect.model.networks.MakeDenseNet(freeze_weights=False).to(device) +if args.classweights: + class_model.load_state_dict(torch.load(args.classweights)) if torch.cuda.is_available(): device = torch.device('cuda') @@ -131,5 +133,5 @@ trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, device=device) -trainer.test(test_data, test_pre=[rect.utils.transforms.z_score()], - test_post=[rect.utils.transforms.Binary(), rect.utils.transforms.KeepLargestComponent()]) +trainer.test(test_data, test_pre=[rect.utils.transforms.z_score()], oname='run', + test_post=[rect.utils.transforms.Binary()], overlap='mask') From 91f3e473b85f0950454b86f0b92485eb6c4f9ede Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 16:24:07 +0100 Subject: [PATCH 46/62] Fix bug in test.py --- test.py | 26 ++++++++++++++++---------- 1 file changed, 16 insertions(+), 10 deletions(-) diff --git a/test.py b/test.py index 82d46ef..308a254 100644 --- a/test.py +++ b/test.py @@ -91,7 +91,10 @@ if args.ensemble: ensemble = int(args.ensemble) else: - ensemble = None + if args.weights: + ensemble=int(len(args.weights)) + else: + ensemble = None ## run training import rectangle as rect @@ -107,24 +110,27 @@ random.seed(seed) np.random.seed(seed) -model = [rect.model.networks.UNet(n_layers=int(args.depth), device=device, - gate=args.gate) for e in int(args.ensemble)] +if torch.cuda.is_available(): + device = torch.device('cuda') + torch.backends.cudnn.benchmark = True +else: + device = torch.device('cpu') + +if ensemble: + model = [rect.model.networks.UNet(n_layers=int(args.depth), device=device, + gate=args.gate) for e in range(ensemble)] +else: + model = rect.model.networks.UNet(n_layers=int(args.depth), device=device, + gate=args.gate) for n, m in enumerate(model): m.load_state_dict(torch.load(args.weights[n])) - if args.classifier: class_model = rect.model.networks.MakeDenseNet(freeze_weights=False).to(device) if args.classweights: class_model.load_state_dict(torch.load(args.classweights)) -if torch.cuda.is_available(): - device = torch.device('cuda') - torch.backends.cudnn.benchmark = True -else: - device = torch.device('cpu') - f_test = h5py.File(args.test, 'r') if args.classifier: train_data = rect.utils.io.PreScreenLoader(class_model.eval(), f_test, label=args.label, threshold=float(args.thresh)) From d72bbdcc4b2c513a17f2b1b289be454aa1d2e85b Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 17:02:55 +0100 Subject: [PATCH 47/62] Fix compatibility with weights on CPU --- test.py | 12 ++++-------- 1 file changed, 4 insertions(+), 8 deletions(-) diff --git a/test.py b/test.py index 308a254..d22c9fa 100644 --- a/test.py +++ b/test.py @@ -94,7 +94,7 @@ if args.weights: ensemble=int(len(args.weights)) else: - ensemble = None + ensemble = 1 ## run training import rectangle as rect @@ -116,15 +116,11 @@ else: device = torch.device('cpu') -if ensemble: - model = [rect.model.networks.UNet(n_layers=int(args.depth), device=device, - gate=args.gate) for e in range(ensemble)] -else: - model = rect.model.networks.UNet(n_layers=int(args.depth), device=device, - gate=args.gate) +model = [rect.model.networks.UNet(n_layers=int(args.depth), device=device, + gate=args.gate) for e in range(ensemble)] for n, m in enumerate(model): - m.load_state_dict(torch.load(args.weights[n])) + m.load_state_dict(torch.load(args.weights[n], map_location=device)) if args.classifier: class_model = rect.model.networks.MakeDenseNet(freeze_weights=False).to(device) From ea3376d8bfb4123362234d3c0b2a6e046a94cb17 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 17:04:56 +0100 Subject: [PATCH 48/62] Set classifier default=False --- test.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test.py b/test.py index d22c9fa..7758149 100644 --- a/test.py +++ b/test.py @@ -58,7 +58,7 @@ metavar='classifier', type=bool, action='store', - default=True, + default=False, help='Use of classifier for pre-screening. If selected will train without and then perform test without + with.') parser.add_argument('--classweights', From c287245b7d34bd6b28463070e00557fd4a843f4c Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 17:06:48 +0100 Subject: [PATCH 49/62] FIx defaults for class stuff in test.py --- test.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/test.py b/test.py index 7758149..700698d 100644 --- a/test.py +++ b/test.py @@ -66,7 +66,7 @@ metavar='classweights', type=str, action='store', - default=True, + default=None, help='Path to trained weights for classifier.') parser.add_argument('--threshold', @@ -122,7 +122,7 @@ for n, m in enumerate(model): m.load_state_dict(torch.load(args.weights[n], map_location=device)) -if args.classifier: +if args.classifier==True: class_model = rect.model.networks.MakeDenseNet(freeze_weights=False).to(device) if args.classweights: class_model.load_state_dict(torch.load(args.classweights)) From 9636da17c4557b6e546a8061a809256d60cd175c Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 17:17:12 +0100 Subject: [PATCH 50/62] Fix eval() in trainer test --- src/rectangle/utils/train.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index bf78c2b..5679677 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -381,7 +381,10 @@ def test(self, test_data, oname=None, rec_log = [] precision = Precision() recall = Recall() - self.model.eval() + if self.ensemble: + self.model = [model.eval() for model in self.model] + else: + self.model.eval() with torch.no_grad(): for i, (input, label) in enumerate(test): input, label = input.to(self.device), label.to(self.device) From 15c83d62a21f518f403c4a66b40c32199a064dfe Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 17:46:13 +0100 Subject: [PATCH 51/62] Improve plot settings for testing --- src/rectangle/utils/train.py | 101 ++++++++++++++++++----------------- 1 file changed, 52 insertions(+), 49 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index 5679677..3a2a69b 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -375,6 +375,11 @@ def test(self, test_data, oname=None, oname = date.today() oname = oname.strftime("%b-%d-%Y") + path_ = path.join(self.outdir,\ + 'testing/plots') + if not path.exists(path_): + makedirs(path_) + test = DataLoader(test_data, 1, shuffle=False) dice_log = [] prec_log = [] @@ -407,7 +412,7 @@ def test(self, test_data, oname=None, for aug in test_post: pred = aug(pred) dice_metric = self.metric(pred, label) - dice_log.append(1-dice_metric.item().detach().cpu().numpy()) + dice_log.append(1-dice_metric.detach().cpu().numpy()) prec_log.append(precision(pred, label).detach().cpu().numpy()) rec_log.append(recall(pred, label).detach().cpu().numpy()) @@ -419,54 +424,52 @@ def test(self, test_data, oname=None, pred_img = np.squeeze(pred_img) label_img = np.squeeze(label_img) - if overlap=='contour': - input_img -= input_img.min() - input_img *= 1.0/input_img.max() - label_img = laplace(label_img) - pred_img = laplace(pred_img) - label_img = (label_img != 0) - pred_img = (pred_img != 0) - label_img = np.ma.masked_where(label_img == 0, label_img) - pred_img = np.ma.masked_where(pred_img == 0, pred_img) - plt.figure() - plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1) - plt.axis('off') - elif overlap=='mask': - input_img -= input_img.min() - input_img *= 1.0/input_img.max() - label_img = np.ma.masked_where(label_img == 0, label_img) - pred_img = np.ma.masked_where(pred_img == 0, pred_img) - plt.figure() - plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1) - plt.axis('off') - else: - plt.figure() - plt.subplot(131) - plt.imshow(input_img, cmap='gray') - plt.axis('off') - plt.title('Image') - plt.subplot(132) - plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.subplot(133) - plt.imshow(label_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.title('Ground Truth') - - path_ = path.join(self.outdir,\ - 'testing/plots') - if not path.exists(path_): - makedirs(path_) - plt.savefig(path.join(path_, 'pred{}_{}.png'.format(i, oname))) + if overlap: + if overlap=='contour': + input_img -= input_img.min() + input_img *= 1.0/input_img.max() + label_img = laplace(label_img) + pred_img = laplace(pred_img) + label_img = (label_img != 0) + pred_img = (pred_img != 0) + label_img = np.ma.masked_where(label_img == 0, label_img) + pred_img = np.ma.masked_where(pred_img == 0, pred_img) + plt.figure() + plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1) + plt.axis('off') + plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1) + plt.axis('off') + elif overlap=='mask': + input_img -= input_img.min() + input_img *= 1.0/input_img.max() + label_img = np.ma.masked_where(label_img == 0, label_img) + pred_img = np.ma.masked_where(pred_img == 0, pred_img) + plt.figure() + plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1, alpha=0.5) + plt.axis('off') + plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1, alpha=0.5) + plt.axis('off') + + plt.savefig(path.join(path_, 'pred_overlap_{}_{}.png'.format(i, oname))) + + plt.figure() + plt.subplot(131) + plt.imshow(input_img, cmap='gray') + plt.axis('off') + plt.title('Image') + plt.subplot(132) + plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) + plt.subplot(133) + plt.imshow(label_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.title('Ground Truth') + plt.savefig(path.join(path_, 'pred_{}_{}.png'.format(i, oname))) dice_log = np.array(dice_log, dtype=float) prec_log = np.array(prec_log, dtype=float) From 25126ac4f83212187c1f63061e070c59c2c6fddb Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Tue, 29 Jun 2021 18:23:22 +0100 Subject: [PATCH 52/62] Adjust test.py to new settings --- src/rectangle/utils/train.py | 27 ++++++++++----------------- 1 file changed, 10 insertions(+), 17 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index 3a2a69b..b267b0c 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -424,6 +424,12 @@ def test(self, test_data, oname=None, pred_img = np.squeeze(pred_img) label_img = np.squeeze(label_img) + plt.figure() + plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) + plt.savefig(path.join(path_, 'pred_{}_{}.png'.format(i, oname))) + if overlap: if overlap=='contour': input_img -= input_img.min() @@ -440,6 +446,7 @@ def test(self, test_data, oname=None, plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1) plt.axis('off') plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1) + plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) plt.axis('off') elif overlap=='mask': input_img -= input_img.min() @@ -449,28 +456,14 @@ def test(self, test_data, oname=None, plt.figure() plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) plt.axis('off') - plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1, alpha=0.5) + plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1, alpha=0.3) plt.axis('off') - plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1, alpha=0.5) + plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1, alpha=0.3) + plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) plt.axis('off') plt.savefig(path.join(path_, 'pred_overlap_{}_{}.png'.format(i, oname))) - plt.figure() - plt.subplot(131) - plt.imshow(input_img, cmap='gray') - plt.axis('off') - plt.title('Image') - plt.subplot(132) - plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.subplot(133) - plt.imshow(label_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.title('Ground Truth') - plt.savefig(path.join(path_, 'pred_{}_{}.png'.format(i, oname))) - dice_log = np.array(dice_log, dtype=float) prec_log = np.array(prec_log, dtype=float) rec_log = np.array(rec_log, dtype=float) From 007909d915857fab44aaf1ebd6e5295cf6c6a18b Mon Sep 17 00:00:00 2001 From: sophmrtn Date: Tue, 29 Jun 2021 18:29:42 +0100 Subject: [PATCH 53/62] adding plot example function --- src/rectangle/utils/io.py | 38 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 38 insertions(+) diff --git a/src/rectangle/utils/io.py b/src/rectangle/utils/io.py index f028cbd..24b4fab 100644 --- a/src/rectangle/utils/io.py +++ b/src/rectangle/utils/io.py @@ -3,7 +3,45 @@ import numpy as np import random +def plot_example(frame, labels, gt_method=None, savefig=None): + ''' + frame = image array + labels = single label or list of labels + savefig= filepath to save, if None (default) use plt.show() + gt_method = String to describe ground truth method used + ''' + + colors=['lime', 'red', 'blue', 'orange'] # cycles through colors in this order + if gt_method is None: + gt_method = 'Vote' + + legend_names = ['Label 1', 'Label 2', 'Label 3', gt_method] + # 0 = label 1 = lime + # 1 = label 2 = red + # 3 = label 3 = blue + # 4 = ground truth = orange (if included in list) + + plt.figure(figsize=(12, 12)) + plt.imshow(frame, cmap='gray') + if type(labels) == list: + for i, label in enumerate(labels): + plt.contour(label, colors=colors[i], linewidths=1) + else: + plt.contour(label, colors=colors[0], linewidths=1) + plt.tight_layout() + plt.axis('off') + + # make legend + patches = [ mpatches.Patch(color=colors[i], label=legend_names[i]) for i in range(len(labels) ) ] + # put those patched as legend-handles into the legend + plt.legend(handles=patches, bbox_to_anchor=(0.97, 0.97), loc=1, borderaxespad=0., fontsize=20) + if savefig is None: + plt.show() + else: + plt.savefig(savefig) + + def train_val_test(file, ratio=(0.6, 0.2, 0.2)): """ Generate list of keys for file based on index values Input arguments: From 02e3472c2362d0d4277f165dfa8ac8937258f3c5 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Wed, 30 Jun 2021 09:47:21 +0100 Subject: [PATCH 54/62] Reduce frequency of plotting stuff in test --- src/rectangle/utils/train.py | 79 ++++++++++++++++++------------------ 1 file changed, 40 insertions(+), 39 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index b267b0c..e6d57f5 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -424,45 +424,46 @@ def test(self, test_data, oname=None, pred_img = np.squeeze(pred_img) label_img = np.squeeze(label_img) - plt.figure() - plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.savefig(path.join(path_, 'pred_{}_{}.png'.format(i, oname))) - - if overlap: - if overlap=='contour': - input_img -= input_img.min() - input_img *= 1.0/input_img.max() - label_img = laplace(label_img) - pred_img = laplace(pred_img) - label_img = (label_img != 0) - pred_img = (pred_img != 0) - label_img = np.ma.masked_where(label_img == 0, label_img) - pred_img = np.ma.masked_where(pred_img == 0, pred_img) - plt.figure() - plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1) - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.axis('off') - elif overlap=='mask': - input_img -= input_img.min() - input_img *= 1.0/input_img.max() - label_img = np.ma.masked_where(label_img == 0, label_img) - pred_img = np.ma.masked_where(pred_img == 0, pred_img) - plt.figure() - plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1, alpha=0.3) - plt.axis('off') - plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1, alpha=0.3) - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.axis('off') - - plt.savefig(path.join(path_, 'pred_overlap_{}_{}.png'.format(i, oname))) + if i % 100==0: + plt.figure() + plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) + plt.savefig(path.join(path_, 'pred_{}_{}.png'.format(i, oname))) + + if overlap: + if overlap=='contour': + input_img -= input_img.min() + input_img *= 1.0/input_img.max() + label_img = laplace(label_img) + pred_img = laplace(pred_img) + label_img = (label_img != 0) + pred_img = (pred_img != 0) + label_img = np.ma.masked_where(label_img == 0, label_img) + pred_img = np.ma.masked_where(pred_img == 0, pred_img) + plt.figure() + plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1) + plt.axis('off') + plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1) + plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) + plt.axis('off') + elif overlap=='mask': + input_img -= input_img.min() + input_img *= 1.0/input_img.max() + label_img = np.ma.masked_where(label_img == 0, label_img) + pred_img = np.ma.masked_where(pred_img == 0, pred_img) + plt.figure() + plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1, alpha=0.3) + plt.axis('off') + plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1, alpha=0.3) + plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) + plt.axis('off') + + plt.savefig(path.join(path_, 'pred_overlap_{}_{}.png'.format(i, oname))) dice_log = np.array(dice_log, dtype=float) prec_log = np.array(prec_log, dtype=float) From 9998007232f620ede129341634d4ae37faaa1d70 Mon Sep 17 00:00:00 2001 From: sophmrtn Date: Wed, 30 Jun 2021 15:22:54 +0100 Subject: [PATCH 55/62] add weighted bce custom loss --- src/rectangle/utils/metrics.py | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/src/rectangle/utils/metrics.py b/src/rectangle/utils/metrics.py index 4163d52..61f5949 100644 --- a/src/rectangle/utils/metrics.py +++ b/src/rectangle/utils/metrics.py @@ -3,6 +3,26 @@ # Loss function +class WeightedBCE(nn.Module): + def __init__(self, weights=None): + super().__init__() + self.weights = weights + + def forward(self, inputs, targets): + inputs = inputs.view(-1).float() + targets = targets.view(-1).float() + + if self.weights is not None: + assert len(self.weights) == 2 + + loss = weights[1] * (targets * torch.log(inputs)) + \ + weights[0] * ((1 - targets) * torch.log(1 - inputs)) + else: + loss = targets * torch.log(inputs) + (1 - targets) * torch.log(1 - inputs) + + return torch.neg(torch.mean(loss)) + + class DiceLoss(nn.Module): """ Loss function based on Dice-Sorensen Coefficient (L = 1 - Dice) Input arguments: From 3a2efbcbde55d5cefd6eb24bd0df789833b10a40 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Wed, 30 Jun 2021 18:09:25 +0100 Subject: [PATCH 56/62] Add logging of positive and negatives to test Also fixed contour plot --- src/rectangle/utils/train.py | 127 ++++++++++++++++++++--------------- test.py | 2 +- 2 files changed, 72 insertions(+), 57 deletions(-) diff --git a/src/rectangle/utils/train.py b/src/rectangle/utils/train.py index efa694a..8293f06 100644 --- a/src/rectangle/utils/train.py +++ b/src/rectangle/utils/train.py @@ -386,6 +386,8 @@ def test(self, test_data, oname=None, dice_log = [] prec_log = [] rec_log = [] + neg_log = [] + pos_log = [] precision = Precision() recall = Recall() if self.ensemble: @@ -413,74 +415,83 @@ def test(self, test_data, oname=None, if test_post: for aug in test_post: pred = aug(pred) - dice_metric = self.metric(pred, label) - dice_log.append(1-dice_metric.detach().cpu().numpy()) - prec_log.append(precision(pred, label).detach().cpu().numpy()) - rec_log.append(recall(pred, label).detach().cpu().numpy()) - - input_img = input.detach().cpu().numpy() - pred_img = pred.detach().cpu().numpy() - label_img = label.detach().cpu().numpy() - - input_img = np.squeeze(input_img) - pred_img = np.squeeze(pred_img) - label_img = np.squeeze(label_img) - - if i % 100==0: - plt.figure() - plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.savefig(path.join(path_, 'pred_{}_{}.png'.format(i, oname))) - - if overlap: - if overlap=='contour': - input_img -= input_img.min() - input_img *= 1.0/input_img.max() - label_img = laplace(label_img) - pred_img = laplace(pred_img) - label_img = (label_img != 0) - pred_img = (pred_img != 0) - label_img = np.ma.masked_where(label_img == 0, label_img) - pred_img = np.ma.masked_where(pred_img == 0, pred_img) - plt.figure() - plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1) - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.axis('off') - elif overlap=='mask': - input_img -= input_img.min() - input_img *= 1.0/input_img.max() - label_img = np.ma.masked_where(label_img == 0, label_img) - pred_img = np.ma.masked_where(pred_img == 0, pred_img) - plt.figure() - plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) - plt.axis('off') - plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1, alpha=0.3) - plt.axis('off') - plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1, alpha=0.3) - plt.title('Prediction (DSC={:.2f})'.format(dice_log[i])) - plt.axis('off') - - plt.savefig(path.join(path_, 'pred_overlap_{}_{}.png'.format(i, oname))) + + if pred.sum() == 0: + if label.sum() == 0: + neg_log.append(1.) + else: + neg_log.append(0.) + else: + if label.sum() == 0: + pos_log.append(0.) + else: + pos_log.append(1.) + dice_metric = self.metric(pred, label) + dice_log.append(1-dice_metric.detach().cpu().numpy()) + prec_log.append(precision(pred, label).detach().cpu().numpy()) + rec_log.append(recall(pred, label).detach().cpu().numpy()) + + input_img = input.detach().cpu().numpy() + pred_img = pred.detach().cpu().numpy() + label_img = label.detach().cpu().numpy() + + input_img = np.squeeze(input_img) + pred_img = np.squeeze(pred_img) + label_img = np.squeeze(label_img) + + if i % 50==0: + plt.figure() + plt.imshow(pred_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.title('Prediction (DSC={:.2f})'.format(dice_log[-1])) + plt.savefig(path.join(path_, 'pred_{}_{}.png'.format(i, oname))) + + if overlap: + if overlap=='contour': + input_img -= input_img.min() + input_img *= 1.0/input_img.max() + # label_img = np.ma.masked_where(label_img == 0, label_img) + # pred_img = np.ma.masked_where(pred_img == 0, pred_img) + plt.figure() + plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.contour(label_img, cmap='Greens', linewidths=1) + plt.axis('off') + plt.contour(pred_img, cmap='Reds', linewidths=1) + plt.title('Prediction (DSC={:.2f})'.format(dice_log[-1])) + plt.axis('off') + elif overlap=='mask': + input_img -= input_img.min() + input_img *= 1.0/input_img.max() + label_img = np.ma.masked_where(label_img == 0, label_img) + pred_img = np.ma.masked_where(pred_img == 0, pred_img) + plt.figure() + plt.imshow(input_img, cmap='gray', vmin=0, vmax=1) + plt.axis('off') + plt.imshow(label_img, cmap='Greens', vmin=0, vmax=1, alpha=0.3) + plt.axis('off') + plt.imshow(pred_img, cmap='Reds', vmin=0, vmax=1, alpha=0.3) + plt.title('Prediction (DSC={:.2f})'.format(dice_log[-1])) + plt.axis('off') + + plt.savefig(path.join(path_, 'pred_overlap_{}_{}.png'.format(i, oname))) dice_log = np.array(dice_log, dtype=float) prec_log = np.array(prec_log, dtype=float) rec_log = np.array(rec_log, dtype=float) + neg_log = np.array(neg_log, dtype=float) + pos_log = np.array(pos_log, dtype=float) plt.figure() - plt.scatter(rec_log, prec_log) - plt.plot([0,0.5,1], [0.5,0.5,0.5], '--') + plt.scatter(rec_log, prec_log, alpha=0.4) + # plt.plot([0,0.5,1], [0.5,0.5,0.5], '--') plt.xlabel('Recall') plt.ylabel('Precision') plt.title('AUC = {:.2f}'.format(np.sum(prec_log * rec_log)/np.size(prec_log))) plt.savefig(path.join(path_, 'prec_rec_{}'.format(oname))) - print('Mean Dice score: {:.2f}±{:.3f}, Mean Precision: {:.2f}±{:.3f}, Mean Recall: {:.2f}±{:.3f}'.format(np.mean(dice_log), np.std(dice_log), np.mean(prec_log), np.std(prec_log), np.mean(rec_log), np.std(rec_log))) + print('Mean Dice score: {:.2f}±{:.3f}, Mean Precision: {:.2f}±{:.3f}, Mean Recall: {:.2f}±{:.3f} \n TP Rate: {:.2f}±{:.3f}, TN Rate: {:.2f}±{:.3f}'.format(np.mean(dice_log), np.std(dice_log), np.mean(prec_log), np.std(prec_log), np.mean(rec_log), np.std(rec_log), np.mean(pos_log), np.std(pos_log), np.mean(neg_log), np.std(neg_log))) path_ = path.join(self.outdir,\ 'testing/table') if not path.exists(path_): @@ -491,6 +502,10 @@ def test(self, test_data, oname=None, prec_log, delimiter=',') np.savetxt(path.join(path_, 'recall_{}.csv'.format(oname)),\ rec_log, delimiter=',') + np.savetxt(path.join(path_, 'negative_{}.csv'.format(oname)),\ + neg_log, delimiter=',') + np.savetxt(path.join(path_, 'positive_{}.csv'.format(oname)),\ + pos_log, delimiter=',') print('Testing complete') diff --git a/test.py b/test.py index 700698d..9845da6 100644 --- a/test.py +++ b/test.py @@ -136,4 +136,4 @@ trainer = rect.utils.train.Trainer(model, ensemble=ensemble, outdir=args.odir, device=device) trainer.test(test_data, test_pre=[rect.utils.transforms.z_score()], oname='run', - test_post=[rect.utils.transforms.Binary()], overlap='mask') + test_post=[rect.utils.transforms.Binary()], overlap='contour') From da7014cfc029509230c0862499178a90447d9339 Mon Sep 17 00:00:00 2001 From: Liam Chalcroft Date: Thu, 1 Jul 2021 17:17:41 +0100 Subject: [PATCH 57/62] Fix typo in test.py --- test.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/test.py b/test.py index 9845da6..ef4175d 100644 --- a/test.py +++ b/test.py @@ -129,7 +129,7 @@ f_test = h5py.File(args.test, 'r') if args.classifier: - train_data = rect.utils.io.PreScreenLoader(class_model.eval(), f_test, label=args.label, threshold=float(args.thresh)) + test_data = rect.utils.io.PreScreenLoader(class_model.eval(), f_test, label=args.label, threshold=float(args.thresh)) else: test_data = rect.utils.io.H5DataLoader(f_test, label='vote') From 306eec6ac5c62b9c7d06d1e07333a554eee943ae Mon Sep 17 00:00:00 2001 From: sophmrtn Date: Thu, 1 Jul 2021 21:04:10 +0100 Subject: [PATCH 58/62] correcting loss --- src/rectangle/utils/metrics.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/rectangle/utils/metrics.py b/src/rectangle/utils/metrics.py index 61f5949..74e75ee 100644 --- a/src/rectangle/utils/metrics.py +++ b/src/rectangle/utils/metrics.py @@ -20,7 +20,7 @@ def forward(self, inputs, targets): else: loss = targets * torch.log(inputs) + (1 - targets) * torch.log(1 - inputs) - return torch.neg(torch.mean(loss)) + return loss class DiceLoss(nn.Module): From d44a1123c74e59f004783d845d2feae6d0a583c1 Mon Sep 17 00:00:00 2001 From: Jiongqi <72549351+Jiongqi@users.noreply.github.com> Date: Sat, 3 Jul 2021 00:00:54 +0800 Subject: [PATCH 59/62] Update README.md --- README.md | 2 ++ 1 file changed, 2 insertions(+) diff --git a/README.md b/README.md index b749bf7..d5e7af2 100644 --- a/README.md +++ b/README.md @@ -28,3 +28,5 @@ Following this, training/inference may be performed using objects in the *train* To familiarise with the code used, an interactive notebook used for experiments in the associated report is available below. Please note that data used is proprietary and so has been withheld from the published repository. Open In Colab + +(The code relevant to different label sampling methods is in sub-branch: label_method.) From c29b12c40a2114ae3731f4282bb7b065bcc0d197 Mon Sep 17 00:00:00 2001 From: sophmrtn Date: Fri, 2 Jul 2021 19:52:08 +0100 Subject: [PATCH 60/62] updating legend sizes and linewidth for plot_example --- src/rectangle/utils/io.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/rectangle/utils/io.py b/src/rectangle/utils/io.py index 24b4fab..962dcb6 100644 --- a/src/rectangle/utils/io.py +++ b/src/rectangle/utils/io.py @@ -25,16 +25,16 @@ def plot_example(frame, labels, gt_method=None, savefig=None): plt.imshow(frame, cmap='gray') if type(labels) == list: for i, label in enumerate(labels): - plt.contour(label, colors=colors[i], linewidths=1) + plt.contour(label, colors=colors[i], linewidths=2) else: - plt.contour(label, colors=colors[0], linewidths=1) + plt.contour(label, colors=colors[0], linewidths=2) plt.tight_layout() plt.axis('off') # make legend patches = [ mpatches.Patch(color=colors[i], label=legend_names[i]) for i in range(len(labels) ) ] # put those patched as legend-handles into the legend - plt.legend(handles=patches, bbox_to_anchor=(0.97, 0.97), loc=1, borderaxespad=0., fontsize=20) + plt.legend(handles=patches, bbox_to_anchor=(0.97, 0.97), loc=1, borderaxespad=0., fontsize=30) if savefig is None: plt.show() From d72fa6c82d6c76b02096de04b7cb9db93b76c90d Mon Sep 17 00:00:00 2001 From: Sophie Martin <44570734+sophmrtn@users.noreply.github.com> Date: Tue, 27 Jul 2021 22:28:35 +0100 Subject: [PATCH 61/62] Update README.md --- README.md | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index d5e7af2..e0b46b0 100644 --- a/README.md +++ b/README.md @@ -1,9 +1,9 @@ # RectAngle Segmentation and classification tool for trans-rectal B-mode ultrasound images. -*Submitted as coursework for MPHY0041: Machine Learning in Medical Imaging.* +*Submitted as part of the ASMUS2021 Conference for the paper 'Development and evaluation of intraoperative ultrasound segmentation with negative image frames and multiple observer labels'.* -This package contains PyTorch-based implementations of a U-Net based segmentation model, and a DenseNet-based classification model, for the simultaneous detection and segmentation of prostate in rectal b-mode ultrasound images. +This package contains a PyTorch-based implementation of a U-Net based segmentation model, and a DenseNet-based classification model, for the detection and segmentation of prostate in rectal b-mode ultrasound images. ## Installation From 941138fb63bdc3f3cb297a94fa057a16b88b00be Mon Sep 17 00:00:00 2001 From: Sophie Martin <44570734+sophmrtn@users.noreply.github.com> Date: Tue, 27 Jul 2021 22:29:14 +0100 Subject: [PATCH 62/62] Update README.md --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index e0b46b0..5ab8fc2 100644 --- a/README.md +++ b/README.md @@ -25,7 +25,7 @@ Once this is activated, the package may be installed using the *setup.py* file: Following this, training/inference may be performed using objects in the *train* module. -To familiarise with the code used, an interactive notebook used for experiments in the associated report is available below. Please note that data used is proprietary and so has been withheld from the published repository. +To familiarise yourself with the code used, an interactive notebook used for experiments in the associated report is available below. Please note that data used is proprietary and so has been withheld from the published repository. Open In Colab