MCPcopy Create free account
hub / github.com/tdrussell/diffusion-pipe / post_init

Method post_init

utils/dataset.py:951–985  ·  view source on GitHub ↗
(self, data_parallel_rank, data_parallel_world_size, per_device_batch_size: dict, gradient_accumulation_steps, per_device_batch_size_image: dict)

Source from the content-addressed store, hash-verified

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

Callers 2

train.pyFile · 0.45
dataset.pyFile · 0.45

Calls 4

shuffle_with_seedFunction · 0.85
getMethod · 0.80

Tested by

no test coverage detected