(self, global_batch_size, micro_batch_size, data_parallel_size)
| 85 | |
| 86 | class ConstantNumMicroBatches(NumMicroBatchesCalculator): |
| 87 | def __init__(self, global_batch_size, micro_batch_size, data_parallel_size): |
| 88 | micro_batch_times_data_parallel = micro_batch_size * data_parallel_size |
| 89 | assert global_batch_size % micro_batch_times_data_parallel == 0, ( |
| 90 | "global batch size ({}) is not divisible by micro batch size ({})" |
| 91 | " times data parallel size ({})".format( |
| 92 | global_batch_size, micro_batch_size, data_parallel_size |
| 93 | ) |
| 94 | ) |
| 95 | self.num_micro_batches = global_batch_size // micro_batch_times_data_parallel |
| 96 | assert self.num_micro_batches >= 1 |
| 97 | self.current_global_batch_size = global_batch_size |
| 98 | |
| 99 | def update(self, consumed_samples, consistency_check): |
| 100 | pass |
nothing calls this directly
no outgoing calls
no test coverage detected