CLYang617/RemoteSensingChangeDetection-RSCD.HA2F
0
1import sys2 3from model.trainer import Trainer4 5sys.path.insert(0, '.')6 7import torch8import torch.nn.functional as F9import torch.backends.cudnn as cudnn10from torch.nn.parallel import gather11import torch.optim.lr_scheduler12 13import dataset.dataset as myDataLoader14import dataset.Transforms as myTransforms15from model.metric_tool import ConfuseMatrixMeter16from model.utils import BCEDiceLoss, init_seed, adjust_learning_rate17 18import os, time19import numpy as np20from argparse import ArgumentParser21 22 23 24@torch.no_grad()25def val(args, val_loader, model):26 model.eval()27 28 salEvalVal = ConfuseMatrixMeter(n_class=2)29 30 epoch_loss = []31 32 total_batches = len(val_loader)33 print(len(val_loader))34 for iter, batched_inputs in enumerate(val_loader):35 36 img, target = batched_inputs37 pre_img = img[:, 0:3]38 post_img = img[:, 3:6]39 40 start_time = time.time()41 42 if args.onGPU == True:43 pre_img = pre_img.cuda()44 target = target.cuda()45 post_img = post_img.cuda()46 47 pre_img_var = torch.autograd.Variable(pre_img).float()48 post_img_var = torch.autograd.Variable(post_img).float()49 target_var = torch.autograd.Variable(target).float()50 51 # run the mdoel52 output = model(pre_img_var, post_img_var)53 loss = BCEDiceLoss(output, target_var)54 55 pred = torch.where(output > 0.5, torch.ones_like(output), torch.zeros_like(output)).long()56 57 # torch.cuda.synchronize()58 time_taken = time.time() - start_time59 60 epoch_loss.append(loss.data.item())61 62 # compute the confusion matrix63 if args.onGPU and torch.cuda.device_count() > 1:64 output = gather(pred, 0, dim=0)65 # salEvalVal.addBatch(pred, target_var)66 f1 = salEvalVal.update_cm(pr=pred.cpu().numpy(), gt=target_var.cpu().numpy())67 if iter % 5 == 0:68 print('\r[%d/%d] F1: %3f loss: %.3f time: %.3f' % (iter, total_batches, f1, loss.data.item(), time_taken),69 end='')70 71 average_epoch_loss_val = sum(epoch_loss) / len(epoch_loss)72 scores = salEvalVal.get_scores()73 74 return average_epoch_loss_val, scores75 76 77def train(args, train_loader, model, optimizer, epoch, max_batches, cur_iter=0, lr_factor=1.):78 # switch to train mode79 model.train()80 81 salEvalVal = ConfuseMatrixMeter(n_class=2)82 epoch_loss = []83 84 for iter, batched_inputs in enumerate(train_loader):85 86 img, target = batched_inputs87 pre_img = img[:, 0:3]88 post_img = img[:, 3:6]89 90 start_time = time.time()91 92 # adjust the learning rate93 lr = adjust_learning_rate(args, optimizer, epoch, iter + cur_iter, max_batches, lr_factor=lr_factor)94 95 if args.onGPU == True:96 pre_img = pre_img.cuda()97 target = target.cuda()98 post_img = post_img.cuda()99 100 pre_img_var = torch.autograd.Variable(pre_img).float()101 post_img_var = torch.autograd.Variable(post_img).float()102 target_var = torch.autograd.Variable(target).float()103 104 # run the model105 output = model(pre_img_var, post_img_var)106 loss = BCEDiceLoss(output, target_var)107 108 pred = torch.where(output > 0.5, torch.ones_like(output), torch.zeros_like(output)).long()109 110 optimizer.zero_grad()111 loss.backward()112 optimizer.step()113 114 epoch_loss.append(loss.data.item())115 time_taken = time.time() - start_time116 res_time = (max_batches * args.max_epochs - iter - cur_iter) * time_taken / 3600117 118 if args.onGPU and torch.cuda.device_count() > 1:119 output = gather(pred, 0, dim=0)120 121 # Computing F-measure and IoU on GPU122 with torch.no_grad():123 f1 = salEvalVal.update_cm(pr=pred.cpu().numpy(), gt=target_var.cpu().numpy())124 125 if iter % 5 == 0:126 print('\riteration: [%d/%d] f1: %.3f lr: %.7f loss: %.3f time:%.3f h' % (127 iter + cur_iter, max_batches * args.max_epochs, f1, lr, loss.data.item(),128 res_time),129 end='')130 131 average_epoch_loss_train = sum(epoch_loss) / len(epoch_loss)132 scores = salEvalVal.get_scores()133 134 return average_epoch_loss_train, scores, lr135 136 137def trainValidateSegmentation(args):138 os.environ['CUDA_VISIBLE_DEVICES'] = str(args.gpu_id)139 140 torch.backends.cudnn.benchmark = True141 142 init_seed(args.seed)143 144 args.savedir = args.savedir + '_' + args.file_root + '_iter_' + str(args.max_steps) + '_lr_' + str(args.lr) + '/'145 146 if args.file_root == 'LEVIR':147 args.file_root = './levir_cd_256'148 elif args.file_root == 'WHU':149 args.file_root = './whu_cd_256'150 elif args.file_root == 'CLCD':151 args.file_root = './clcd_256'152 elif args.file_root == 'SYSU':153 args.file_root = './sysu_256'154 elif args.file_root == 'OSCD':155 args.file_root = 'oscd_256'156 else:157 raise TypeError('%s has not defined' % args.file_root)158 159 if not os.path.exists(args.savedir):160 os.makedirs(args.savedir)161 162 163 model = Trainer(args.model_type).float()164 if args.onGPU:165 model = model.cuda()166 167 # mean = [0.5, 0.5, 0.5, 0.5, 0.5, 0.5]168 # std = [0.5, 0.5, 0.5, 0.5, 0.5, 0.5]169 170 mean = [0.406, 0.456, 0.485, 0.406, 0.456, 0.485]171 std = [0.225, 0.224, 0.229, 0.225, 0.224, 0.229]172 173 # compose the data with transforms174 trainDataset_main = myTransforms.Compose([175 myTransforms.Normalize(mean=mean, std=std),176 myTransforms.Scale(args.inWidth, args.inHeight),177 myTransforms.RandomCropResize(int(7. / 224. * args.inWidth)),178 myTransforms.RandomFlip(),179 myTransforms.RandomExchange(),180 myTransforms.ToTensor()181 ])182 183 valDataset = myTransforms.Compose([184 myTransforms.Normalize(mean=mean, std=std),185 myTransforms.Scale(args.inWidth, args.inHeight),186 myTransforms.ToTensor()187 ])188 189 train_data = myDataLoader.Dataset(file_root=args.file_root, mode="train", transform=trainDataset_main)190 191 trainLoader = torch.utils.data.DataLoader(192 train_data,193 batch_size=args.batch_size, shuffle=True,194 num_workers=args.num_workers, pin_memory=True, drop_last=False195 )196 197 test_data = myDataLoader.Dataset(file_root=args.file_root, mode="test", transform=valDataset)198 testLoader = torch.utils.data.DataLoader(199 test_data, shuffle=False,200 batch_size=args.batch_size, num_workers=args.num_workers, pin_memory=True)201 202 203 max_batches = len(trainLoader)204 print('For each epoch, we have {} batches'.format(max_batches))205 206 if args.onGPU:207 cudnn.benchmark = True208 209 args.max_epochs = int(np.ceil(args.max_steps / max_batches))210 start_epoch = 0211 cur_iter = 0212 max_F1_val = 0213 214 if args.resume is not None:215 args.resume = args.savedir + 'checkpoint.pth.tar'216 if os.path.isfile(args.resume):217 print("=> loading checkpoint '{}'".format(args.resume))218 checkpoint = torch.load(args.resume)219 start_epoch = checkpoint['epoch']220 cur_iter = start_epoch * len(trainLoader)221 # args.lr = checkpoint['lr']222 model.load_state_dict(checkpoint['state_dict'])223 print("=> loaded checkpoint '{}' (epoch {})"224 .format(args.resume, checkpoint['epoch']))225 else:226 print("=> no checkpoint found at '{}'".format(args.resume))227 228 logFileLoc = args.savedir + args.logFile229 if os.path.isfile(logFileLoc):230 logger = open(logFileLoc, 'a')231 else:232 logger = open(logFileLoc, 'w')233 logger.write(234 "\n%s\t%s\t%s\t%s\t%s\t%s\t%s" % ('Epoch', 'Kappa (val)', 'IoU (val)', 'F1 (val)', 'R (val)', 'P (val)', 'OA (val)'))235 logger.flush()236 237 optimizer = torch.optim.Adam(model.parameters(), args.lr, (0.9, 0.99), eps=1e-08, weight_decay=1e-4)238 239 for epoch in range(start_epoch, args.max_epochs):240 lossTr, score_tr, lr = \241 train(args, trainLoader, model, optimizer, epoch, max_batches, cur_iter)242 cur_iter += len(trainLoader)243 244 torch.cuda.empty_cache()245 246 # evaluate on validation set247 if epoch == 0:248 continue249 250 lossVal, score_val = val(args, testLoader, model)251 torch.cuda.empty_cache()252 logger.write("\n%d\t\t%.4f\t\t%.4f\t\t%.4f\t\t%.4f\t\t%.4f\t\t%.4f" % (epoch, score_val['Kappa'], score_val['IoU'],253 score_val['F1'], score_val['recall'],254 score_val['precision'], score_val['OA']))255 logger.flush()256 257 torch.save({258 'epoch': epoch + 1,259 'arch': str(model),260 'state_dict': model.state_dict(),261 'optimizer': optimizer.state_dict(),262 'lossTr': lossTr,263 'lossVal': lossVal,264 'F_Tr': score_tr['F1'],265 'F_val': score_val['F1'],266 'lr': lr267 }, args.savedir + 'checkpoint.pth.tar')268 269 # save the model also270 model_file_name = args.savedir + 'best_model.pth'271 if epoch % 1 == 0 and max_F1_val <= score_val['F1']:272 max_F1_val = score_val['F1']273 torch.save(model.state_dict(), model_file_name)274 275 print("Epoch " + str(epoch) + ': Details')276 print("\nEpoch No. %d:\tTrain Loss = %.4f\tVal Loss = %.4f\t F1(tr) = %.4f\t F1(val) = %.4f" \277 % (epoch, lossTr, lossVal, score_tr['F1'], score_val['F1']))278 torch.cuda.empty_cache()279 280 state_dict = torch.load(model_file_name)281 model.load_state_dict(state_dict)282 283 loss_test, score_test = val(args, testLoader, model)284 print("\nTest :\t Kappa (te) = %.4f\t IoU (te) = %.4f\t F1 (te) = %.4f\t R (te) = %.4f\t P (te) = %.4f" \285 % (score_test['Kappa'], score_test['IoU'], score_test['F1'], score_test['recall'], score_test['precision']))286 logger.write("\n%s\t\t%.4f\t\t%.4f\t\t%.4f\t\t%.4f\t\t%.4f" % ('Test', score_test['Kappa'], score_test['IoU'],287 score_test['F1'], score_test['recall'],288 score_test['precision'], score_test['OA']))289 logger.flush()290 logger.close()291 292 293if __name__ == '__main__':294 parser = ArgumentParser()295 parser.add_argument('--file_root', default="LEVIR", help='Data directory | LEVIR | WHU | CLCD | SYSU | OSCD ')296 parser.add_argument('--inWidth', type=int, default=256, help='Width of RGB image')297 parser.add_argument('--inHeight', type=int, default=256, help='Height of RGB image')298 parser.add_argument('--max_steps', type=int, default=80000, help='Max. number of iterations')299 parser.add_argument('--num_workers', type=int, default=4, help='No. of parallel threads')300 parser.add_argument('--model_type', type=str, default='small', help='select vit model type | tiny | small')301 parser.add_argument('--batch_size', type=int, default=16, help='Batch size')302 parser.add_argument('--step_loss', type=int, default=100, help='Decrease learning rate after how many epochs')303 parser.add_argument('--lr', type=float, default=2e-4, help='Initial learning rate')304 parser.add_argument('--lr_mode', default='poly', help='Learning rate policy, step or poly')305 parser.add_argument('--seed', default=16, help='initialization seed number')306 parser.add_argument('--savedir', default='./results', help='Directory to save the results')307 parser.add_argument('--resume', default=None, help='Use this checkpoint to continue training | '308 './results_ep100/checkpoint.pth.tar')309 parser.add_argument('--logFile', default='trainValLog.txt',310 help='File that stores the training and validation logs')311 parser.add_argument('--onGPU', default=True, type=lambda x: (str(x).lower() == 'true'),312 help='Run on CPU or GPU. If TRUE, then GPU.')313 parser.add_argument('--gpu_id', default=0, type=int, help='GPU id number')314 315 args = parser.parse_args()316 print('Called with args:')317 print(args)318 319 trainValidateSegmentation(args)320 