Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
base_options.py98 linesDownload Raw Back to net_options
1#configuration based on neural network
2#action=store_true: presence of this argument in the command line will set the corresponding variable to True. If the argument is not specified, the variable will be set to False.Opposite of store_false
3import sys
4
5from options.base_option import BaseOptions as Base
6from util import util
7import os
8import torch
9import numpy as np
10import random
11
12
13class BaseOptions(Base):
14    def initialize(self):
15        Base.initialize(self)
16        # experiment specifics
17        self.parser.add_argument('--inet', type=str, default='net_options',
18                                 help='chooses which architecture to use for inet.')
19        self.parser.add_argument('--weight_path', type=str, default=None, help='pretrained checkpoint to use.')
20        #Path to a pretrained checkpoint for initialization.
21        self.parser.add_argument('--init_type', type=str, default='edsr',
22                                 help='network initialization [normal|xavier|kaiming|orthogonal|uniform]')
23        #(mingcv)for network
24        self.parser.add_argument('--hyper', action='store_true',
25                                 help='if true, augment input with vgg hypercolumn feature')
26        # A hypercolumn refers to a vector formed by concatenating the outputs of all feature maps at a particular spatial location across different layers of the VGG network.
27        #Benefits: the network can leverage both low-level details and high-level semantic information simultaneously.provides a richer representation of the input image.
28        #Implementation: During the forward pass of the neural network:
29        #1.Extract feature maps from various layers of the VGG network.
30        #2.Concatenate these feature maps spatially (typically by resizing them to a common size).
31        #3.Use this concatenated feature vector as an augmented input alongside the original input to the neural network.
32
33        self.initialized = True
34
35    def parse(self):
36        if not self.initialized:
37            self.initialize()
38        self.opt = self.parser.parse_args() #Stores parsed command-line arguments.
39        self.opt.isTrain = self.isTrain  # (mingcv)train or test
40        #to use train_options.py or not
41        if self.opt.seed == 0:#for reproducibility across modules.
42            seed = random.randrange(2 ** 12 - 1)
43            self.opt.seed = seed
44
45        torch.backends.cudnn.deterministic = True
46        #PyTorch uses CUDA libraries (like cuDNN) for GPU-accelerated computations. Setting torch.backends.cudnn.deterministic to True ensures that cuDNN uses deterministic algorithms. Operations that rely on cuDNN (like certain convolution operations) will produce the same results on the same input data and configuration.
47        torch.manual_seed(self.opt.seed)#Ensure that operations like initializing weights in neural networks or shuffling data batches produce the same results 
48        np.random.seed(self.opt.seed)  #(mingcv) seed for every module 
49        #np.random.seed:seed for the random number generator in NumPy
50        random.seed(self.opt.seed)#seed for the built-in Python random number generator
51        #gpu_ids: for multigpu system for parallelization e.g. 0,1,2
52        str_ids = self.opt.gpu_ids.split(',')
53        self.opt.gpu_ids = []
54        for str_id in str_ids:
55            id = int(str_id)
56            if id >= 0:
57                self.opt.gpu_ids.append(id)
58        #converting string of gpus ids into a list of ints
59        #(mingcv)set gpu ids
60        if len(self.opt.gpu_ids) > 0: #if atleast 1 gpu, set first gpu as cuda device 
61            torch.cuda.set_device(self.opt.gpu_ids[0])
62
63        args = vars(self.opt)
64        #vars() is a built-in function that returns the __dict__ attribute of an object if it exists. 
65        # args = vars(self.opt) converts the attributes of self.opt into  a dictionary args for iteration and accessibility
66        print('------------ Options -------------')
67        for k, v in sorted(args.items()): # prints all keys and values of dict args
68            print('%s: %s' % (str(k), str(v)))
69        print('-------------- End ----------------')
70
71        #(mingcv) save to the disk
72        self.opt.name = self.opt.name or '_'.join([self.opt.model]) 
73        # '_'.join([self.opt.model]) if self.opt.name is None
74        expr_dir = os.path.join(self.opt.checkpoints_dir, self.opt.name)
75        #experiment directory
76        util.mkdirs(expr_dir)
77        file_name = os.path.join(expr_dir, 'opt.txt')
78        with open(file_name, 'wt') as opt_file:
79            opt_file.write('------------ Options -------------\n') # write key and values of options in opt.txt
80            for k, v in sorted(args.items()):
81                opt_file.write('%s: %s\n' % (str(k), str(v)))
82            opt_file.write('-------------- End ----------------\n')
83
84        if self.opt.debug: #debugging mode with small dataset no multithreading,no flipping
85            self.opt.display_freq = 20 #display: visual outputs related to training like plots/images
86            self.opt.print_freq = 20 #print: training progress information printed in console  
87             #after processing 20 batches of images, the training results will be printed/displayed.
88            self.opt.nEpochs = 40
89            self.opt.max_dataset_size = 100
90            self.opt.no_log = False 
91            # controls whether logging of training progress or other information is enabled (False means logging is enabled).
92            self.opt.nThreads = 0 #operations are executed in a single thread
93            self.opt.decay_iter = 0 # iteration count after which learning rate decay might occur.  0 means no decay
94            self.opt.serial_batches = True #whether batches of data are processed in serial (one after another) or in parallel during training. serial:each sample is seen exactly once per epoch in the specified order.
95            self.opt.no_flip = True # disables flipping of images during data augmentation.
96
97        return self.opt
98