| 76 | |
| 77 | |
| 78 | class FalconIterator: |
| 79 | def __init__(self, filenames, seed, shuffle, tokenizer, rank_id, worker_id, state_dict): |
| 80 | self._seed = seed |
| 81 | self._shuffle = shuffle |
| 82 | self._rng = np.random.default_rng(seed) if shuffle else None |
| 83 | |
| 84 | self.print_head = f"[WORKER] R{rank_id:2d}W{worker_id:2d}: " |
| 85 | self.worker_id = worker_id |
| 86 | self.rank_id = rank_id |
| 87 | |
| 88 | self._filenames = filenames |
| 89 | self._file_idx = -1 |
| 90 | |
| 91 | self._curr_idx = 0 # current index of data item within current contents |
| 92 | |
| 93 | self.tokenizer = tokenizer |
| 94 | |
| 95 | self._curr_contents = None |
| 96 | self._pre_cache = None |
| 97 | self._pre_cache_thread = None |
| 98 | |
| 99 | if len(state_dict) != 0: |
| 100 | self._file_idx = state_dict[self.worker_id]['_file_idx'] - 1 |
| 101 | self._pre_cache_thread = Thread(target=self._preload_cache) |
| 102 | self._pre_cache_thread.start() |
| 103 | self._load_new_file() |
| 104 | assert self._file_idx == state_dict[self.worker_id]['_file_idx'] |
| 105 | self._curr_idx = state_dict[self.worker_id]['_curr_idx'] + 1 |
| 106 | else: |
| 107 | self._pre_cache_thread = Thread(target=self._preload_cache) |
| 108 | self._pre_cache_thread.start() |
| 109 | self._load_new_file() |
| 110 | |
| 111 | def __iter__(self): |
| 112 | return self |
| 113 | |
| 114 | def _preload_cache(self): |
| 115 | if self._file_idx + 1 >= len(self._filenames): |
| 116 | self._pre_cache = None |
| 117 | else: |
| 118 | print(f"{self.print_head} current {self._file_idx}, async load {self._file_idx + 1} {self._filenames[self._file_idx + 1]}") |
| 119 | |
| 120 | with open(self._filenames[self._file_idx + 1], 'rb') as f: |
| 121 | ann = pickle.load(f) |
| 122 | self._pre_cache = ann |
| 123 | |
| 124 | return |
| 125 | |
| 126 | def _load_new_file(self, pre_load=True): |
| 127 | self._pre_cache_thread.join() |
| 128 | |
| 129 | if self._file_idx + 1 >= len(self._filenames): |
| 130 | assert self._pre_cache is None |
| 131 | raise StopIteration |
| 132 | else: |
| 133 | assert self._pre_cache is not None |
| 134 | self._curr_contents = self._pre_cache |
| 135 | |