CoolFace
Modelpublic

Jupitern52/TextBraTS

sourceHugging Facemitupdated 1y agoView on Hugging Face
2likes
trainer.py224 linesDownload Raw Back to root
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