-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathplot_psnr.py
More file actions
52 lines (47 loc) · 1.59 KB
/
Copy pathplot_psnr.py
File metadata and controls
52 lines (47 loc) · 1.59 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
import matplotlib.pyplot as plt
import numpy as np
import cv2
import glob
import argparse
parser=argparse.ArgumentParser()
parser.add_argument('--pred_dirs',type=str,nargs='+',help='Give list of pred directories')
parser.add_argument('--gt_dirs',type=str,nargs='+',help='Give list of gt directories')
args=parser.parse_args()
def psnr(pred, gt,normalize=True):
pred=pred.astype(np.float32)
gt=gt.astype(np.float32)
if normalize:
pred=pred/255.0
gt=gt/255.0
mse=np.mean((pred-gt)**2)
psnr=10*np.log10(1/mse)
return psnr
def psnr_dir(pred_dir,gt_dir,normalize=True):
pred_list=glob.glob(pred_dir+"/*.png")
gt_list=glob.glob(gt_dir+"/*.png")
pred_list.sort()
gt=cv2.imread(gt_list[0])
psnr_list=[]
for pred_path in pred_list:
pred=cv2.imread(pred_path)
psnr_list.append(psnr(pred,gt,normalize))
return np.array(psnr_list)
if __name__=='__main__':
gt_dir=args.gt_dirs[0]
min_len=1e7
colors=['r','g','b','y','c','m','k']
for pred_dir in args.pred_dirs:
psnr_list=psnr_dir(pred_dir,gt_dir)
if len(psnr_list)<min_len:
min_len=len(psnr_list)
for i,pred_dir in enumerate(args.pred_dirs):
psnr_list=psnr_dir(pred_dir,gt_dir)
# print(f"PSNR for {pred_dir} is {np.mean(psnr_list)}")
x_axis=np.arange(0,min_len)*40
plt.plot(x_axis,psnr_list[:min_len],'-o',label=pred_dir,color=colors[i])
print(f"MEAN_PSNR for {pred_dir}:",psnr_list[-1])
plt.title('PSNR vs Epochs')
plt.xlabel('Epochs')
plt.ylabel('PSNR')
plt.legend()
plt.savefig('psnr.png')