| 23 | |
| 24 | |
| 25 | class BlendableDataset(torch.utils.data.Dataset): |
| 26 | def __init__(self, datasets, weights): |
| 27 | |
| 28 | self.datasets = datasets |
| 29 | num_datasets = len(datasets) |
| 30 | assert num_datasets == len(weights) |
| 31 | |
| 32 | self.size = 0 |
| 33 | for dataset in self.datasets: |
| 34 | self.size += len(dataset) |
| 35 | |
| 36 | # Normalize weights. |
| 37 | weights = np.array(weights, dtype=np.float64) |
| 38 | sum_weights = np.sum(weights) |
| 39 | assert sum_weights > 0.0 |
| 40 | weights /= sum_weights |
| 41 | |
| 42 | # Build indecies. |
| 43 | start_time = time.time() |
| 44 | assert num_datasets < 255 |
| 45 | self.dataset_index = np.zeros(self.size, dtype=np.uint8) |
| 46 | self.dataset_sample_index = np.zeros(self.size, dtype=np.int64) |
| 47 | |
| 48 | from megatron.data import helpers |
| 49 | |
| 50 | helpers.build_blending_indices( |
| 51 | self.dataset_index, |
| 52 | self.dataset_sample_index, |
| 53 | weights, |
| 54 | num_datasets, |
| 55 | self.size, |
| 56 | torch.distributed.get_rank() == 0, |
| 57 | ) |
| 58 | print_rank_0( |
| 59 | "> elapsed time for building blendable dataset indices: " |
| 60 | "{:.2f} (sec)".format(time.time() - start_time) |
| 61 | ) |
| 62 | |
| 63 | def __len__(self): |
| 64 | return self.size |
| 65 | |
| 66 | def __getitem__(self, idx): |
| 67 | dataset_idx = self.dataset_index[idx] |
| 68 | sample_idx = self.dataset_sample_index[idx] |
| 69 | return self.datasets[dataset_idx][sample_idx] |
no outgoing calls
no test coverage detected