Loader for IndexedDataset
| 146 | |
| 147 | |
| 148 | class IndexedDataset(torch.utils.data.Dataset): |
| 149 | """Loader for IndexedDataset""" |
| 150 | |
| 151 | _HDR_MAGIC = b"TNTIDX\x00\x00" |
| 152 | |
| 153 | def __init__(self, path): |
| 154 | super().__init__() |
| 155 | self.path = path |
| 156 | self.data_file = None |
| 157 | self.read_index(path) |
| 158 | |
| 159 | def read_index(self, path): |
| 160 | with open(index_file_path(path), "rb") as f: |
| 161 | magic = f.read(8) |
| 162 | assert magic == self._HDR_MAGIC, ( |
| 163 | "Index file doesn't match expected format. " |
| 164 | "Make sure that --dataset-impl is configured properly." |
| 165 | ) |
| 166 | version = f.read(8) |
| 167 | assert struct.unpack("<Q", version) == (1,) |
| 168 | code, self.element_size = struct.unpack("<QQ", f.read(16)) |
| 169 | self.dtype = dtypes[code] |
| 170 | self._len, self.s = struct.unpack("<QQ", f.read(16)) |
| 171 | self.doc_count = struct.unpack("<Q", f.read(8)) |
| 172 | self.dim_offsets = read_longs(f, self._len + 1) |
| 173 | self.data_offsets = read_longs(f, self._len + 1) |
| 174 | self.sizes = read_longs(f, self.s) |
| 175 | self.doc_idx = read_longs(f, self.doc_count) |
| 176 | |
| 177 | def read_data(self, path): |
| 178 | self.data_file = open(data_file_path(path), "rb", buffering=0) |
| 179 | |
| 180 | def check_index(self, i): |
| 181 | if i < 0 or i >= self._len: |
| 182 | raise IndexError("index out of range") |
| 183 | |
| 184 | def __del__(self): |
| 185 | if self.data_file: |
| 186 | self.data_file.close() |
| 187 | |
| 188 | # @lru_cache(maxsize=8) |
| 189 | def __getitem__(self, idx): |
| 190 | if not self.data_file: |
| 191 | self.read_data(self.path) |
| 192 | if isinstance(idx, int): |
| 193 | i = idx |
| 194 | self.check_index(i) |
| 195 | tensor_size = self.sizes[self.dim_offsets[i] : self.dim_offsets[i + 1]] |
| 196 | a = np.empty(tensor_size, dtype=self.dtype) |
| 197 | self.data_file.seek(self.data_offsets[i] * self.element_size) |
| 198 | self.data_file.readinto(a) |
| 199 | return a |
| 200 | elif isinstance(idx, slice): |
| 201 | start, stop, step = idx.indices(len(self)) |
| 202 | if step != 1: |
| 203 | raise ValueError("Slices into indexed_dataset must be contiguous") |
| 204 | sizes = self.sizes[self.dim_offsets[start] : self.dim_offsets[stop]] |
| 205 | size = sum(sizes) |
no outgoing calls
no test coverage detected