Jupitern52/TextBraTS
2
1# Copyright 2020 - 2022 MONAI Consortium2# Licensed under the Apache License, Version 2.0 (the "License");3# you may not use this file except in compliance with the License.4# You may obtain a copy of the License at5# http://www.apache.org/licenses/LICENSE-2.06# Unless required by applicable law or agreed to in writing, software7# distributed under the License is distributed on an "AS IS" BASIS,8# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.9# See the License for the specific language governing permissions and10# limitations under the License.11 12import os13import shutil14import time15 16import numpy as np17import torch18import torch.nn.parallel19import torch.utils.data.distributed20from tensorboardX import SummaryWriter21from torch.amp import GradScaler, autocast22from utils.utils import AverageMeter, distributed_all_gather23 24from monai.data import decollate_batch25 26 27def train_epoch(model, loader, optimizer, scaler, epoch, loss_func, args):28 model.train()29 start_time = time.time()30 run_loss = AverageMeter()31 for idx, batch_data in enumerate(loader):32 if isinstance(batch_data, list):33 data, target, text = batch_data34 else:35 data, target, text = batch_data["image"], batch_data["label"], batch_data["text_feature"]36 data, target, text = data.cuda(args.rank), target.cuda(args.rank), text.cuda(args.rank)37 optimizer.zero_grad(set_to_none=True)38 with autocast('cuda',enabled=args.amp):39 logits = model(data,text)40 loss = loss_func(logits, target)41 if args.amp:42 scaler.scale(loss).backward()43 scaler.step(optimizer)44 scaler.update()45 else:46 loss.backward()47 optimizer.step()48 if args.distributed:49 loss_list = distributed_all_gather([loss], out_numpy=True, is_valid=idx < loader.sampler.valid_length)50 run_loss.update(51 np.mean(np.mean(np.stack(loss_list, axis=0), axis=0), axis=0), n=args.batch_size * args.world_size52 )53 else:54 run_loss.update(loss.item(), n=args.batch_size)55 if args.rank == 0:56 print(57 "Epoch {}/{} {}/{}".format(epoch, args.max_epochs, idx, len(loader)),58 "loss: {:.4f}".format(run_loss.avg),59 "time {:.2f}s".format(time.time() - start_time),60 )61 start_time = time.time()62 '''for param in model.parameters():63 param.grad = None'''64 optimizer.zero_grad(set_to_none=True)65 return run_loss.avg66 67 68def val_epoch(model, loader, epoch, acc_func, args, post_sigmoid=None, post_pred=None):69 model.eval()70 start_time = time.time()71 run_acc = AverageMeter()72 73 with torch.no_grad():74 for idx, batch_data in enumerate(loader):75 data, target, text = batch_data["image"], batch_data["label"], batch_data["text_feature"]76 data, target, text = data.cuda(args.rank), target.cuda(args.rank), text.cuda(args.rank)77 with autocast('cuda',enabled=args.amp):78 logits = model(data,text)79 val_labels_list = decollate_batch(target)80 val_outputs_list = decollate_batch(logits)81 val_output_convert = [post_pred(post_sigmoid(val_pred_tensor)) for val_pred_tensor in val_outputs_list]82 acc_func.reset()83 acc_func(y_pred=val_output_convert, y=val_labels_list)84 acc, not_nans = acc_func.aggregate()85 acc = acc.cuda(args.rank)86 if args.distributed:87 acc_list, not_nans_list = distributed_all_gather(88 [acc, not_nans], out_numpy=True, is_valid=idx < loader.sampler.valid_length89 )90 for al, nl in zip(acc_list, not_nans_list):91 run_acc.update(al, n=nl)92 else:93 run_acc.update(acc.cpu().numpy(), n=not_nans.cpu().numpy())94 95 if args.rank == 0:96 Dice_TC = run_acc.avg[0]97 Dice_WT = run_acc.avg[1]98 Dice_ET = run_acc.avg[2]99 print(100 "Val {}/{} {}/{}".format(epoch, args.max_epochs, idx, len(loader)),101 ", Dice_TC:",102 Dice_TC,103 ", Dice_WT:",104 Dice_WT,105 ", Dice_ET:",106 Dice_ET,107 ", time {:.2f}s".format(time.time() - start_time),108 )109 start_time = time.time()110 111 return run_acc.avg112 113 114def save_checkpoint(model, epoch, args, filename="model.pt", best_acc=0, optimizer=None, scheduler=None):115 state_dict = model.state_dict() if not args.distributed else model.module.state_dict()116 save_dict = {"epoch": epoch, "best_acc": best_acc, "state_dict": state_dict}117 if optimizer is not None:118 save_dict["optimizer"] = optimizer.state_dict()119 if scheduler is not None:120 save_dict["scheduler"] = scheduler.state_dict()121 filename = os.path.join(args.logdir, filename)122 torch.save(save_dict, filename)123 print("Saving checkpoint", filename)124 125 126def run_training(127 model,128 train_loader,129 val_loader,130 optimizer,131 loss_func,132 acc_func,133 args,134 scheduler=None,135 start_epoch=0,136 post_sigmoid=None,137 post_pred=None,138 semantic_classes=None,139):140 writer = None141 if args.logdir is not None and args.rank == 0:142 writer = SummaryWriter(log_dir=args.logdir)143 if args.rank == 0:144 print("Writing Tensorboard logs to ", args.logdir)145 scaler = None146 if args.amp:147 scaler = GradScaler()148 val_acc_max = 0.0149 for epoch in range(start_epoch, args.max_epochs):150 if args.distributed:151 train_loader.sampler.set_epoch(epoch)152 torch.distributed.barrier()153 print(args.rank, time.ctime(), "Epoch:", epoch)154 epoch_time = time.time()155 train_loss = train_epoch(156 model, train_loader, optimizer, scaler=scaler, epoch=epoch, loss_func=loss_func, args=args157 )158 if args.rank == 0:159 print(160 "Final training {}/{}".format(epoch, args.max_epochs - 1),161 "loss: {:.4f}".format(train_loss),162 "time {:.2f}s".format(time.time() - epoch_time),163 )164 if args.rank == 0 and writer is not None:165 writer.add_scalar("train_loss", train_loss, epoch)166 b_new_best = False167 if (epoch + 1) % args.val_every == 0:168 if args.distributed:169 torch.distributed.barrier()170 epoch_time = time.time()171 val_acc = val_epoch(172 model,173 val_loader,174 epoch=epoch,175 acc_func=acc_func,176 args=args,177 post_sigmoid=post_sigmoid,178 post_pred=post_pred,179 )180 181 if args.rank == 0:182 Dice_TC = val_acc[0]183 Dice_WT = val_acc[1]184 Dice_ET = val_acc[2]185 print(186 "Final validation stats {}/{}".format(epoch, args.max_epochs - 1),187 ", Dice_TC:",188 Dice_TC,189 ", Dice_WT:",190 Dice_WT,191 ", Dice_ET:",192 Dice_ET,193 ", time {:.2f}s".format(time.time() - epoch_time),194 )195 196 if writer is not None:197 writer.add_scalar("Mean_Val_Dice", np.mean(val_acc), epoch)198 if semantic_classes is not None:199 for val_channel_ind in range(len(semantic_classes)):200 if val_channel_ind < val_acc.size:201 writer.add_scalar(semantic_classes[val_channel_ind], val_acc[val_channel_ind], epoch)202 val_avg_acc = np.mean(val_acc)203 if val_avg_acc > val_acc_max:204 print("new best ({:.6f} --> {:.6f}). ".format(val_acc_max, val_avg_acc))205 val_acc_max = val_avg_acc206 b_new_best = True207 if args.rank == 0 and args.logdir is not None and args.save_checkpoint:208 save_checkpoint(209 model, epoch, args, best_acc=val_acc_max, optimizer=optimizer, scheduler=scheduler210 )211 if args.rank == 0 and args.logdir is not None and args.save_checkpoint:212 print("Saving")213 save_checkpoint(model, epoch, args, best_acc=val_acc_max, filename="model_final.pt")214 if b_new_best:215 print("Copying to model.pt new best model!!!!")216 shutil.copyfile(os.path.join(args.logdir, "model_final.pt"), os.path.join(args.logdir, "model.pt"))217 218 if scheduler is not None:219 scheduler.step()220 221 print("Training Finished !, Best Accuracy: ", val_acc_max)222 223 return val_acc_max224 