| 354 | |
| 355 | |
| 356 | class MMapIndexedDataset(torch.utils.data.Dataset): |
| 357 | class Index(object): |
| 358 | _HDR_MAGIC = b"MMIDIDX\x00\x00" |
| 359 | |
| 360 | @classmethod |
| 361 | def writer(cls, path, dtype): |
| 362 | class _Writer(object): |
| 363 | def __enter__(self): |
| 364 | self._file = open(path, "wb") |
| 365 | |
| 366 | self._file.write(cls._HDR_MAGIC) |
| 367 | self._file.write(struct.pack("<Q", 1)) |
| 368 | self._file.write(struct.pack("<B", code(dtype))) |
| 369 | |
| 370 | return self |
| 371 | |
| 372 | @staticmethod |
| 373 | def _get_pointers(sizes): |
| 374 | dtype_size = dtype().itemsize |
| 375 | address = 0 |
| 376 | pointers = [] |
| 377 | |
| 378 | for size in sizes: |
| 379 | pointers.append(address) |
| 380 | address += size * dtype_size |
| 381 | |
| 382 | return pointers |
| 383 | |
| 384 | def write(self, sizes, doc_idx): |
| 385 | pointers = self._get_pointers(sizes) |
| 386 | |
| 387 | self._file.write(struct.pack("<Q", len(sizes))) |
| 388 | self._file.write(struct.pack("<Q", len(doc_idx))) |
| 389 | |
| 390 | sizes = np.array(sizes, dtype=np.int32) |
| 391 | self._file.write(sizes.tobytes(order="C")) |
| 392 | del sizes |
| 393 | |
| 394 | pointers = np.array(pointers, dtype=np.int64) |
| 395 | self._file.write(pointers.tobytes(order="C")) |
| 396 | del pointers |
| 397 | |
| 398 | doc_idx = np.array(doc_idx, dtype=np.int64) |
| 399 | self._file.write(doc_idx.tobytes(order="C")) |
| 400 | |
| 401 | def __exit__(self, exc_type, exc_val, exc_tb): |
| 402 | self._file.close() |
| 403 | |
| 404 | return _Writer() |
| 405 | |
| 406 | def __init__(self, path, skip_warmup=False): |
| 407 | with open(path, "rb") as stream: |
| 408 | magic_test = stream.read(9) |
| 409 | assert self._HDR_MAGIC == magic_test, ( |
| 410 | "Index file doesn't match expected format. " |
| 411 | "Make sure that --dataset-impl is configured properly." |
| 412 | ) |
| 413 | version = struct.unpack("<Q", stream.read(8)) |