Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
util.py436 linesDownload Raw Back to util
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