-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathexperiment_sidd.py
More file actions
106 lines (92 loc) · 3.98 KB
/
Copy pathexperiment_sidd.py
File metadata and controls
106 lines (92 loc) · 3.98 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
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
import argparse
import pprint
import json
parser = argparse.ArgumentParser()
parser.add_argument('--gpu', type=bool, default=True)
parser.add_argument('--block_size', type=int, default=16)
parser.add_argument('--num_block_per_row', type=int, default=21)
parser.add_argument('--overlap', type=int, default=4)
parser.add_argument('--use_pretrain', type=int, default=1)
parser.add_argument('--file_batch', type=int, default=5)
parser.add_argument('--epoch', type=int, default=50)
parser.add_argument('--num_filter', type=int, default=4)
parser.add_argument('--train_files_path', type=str, default='')
parser.add_argument('--test_files_path', type=str, default='')
def main(args):
# Enable GPU
if args.gpu:
#%tensorflow_version 2.x
import tensorflow as tf
device_name = tf.test.gpu_device_name()
if device_name != '/device:GPU:0':
raise SystemError('GPU device not found')
print('Found GPU at: {}'.format(device_name))
BLOCK_SIZE = args.block_size
NUM_BLOCK = args.num_block_per_row
BLOCK_PER_IMAGE = NUM_BLOCK * NUM_BLOCK
OVERLAP = args.overlap
NUM_CLUSTER = args.num_filter
SHAPE = (BLOCK_SIZE, BLOCK_SIZE, 3)
EPOCH = args.epoch
FILE_BATCH = args.file_batch
from train import *
from evaluation import *
'''
Load Images
'''
train_files = []
valid_files = []
for root_path, sub, files in os.walk(args.train_files_path):
contents = files
contents.sort()
for f in contents:
file_path = os.path.join(root_path,f)
if os.path.isfile(file_path) and "GT" in f:
train_files.append(file_path)
if os.path.isfile(file_path) and "NOISY" in f:
valid_files.append(file_path)
if len(train_files) != len(valid_files):
raise ValueError('Train and Validation file must have same length', len(train_files), len(valid_files))
decoders = []
'''
Train Network
'''
for fb in range(0, len(train_files), FILE_BATCH):
train = train_files[fb:fb+FILE_BATCH]
valid = valid_files[fb:fb+FILE_BATCH]
clear_images = crop_square(read_images(valid))
noise_images = crop_square(read_images(train))
WIDTH = len(clear_images[0][0])
HEIGHT = len(clear_images[0])
clear_images, noise_images = gen_train_set(clear_images, noise_images, SHAPE,
BLOCK_SIZE, NUM_BLOCK, OVERLAP)
model = train_encoder(noise_images, clear_images, NUM_CLUSTER, EPOCH)
clus, label_clus = clustering(noise_images, clear_images, model.encoder, model.classifier, NUM_CLUSTER)
decoders = train_decoders(clus, label_clus, encoder, EPOCH, decoder.get_weights())
'''
Evaluation
'''
avg_psnr, avg_ssim, avg_uqi = 0, 0, 0
BATCH = 4
NUM_BLOCK = 15
for i in range(BATCH):
clear_images = sidd_test_data(args.test_validation_files_path, 'ValidationGtBlocksSrgb', i)
test_images = np.array(clear_images)
noise_images = sidd_test_data(args.test_files_path, 'ValidationNoisyBlocksSrgb', i)
WIDTH = len(clear_images[0][0])
HEIGHT = len(clear_images[0])
clear_images, noise_images = gen_train_set(clear_images, noise_images, SHAPE,
BLOCK_SIZE, NUM_BLOCK, OVERLAP)
recons_images = reconstruct_image(noise_images, model.encoder, model.classifier,
decoders, BLOCK_PER_IMAGE, WIDTH, HEIGHT, BLOCK_SIZE, OVERLAP)
avg_psnr += quality_evaluation(recons_images, test_images, metric='PSNR')
avg_ssim += quality_evaluation(recons_images, test_images, metric='SSIM')
avg_uqi += quality_evaluation(recons_images, test_images, metric='UQI')
print('***********************')
print('Overall Results')
print('PSNR: ', avg_psnr/BATCH)
print('SSIM: ', avg_ssim/BATCH)
print('UQI: ', avg_uqi/BATCH)
print('***********************')
if __name__ == '__main__':
main(parser.parse_args())