MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / IndexedDataset

Class IndexedDataset

codegeex/megatron/data/indexed_dataset.py:148–230  ·  view source on GitHub ↗

Loader for IndexedDataset

Source from the content-addressed store, hash-verified

146
147
148class 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)

Callers 2

make_datasetFunction · 0.85
merge_file_Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected