(epoch_number, model, opt, test_dataset, epoch, max_psnr, iters=None)
| 27 | io.imsave(path, img) |
| 28 | |
| 29 | def evaluateModel(epoch_number, model, opt, test_dataset, epoch, max_psnr, iters=None): |
| 30 | |
| 31 | model.netG.eval() |
| 32 | |
| 33 | if iters is not None: |
| 34 | eval_path = os.path.join(opt.checkpoints_dir, opt.name, 'Eval_%s_iter%d.csv' % (epoch, iters)) # define the website directory |
| 35 | else: |
| 36 | eval_path = os.path.join(opt.checkpoints_dir, opt.name, 'Eval_%s.csv' % (epoch)) # define the website directory |
| 37 | eval_results_fstr = open(eval_path, 'w') |
| 38 | eval_results = {'mask': [], 'mse': [], 'psnr': [], 'fmse':[], 'ssim':[]} |
| 39 | |
| 40 | for i, data in tqdm(enumerate(test_dataset), total=len(train_dataloader)): |
| 41 | model.set_input(data) # unpack data from data loader |
| 42 | model.test() # inference |
| 43 | visuals = model.get_current_visuals() # get image results |
| 44 | output = visuals['attentioned'] |
| 45 | real = visuals['real'] |
| 46 | |
| 47 | for i_img in range(real.size(0)): |
| 48 | gt, pred = real[i_img:i_img+1], output[i_img:i_img+1] |
| 49 | fore_nums = data['mask'][i_img].sum().item() |
| 50 | mse_score_op = mean_squared_error(util.tensor2im(pred), util.tensor2im(gt)) |
| 51 | psnr_score_op = peak_signal_noise_ratio(util.tensor2im(gt), util.tensor2im(pred), data_range=255) |
| 52 | fmse_score_op = mean_squared_error(util.tensor2im(pred), util.tensor2im(gt)) * 256 * 256 / fore_nums |
| 53 | ssim_score = ssim(util.tensor2im(pred), util.tensor2im(gt), data_range=255, channel_axis=-1) |
| 54 | |
| 55 | if epoch >= 100: |
| 56 | pred_rgb = util.tensor2im(pred) |
| 57 | img_path = data['img_path'][i_img] |
| 58 | basename, imagename = os.path.split(img_path) |
| 59 | basename = basename.split('/')[-2] |
| 60 | save_img(os.path.join('evaluate', str(epoch_number), 'results',basename, imagename.split('.')[0] + '.png'), pred_rgb) |
| 61 | |
| 62 | # update calculator |
| 63 | eval_results['mse'].append(mse_score_op) |
| 64 | eval_results['psnr'].append(psnr_score_op) |
| 65 | eval_results['fmse'].append(fmse_score_op) |
| 66 | eval_results['ssim'].append(ssim_score) |
| 67 | eval_results['mask'].append(data['mask'][i_img].mean().item()) |
| 68 | eval_results_fstr.writelines('%s,%.3f,%.3f,%.3f\n' % (data['img_path'][i_img], eval_results['mask'][-1],mse_score_op, psnr_score_op)) |
| 69 | if i + 1 % 100 == 0: |
| 70 | # print('%d images have been processed' % (i + 1)) |
| 71 | eval_results_fstr.flush() |
| 72 | eval_results_fstr.flush() |
| 73 | eval_results_fstr.close() |
| 74 | |
| 75 | all_mse, all_psnr, all_fmse, all_ssim = calculateMean(eval_results['mse']), calculateMean(eval_results['psnr']), calculateMean(eval_results['fmse']), calculateMean(eval_results['ssim']) |
| 76 | |
| 77 | print('MSE:%.3f, PSNR:%.3f, fMSE:%.3f, SSIM:%.3f' % (all_mse, all_psnr, all_fmse, all_ssim)) |
| 78 | model.netG.train() |
| 79 | return all_mse, all_psnr, resolveResults(eval_results) |
| 80 | |
| 81 | def resolveResults(results): |
| 82 | interval_metrics = {} |
no test coverage detected