| 38 | |
| 39 | |
| 40 | class PackedDataset(IterableDataset): |
| 41 | def __init__( |
| 42 | self, |
| 43 | filenames, |
| 44 | n_chunks, |
| 45 | block_size, |
| 46 | seed=12345, |
| 47 | shuffle=True, |
| 48 | wrap=False, |
| 49 | num_processes=1, |
| 50 | process_rank=0, |
| 51 | ): |
| 52 | self._filenames = filenames |
| 53 | self._n_chunks = n_chunks |
| 54 | self._block_size = block_size |
| 55 | self._seed = seed |
| 56 | self._shuffle = shuffle |
| 57 | self._wrap = wrap |
| 58 | self._num_processes = num_processes |
| 59 | self._process_rank = process_rank |
| 60 | |
| 61 | def __iter__(self): |
| 62 | worker_info = get_worker_info() |
| 63 | num_workers = worker_info.num_workers if worker_info is not None else 1 |
| 64 | worker_id = worker_info.id if worker_info is not None else 0 |
| 65 | num_shards = num_workers * self._num_processes |
| 66 | shard_id = self._process_rank * num_workers + worker_id |
| 67 | |
| 68 | max_num_files = len(self._filenames) // num_shards * num_shards |
| 69 | filenames = self._filenames[shard_id:max_num_files:num_shards] |
| 70 | |
| 71 | return PackedDatasetIterator( |
| 72 | filenames=filenames, |
| 73 | n_chunks=self._n_chunks, |
| 74 | block_size=self._block_size, |
| 75 | seed=self._seed, |
| 76 | shuffle=self._shuffle, |
| 77 | wrap=self._wrap, |
| 78 | ) |
| 79 | |
| 80 | |
| 81 | class PackedDatasetBuilder(object): |