MCPcopy Create free account
hub / github.com/chenhaoxing/HDNet / evaluateModel

Function evaluateModel

train_evaluate.py:29–79  ·  view source on GitHub ↗
(epoch_number, model, opt, test_dataset, epoch, max_psnr, iters=None)

Source from the content-addressed store, hash-verified

27 io.imsave(path, img)
28
29def 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
81def resolveResults(results):
82 interval_metrics = {}

Callers 1

train_evaluate.pyFile · 0.85

Calls 7

save_imgFunction · 0.85
calculateMeanFunction · 0.85
resolveResultsFunction · 0.85
evalMethod · 0.80
testMethod · 0.80
get_current_visualsMethod · 0.80
set_inputMethod · 0.45

Tested by

no test coverage detected