ReflectionEraser/ReflectionEraserApp
0
1from __future__ import print_function #if running Python2 makes it work like python3
2
3import math
4import os
5import sys
6import time
7
8import numpy as np
9import torch
10import torch.nn as nn
11import yaml #Imports PyYAML for parsing YAML files.
12from PIL import Image #Imports the Python Imaging Library for image processing.
13from skimage.metrics import peak_signal_noise_ratio as compare_psnr #Imports the PSNR metric from skimage to quantify reconstruction quality for images and video subject to lossy compression. Higher is better
14from skimage.metrics import structural_similarity #Imports the SSIM metric from skimage.A full-reference image quality evaluation index that measures image similarity from three aspects: brightness, contrast, and structure.
15
16
17def get_config(config):#apth to configuration file
18 with open(config, 'r') as stream: #opens the file specified by the config variable in read mode and assigned to variable stream
19 return yaml.load(stream) #call loads and parses the YAML content from the file object stream and returns the parsed content.
20
21
22# Converts a Tensor into a Numpy array then desired imtype
23# |imtype|: the desired type of the converted numpy array
24def tensor2im(image_tensor, imtype=np.uint8):
25 image_numpy = image_tensor[0].cpu().float().numpy() #image_tensor[0]: first image in batch in tensor
26 #.cpu() moves tensor to CPU if in GPU
27 #.float() converts to float
28 #.numpy() converts to numpy
29 if image_numpy.shape[0] == 1: #.shape[0] access channels ,1 = grayscale, 3=RGB
30 image_numpy = np.tile(image_numpy, (3, 1, 1))
31 #np.tile(A,reps):Construct an array by repeating A the number of times given by reps.
32 #repeating graysvale image here in 3 channels to get RGB image
33 image_numpy = (np.transpose(image_numpy, (1, 2, 0)) + 1) / 2.0 * 255.0
34 #np.transpose(image_numpy, (1, 2, 0)):The image is transposed from the shape (C, H, W) to (H, W, C).
35 # + 1: shift the pixel intensity range from [-1, 1] to [0, 2].
36 # / 2.0: This scales the values to the range [0, 1].
37 # * 255.0: This scales the values to the range [0, 255].
38 image_numpy = image_numpy.astype(imtype) #numpy array recast to desired imtype
39 if image_numpy.shape[-1] == 6: #after transpose -1 is no.channels and if this is 6
40 image_numpy = np.concatenate([image_numpy[:, :, :3], image_numpy[:, :, 3:]], axis=1)
41 '''???how does this help '''
42 #recombines along width 2 grps of 0-2 and 3-5
43
44 if image_numpy.shape[-1] == 7:#if transposed image has 7 channels
45 edge_map = np.tile(image_numpy[:, :, 6:7], (1, 1, 3))
46 '''???Tiles (repeats) the seventh channel to create a 3-channel edge_map along the channel axis.'''
47 image_numpy = np.concatenate([image_numpy[:, :, :3], image_numpy[:, :, 3:6], edge_map], axis=1)
48 '''???how does this help '''
49 #recombines along width 2 grps of 0-2 and 3-6
50 return image_numpy
51
52
53def tensor2numpy(image_tensor):
54 image_numpy = torch.squeeze(image_tensor).cpu().float().numpy() #.squeeze(): removes all dimensions of size 1 from the shape of image_tensor
55 #PyTorch tensors can have singleton dimensions, which are dimensions with size 1. These dimensions might not carry meaningful information but are retained due to tensor operations.
56 #proabably rmeove number of batchs since 1 batch in 1 img (batch,channels,height,width)
57 image_numpy = (np.transpose(image_numpy, (1, 2, 0)) + 1) / 2.0 * 255.0
58 #np.transpose(image_numpy, (1, 2, 0)):The image is transposed from the shape (C, H, W) to (H, W, C).
59 # + 1: shift the pixel intensity range from [-1, 1] to [0, 2].
60 # / 2.0: This scales the values to the range [0, 1].
61 # * 255.0: This scales the values to the range [0, 255].
62 image_numpy = image_numpy.astype(np.float32)
63 #numpy array recast to float
64 return image_numpy
65
66
67# Get model list for resume
68#retrieve a specific model checkpoint file based on dirname, epoch,key
69def get_model_list(dirname, key, epoch=None):
70 #key: main file_name
71 #epoch to indicate at which epoch the model was saved.Choosing epoch allows you to continue training from a specific point or to evaluate the model's performance at different stages of training
72 if epoch is None:
73 return os.path.join(dirname, key + '_latest.pt') # return latest checkpoint by default
74 if os.path.exists(dirname) is False: #if directory doesnt exist no checkpoint files can be retrieved.
75 return None
76
77 print(dirname, key)
78
79 #list of generated checkpoints after specific epohs during traininf
80 gen_models = [os.path.join(dirname, f) for f in os.listdir(dirname) if
81 os.path.isfile(os.path.join(dirname, f)) and ".pt" in f and 'latest' not in f]
82 #isfile():regular file,not a directory or a special file
83 #.pt type file with epoch number specifid, latest means epoch=None
84 epoch_index = [int(os.path.basename(model_name).split('_')[-2]) for model_name in gen_models if 'latest' not in model_name]
85 #os.path.basename(model_name): extracts the filename from the full path.
86 #split('_')[-2]: ['dsrnet', 's', 'epoch14.pt']-> s (second last)
87 '''??? why not -1 -> to get 14 chatgpt says int(s) will raise VauleError'''
88 print('[i] available epoch list: %s' % epoch_index, gen_models)
89 i = epoch_index.index(int(epoch))
90
91 return gen_models[i]
92
93# to preprocess a batch of images for compatibility with models pretrained on the ImageNet dataset
94def vgg_preprocess(batch):
95 # normalize using imagenet mean and std
96 #ImageNet normalization helps preprocess images to have zero mean and unit variance across each channel (RGB).
97 mean = batch.new(batch.size()) # Pytorch function creates new tensors (mean and std) with the same shape as the input batch.
98 std = batch.new(batch.size())
99 mean[:, 0, :, :] = 0.485 #assigned to all elements in the first channel of every image in the batch.
100 mean[:, 1, :, :] = 0.456 #2nd channel
101 mean[:, 2, :, :] = 0.406 #3rd channel
102 std[:, 0, :, :] = 0.229 #1st channel
103 std[:, 1, :, :] = 0.224 #2nd channel
104 std[:, 2, :, :] = 0.225 #3rd channel
105 batch = (batch + 1) / 2 #pixel values shifted [-1,1]->[0,2]->[0,1]
106 batch -= mean
107 batch = batch / std
108 # each pixel value in batch will have zero mean and unit variance
109 return batch
110
111
112
113 #print diagnostic information about the gradients of a neural network
114 #useful for inspecting the average magnitude of gradients during training or optimization of a neural network.
115 # It helps in diagnosing potential issues like vanishing or exploding gradients,
116def diagnose_network(net, name='network'): #name: Optional name for the network (default is 'network'), used in printing diagnostic information.
117 mean = 0.0 #average mean absolute gradient across all parameters that have gradients.
118 count = 0 #no. of paramters with gradients
119 for param in net.parameters():
120 if param.grad is not None: #trainable Parameters without gradients won't contribute to mean.
121 #e.g frozen layers, non trainable paramters,bias,batch normalization parameters
122 mean += torch.mean(torch.abs(param.grad.data)) #mean absolute value of the gradient of the current parameter param
123 count += 1
124 if count > 0:
125 mean = mean / count
126 print(name)
127 print(mean)
128
129
130def save_image(image_numpy, image_path): #create a PIL Image object from the numpy array and then saves it using the save() method of the PIL Image object.
131 image_pil = Image.fromarray(image_numpy)
132 image_pil.save(image_path)
133
134#print shape and statistics of a numpy array
135def print_numpy(x, val=True, shp=False): #val is statistics shp is shape
136 x = x.astype(np.float64)
137 if shp:
138 print('shape,', x.shape)
139 if val:
140 x = x.flatten() #converts nD to 1D
141 print('mean = %3.3f, min = %3.3f, max = %3.3f, median = %3.3f, std=%3.3f' % (
142 np.mean(x), np.min(x), np.max(x), np.median(x), np.std(x)))
143
144
145def mkdirs(paths):
146 if isinstance(paths, list) and not isinstance(paths, str):
147 #if list makes path for each string
148 for path in paths:
149 mkdir(path)
150 else: #if string makes path
151 mkdir(paths)
152
153
154def mkdir(path):
155 if not os.path.exists(path):
156 os.makedirs(path)
157
158
159def set_opt_param(optimizer, key, value):
160 for group in optimizer.param_groups:
161 group[key] = value
162'???'
163
164
165def vis(x):#display an image in image viewer of system regardless if tensor or numpy array
166 if isinstance(x, torch.Tensor):
167 Image.fromarray(tensor2im(x)).show()
168 elif isinstance(x, np.ndarray):
169 Image.fromarray(x.astype(np.uint8)).show()
170 else:
171 raise NotImplementedError('vis for type [%s] is not implemented', type(x))
172
173
174"""tensorboard"""
175from tensorboardX import SummaryWriter #for data logging
176from datetime import datetime
177
178
179def get_summary_writer(log_dir): #log_dir:directy to store logs
180 if not os.path.exists(log_dir):
181 os.mkdir(log_dir)
182 log_dir = os.path.join(log_dir, datetime.now().strftime('%b%d_%H-%M-%S') + '_' + socket.gethostname())
183 if not os.path.exists(log_dir):
184 os.mkdir(log_dir)
185 writer = SummaryWriter(log_dir) #each time you run your experiments, logs are saved in a new directory with current time, date and machine hostname
186 return writer
187
188
189
190# keeps track of running avergae of metrics
191class AverageMeters(object): #all classes inherit from object
192 def __init__(self, dic=None, total_num=None):
193 #dic: dictionary of metrics
194 #total_num:track of the counts of each metric)
195 self.dic = dic or {}
196 #(mingcv) self.total_num = total_num
197 self.total_num = total_num or {}
198
199
200 def update(self, new_dic): # method appends self.dic and updates total_num with new values.
201 for key in new_dic:
202 if not key in self.dic:
203 self.dic[key] = new_dic[key]
204 self.total_num[key] = 1
205 else:
206 self.dic[key] += new_dic[key]
207 self.total_num[key] += 1
208 # (mingcv)self.total_num += 1
209
210 def __getitem__(self, key): #overriding [], returns average value for a metric
211 return self.dic[key] / self.total_num[key]
212
213 def __str__(self): #overrider str() and print() to show key and values formatted
214 keys = sorted(self.keys())
215 res = ''
216 for key in keys:
217 res += (key + ': %.4f' % self[key] + ' | ')
218 return res
219
220 def keys(self): #returns all the metric names
221 return self.dic.keys()
222
223'''???why _loss'''
224def write_loss(writer, prefix, avg_meters, iteration):
225 #writer is the writer object for logging,SummaryWriter instance
226 #prefix is keyword added before each metric e.g. "train" it logs "accuracy": 85.4 as "train/accuracy"
227 #avg_meters: A dictionary argument where keys are metric names (like ‘loss’, ‘accuracy’) and values are their corresponding values at the current iteration (e.g., 0.23, 85.4).
228 #iteration: An integer argument representing the current iteration or step number in the training or evaluation process.
229 for key in avg_meters.keys():
230 meter = avg_meters[key] # for each key, we get the corresponding value from the avg_meters dictionary and assign it to the variable meter.
231 writer.add_scalar(os.path.join(prefix, key), meter, iteration)
232 # writer.add_scalar(): This method call logs the scalar value (metric) to the writer
233'''
234avg_meters = {
235 “loss”: 0.23,
236 “accuracy”: 85.4,
237 “precision”: 88.1
238}
239write_loss(writer, “train”, avg_meters, 10)
240
241Logging: train/loss = 0.23 at step 10
242Logging: train/accuracy = 85.4 at step 10
243Logging: train/precision = 88.1 at step 10
244
245'''
246"""(Mingcv)progress bar"""
247import socket
248
249# _, term_width = os.popen('stty size', 'r').read().split()
250term_width = 136 #hardcoded terminal width to 136 characters instead of dymeically reading using stty siz
251
252TOTAL_BAR_LENGTH = 65. #length of the progress bar.
253last_time = time.time() #when progress is finished
254begin_time = last_time # used to calculate the elapsed time for each step and the total process.
255
256
257def progress_bar(current, total, msg=None):
258 #current: The current progress count.
259 #total: The total count that represents 100% progress.
260 #msg: An optional message to display alongside the progress bar.
261 global last_time, begin_time # for persistence in step_time and total_time
262 if current == 0:
263 begin_time = time.time() # Reset for new bar.
264
265 cur_len = int(TOTAL_BAR_LENGTH * current / total)
266 rest_len = int(TOTAL_BAR_LENGTH - cur_len) - 1
267
268 sys.stdout.write(' [')
269 for i in range(cur_len):
270 sys.stdout.write('=')
271 sys.stdout.write('>')
272 for i in range(rest_len):
273 sys.stdout.write('.')
274 sys.stdout.write(']')
275
276 cur_time = time.time()
277 step_time = cur_time - last_time
278 last_time = cur_time
279 tot_time = cur_time - begin_time
280
281 L = []
282 L.append(' Step: %s' % format_time(step_time))
283 L.append(' | Tot: %s' % format_time(tot_time))
284 if msg:
285 L.append(' | ' + msg)
286
287 msg = ''.join(L)
288 sys.stdout.write(msg)
289 for i in range(term_width - int(TOTAL_BAR_LENGTH) - len(msg) - 3):
290 sys.stdout.write(' ')
291
292 #(mingcv) Go back to the center of the bar.
293 for i in range(term_width - int(TOTAL_BAR_LENGTH / 2) + 2):
294 sys.stdout.write('\b')
295 sys.stdout.write(' %d/%d ' % (current + 1, total))
296
297 if current < total - 1:
298 sys.stdout.write('\r')
299 else:
300 sys.stdout.write('\n')
301 sys.stdout.flush()
302
303
304def format_time(seconds): #convert seconds to days,hourse,mintures,seconds,milliseconds
305 days = int(seconds / 3600 / 24)
306 seconds = seconds - days * 3600 * 24
307 hours = int(seconds / 3600)
308 seconds = seconds - hours * 3600
309 minutes = int(seconds / 60)
310 seconds = seconds - minutes * 60
311 secondsf = int(seconds)
312 seconds = seconds - secondsf
313 millis = int(seconds * 1000)
314
315 f = ''
316 i = 1
317 if days > 0:
318 f += str(days) + 'D'
319 i += 1
320 if hours > 0 and i <= 2:
321 f += str(hours) + 'h'
322 i += 1
323 if minutes > 0 and i <= 2:
324 f += str(minutes) + 'm'
325 i += 1
326 if secondsf > 0 and i <= 2:
327 f += str(secondsf) + 's'
328 i += 1
329 if millis > 0 and i <= 2:
330 f += str(millis) + 'ms'
331 i += 1
332 if f == '':
333 f = '0ms'
334 return f
335
336
337def parse_args(args): #converts numeric args seperated by commas into a [] of ints
338 str_args = args.split(',') #str1,atr2,str3 becomes [a,b,c]
339 parsed_args = []
340 for str_arg in str_args:
341 arg = int(str_arg)
342 if arg >= 0:
343 parsed_args.append(arg)
344 return parsed_args
345
346#In order to overcome vanishing/explosding gradient, Xavier initialization was introduced. It tries to keep variance of all the layers equal but assumes linear actiivation
347# Kaiming He initialization, which takes activation function into account.
348
349def weights_init_kaiming(m): #layer or module in NN
350 classname = m.__class__.__name__ # gets the class name of the layer m using its __class__ attribute and __name__ property.
351 if classname.find('Conv') != -1: #returns -1 if the substring is not found, So here if found
352 nn.init.kaiming_normal(m.weight.data, a=0, mode='fan_in')
353
354 #nn.init.kaiming_normal:Fill the input Tensor(m.weight.data) with values using a Kaiming normal distribution.The method is described in Delving deep into rectifiers: Surpassing human-level performance on ImageNet classification - He, K. et al. (2015). The resulting tensor will have values sampled from N(0,std^2) where std=gain/sqrt(fan_mode)
355 #a (float) – the negative slope of the rectifier used after this layer only used with Leaky RelU (0-not used)
356 #mode: fan_in(default) or fan_out.Choosing 'fan_in' preserves the magnitude of the variance of the weights in the forward pass.
357 #here intialising for relu,preserving variance of weights in forward pass
358 elif classname.find('Linear') != -1:
359 nn.init.kaiming_normal(m.weight.data, a=0, mode='fan_in')
360 # here intialising for relu,preserving variance of weights in forward pass
361 elif classname.find('BatchNorm') != -1:
362 # nn.init.uniform(m.weight.data, 1.0, 0.02)
363 m.weight.data.normal_(mean=0, std=math.sqrt(2. / 9. / 64.)).clamp_(-0.025, 0.025)
364 nn.init.constant(m.bias.data, 0.0)
365 "??? why std=math.sqrt(2. / 9. / 64.)"
366 # .clamp_(-0.025, 0.025): Clamps (limits) the values in the tensor to be between -0.025 and 0.025.
367 #bias of the batch normalization layer intiliazed to a constant value of 0.0.
368
369
370def batch_PSNR(img, imclean, data_range): #calculates avg PSNR for 1 batch of images
371 #img batch of images to be evaluated with noise
372 #imclean: batch of ground truth images no noise
373 Img = img.data.cpu().numpy().astype(np.float32) #converts img to float numpy array in cpu
374 #- img.data: Detaches the data from the computation graph.
375 #cpu(): Transfers the tensor from GPU to CPU if it’s on GPU.
376 #numpy(): Converts the tensor to a NumPy array.
377 #astype(np.float32): Converts the array type to float32.
378 Iclean = imclean.data.cpu().numpy().astype(np.float32) #same for ground truth
379 PSNR = 0
380 for i in range(Img.shape[0]): #for all batches add the skmetric.psnr()
381 PSNR += compare_psnr(Iclean[i, :, :, :], Img[i, :, :, :], data_range=data_range)
382 #data range: maximum - inimum possible values
383 return PSNR / Img.shape[0] #average psnr per batch
384
385#calculates avg SSIM for 1 batch of images
386def batch_SSIM(img, imclean):
387 Img = img.data.cpu().permute(0, 2, 3, 1).numpy().astype(np.float32)
388 #converts img to float numpy array in cpu
389 #img.data: Detaches the data from the computation graph.
390 #cpu(): Transfers the tensor from GPU to CPU if it’s on GPU.
391 #numpy(): Converts the tensor to a NumPy array.
392 #astype(np.float32): Converts the array type to float32.
393 Iclean = imclean.data.cpu().permute(0, 2, 3, 1).numpy().astype(np.float32)
394 #permute(0, 2, 3, 1): Changes the order of dimensions from (batch_size, channels, height, width) to (batch_size, height, width, channels).
395 SSIM = 0
396
397 for i in range(Img.shape[0]): #for all batches add the skmetric.ssim()
398 SSIM += structural_similarity(Iclean[i, :, :, :], Img[i, :, :, :], win_size=11,
399 multichannel=True, data_range=1)
400 #win_size:The side-length of the sliding window used in comparison. Must be an odd value
401 #multichannel deprecated ...equivalent to channel-axis
402 "??? multichannel=True equal to what number in channel-axis"
403 return SSIM / Img.shape[0] #average ssim for1 batch
404
405
406def data_augmentation(image, mode): #augment given image based on mode passed by manipulating numpy array
407 out = np.transpose(image, (1, 2, 0))
408 if mode == 0:
409 #(mingcv)original
410 out = out
411 elif mode == 1:
412 #(mingcv)flip up and down
413 out = np.flipud(out)
414 elif mode == 2:
415 #(mingcv)rotate counterwise 90 degree
416 out = np.rot90(out)
417 elif mode == 3:
418 #(mingcv)rotate 90 degree and flip up and down
419 out = np.rot90(out)
420 out = np.flipud(out)
421 elif mode == 4:
422 #(mingcv)rotate 180 degree
423 out = np.rot90(out, k=2)
424 elif mode == 5:
425 #(mingcv)rotate 180 degree and flip
426 out = np.rot90(out, k=2)
427 out = np.flipud(out)
428 elif mode == 6:
429 #(mingcv)rotate 270 degree
430 out = np.rot90(out, k=3)
431 elif mode == 7:
432 #(mingcv)rotate 270 degree and flip
433 out = np.rot90(out, k=3)
434 out = np.flipud(out)
435 return np.transpose(out, (2, 0, 1))
436 