| 522 | |
| 523 | |
| 524 | class MixtureDataset(IterableDataset): |
| 525 | def __init__(self, data_iters, weights, modality_info): |
| 526 | self.orig_data_iters = data_iters |
| 527 | self.data_iters = [iter(data_iter) for data_iter in data_iters] # Create initial iterators |
| 528 | self.sampling_probs = np.array(weights) / sum(weights) |
| 529 | self.modality_info = modality_info |
| 530 | |
| 531 | def reset_iterator(self, idx): |
| 532 | """ Reset the iterator when exhausted. """ |
| 533 | self.data_iters[idx] = iter(self.orig_data_iters[idx]) |
| 534 | |
| 535 | def __iter__(self): |
| 536 | while True: |
| 537 | dataset_idx = np.random.choice(len(self.sampling_probs), p=self.sampling_probs) |
| 538 | try: |
| 539 | data = next(self.data_iters[dataset_idx]) |
| 540 | except StopIteration: # If the iterator is exhausted |
| 541 | self.reset_iterator(dataset_idx) # Reset it |
| 542 | data = next(self.data_iters[dataset_idx]) |
| 543 | |
| 544 | mod_dict = make_empty_mod_dict(self.modality_info) |
| 545 | mod_dict.update(data) |
| 546 | yield mod_dict |
| 547 | |
| 548 | |
| 549 | def build_mixture_dataloader(data_iters, weights, modality_info, batch_size, num_workers, epoch_size, num_gpus): |
no outgoing calls
no test coverage detected