(self, data_parallel_rank, data_parallel_world_size, per_device_batch_size: dict, gradient_accumulation_steps, per_device_batch_size_image: dict)
| 949 | self.directory_datasets.append(directory_dataset) |
| 950 | |
| 951 | def post_init(self, data_parallel_rank, data_parallel_world_size, per_device_batch_size: dict, gradient_accumulation_steps, per_device_batch_size_image: dict): |
| 952 | self.data_parallel_rank = data_parallel_rank |
| 953 | self.data_parallel_world_size = data_parallel_world_size |
| 954 | global_batch_size = {size: bs * gradient_accumulation_steps * self.data_parallel_world_size for size, bs in per_device_batch_size.items()} |
| 955 | global_batch_size_image = {size: bs * gradient_accumulation_steps * self.data_parallel_world_size for size, bs in per_device_batch_size_image.items()} |
| 956 | |
| 957 | # group same size_bucket together |
| 958 | datasets_by_size_bucket = defaultdict(list) |
| 959 | for directory_dataset in self.directory_datasets: |
| 960 | for size_bucket_dataset in directory_dataset.get_size_bucket_datasets(): |
| 961 | datasets_by_size_bucket[size_bucket_dataset.size_bucket].append(size_bucket_dataset) |
| 962 | self.buckets = [] |
| 963 | for datasets in datasets_by_size_bucket.values(): |
| 964 | self.buckets.append(ConcatenatedBatchedDataset(datasets)) |
| 965 | |
| 966 | for bucket in self.buckets: |
| 967 | bucket.post_init(global_batch_size, global_batch_size_image, data_parallel_rank, data_parallel_world_size) |
| 968 | |
| 969 | iteration_order = [] |
| 970 | for i, bucket in enumerate(self.buckets): |
| 971 | iteration_order.extend([i]*(len(bucket))) |
| 972 | shuffle_with_seed(iteration_order, 0) |
| 973 | cumulative_sums = [0] * len(self.buckets) |
| 974 | for k, dataset_idx in enumerate(iteration_order): |
| 975 | iteration_order[k] = (dataset_idx, cumulative_sums[dataset_idx]) |
| 976 | cumulative_sums[dataset_idx] += 1 |
| 977 | self.iteration_order = iteration_order |
| 978 | if DEBUG: |
| 979 | print(f'Dataset iteration_order: {self.iteration_order}') |
| 980 | |
| 981 | self.post_init_called = True |
| 982 | |
| 983 | if subsample_ratio := self.dataset_config.get('subsample_ratio', None): |
| 984 | new_len = int(len(self) * subsample_ratio) |
| 985 | self.iteration_order = self.iteration_order[:new_len] |
| 986 | |
| 987 | def set_eval_quantile(self, quantile): |
| 988 | self.eval_quantile = quantile |
no test coverage detected