MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / ConstantNumMicroBatches

Class ConstantNumMicroBatches

codegeex/megatron/microbatches.py:86–100  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

84
85
86class 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
101
102
103class RampupBatchsizeNumMicroBatches(NumMicroBatchesCalculator):

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected