Team Ai
Apppublic

ReflectionEraser/ReflectionEraserApp

sourceHugging Faceotherupdated 2y agoView on Hugging Face
0likes
torchdata.py74 linesDownload Raw Back to data
1import bisect #maintains a list in sorted order without having sort after each insertion using bisection algorithm
2import warnings #issue warning message
3
4
5class Dataset(object): # all python classes inherit from object
6    """An abstract class representing a Dataset.
7
8    All other datasets should subclass it. All subclasses should override
9    ``__len__``, that provides the size of the dataset, and ``__getitem__``,
10    supporting integer indexing in range from 0 to len(self) exclusive.
11    """
12    #traditional overloading is not in python
13
14    def __getitem__(self, index): #__getitem__ for indexing with [ ] (overrriding)
15        raise NotImplementedError
16
17    def __len__(self): # for overriding len()
18        raise NotImplementedError #when you call len(obj), Python internally calls obj.__len__().
19
20    def __add__(self, other):  #to override an binary + operator
21        return ConcatDataset([self, other])
22
23    def reset(self): #currently is a placeholder and does nothing
24        return 
25
26
27class ConcatDataset(Dataset): #inheriting from Dataset
28    """
29    Dataset to concatenate multiple datasets.
30    Purpose: useful to assemble different existing datasets, possibly
31    large-scale datasets as the concatenation operation is done in an
32    on-the-fly manner.
33
34    Arguments:
35        datasets (sequence): List of datasets to be concatenated
36    """
37
38    @staticmethod
39    def cumsum(sequence): #sequence: an ordered collection of items, where each item holds a relative position.Must be indexable and can be iterated. All sequences are iterables but reverse is not true e.g. sets   not indexable
40        #sequence   are strings, lists, tuples, byte sequences, byte arrays and range objects  
41        r, s = [], 0  #r:list to hold cumulative sum of iterated datasets in sequence
42                      #s: cumulative sum of all datasets in sequence
43        for e in sequence: #e: current dataset in sequence
44            l = len(e)      # length of current dataset in sequence
45            r.append(l + s)
46            s += l
47        return r
48
49    def __init__(self, datasets):
50        super(ConcatDataset, self).__init__()
51        assert len(datasets) > 0, 'datasets should not be an empty iterable'
52                                  #halts the program and provides a diagnostic message.
53                                  #asserts are usually stripped out or disabled in production code, as they are meant for development and debugging
54        self.datasets = list(datasets) #store list of datasets
55        self.cumulative_sizes = self.cumsum(self.datasets) #store result of static method cumsum, list of cumulative length of previous datasets
56
57    def __len__(self):
58        return self.cumulative_sizes[-1] #length of concatenated datasets after cumsum
59
60    def __getitem__(self, idx): #dynamically selecting the correct dataset and item within that dataset based on the given index in the self.datasets
61        dataset_idx = bisect.bisect_right(self.cumulative_sizes, idx) #bisect.bisect_right returns the index of the first element in the list that is greater than x (upper bound)
62        # line determines which dataset in self.datasets the item at idx belongs to.
63        if dataset_idx == 0: #idx belongs to first dataset
64            sample_idx = idx #so idx and cumulative idx same
65        else:#finds idx within the dataset
66            sample_idx = idx - self.cumulative_sizes[dataset_idx - 1]
67        return self.datasets[dataset_idx][sample_idx]
68
69    @property #property of the class. Can be accessed without ()
70    def cummulative_sizes(self): # basically a getter for encapsulation
71        warnings.warn("cummulative_sizes attribute is renamed to " #basically if user uses wrong function
72                      "cumulative_sizes", DeprecationWarning, stacklevel=2) #If the stacklevel parameter is not set, the warning appears to originate from within the cummulative_sizes method itself.
73         #By setting stacklevel=2, the warning will appear to originate from the caller of the cummulative_sizes method.
74        return self.cumulative_sizes