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

Function update_train_iters

codegeex/megatron/training.py:221–247  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

219
220
221def update_train_iters(args):
222
223 # For iteration-based training, we don't need to do anything
224 if args.train_iters:
225 return
226
227 # Constant batch size with sample-based training.
228 if args.rampup_batch_size is None:
229 args.train_iters = args.train_samples // args.global_batch_size
230
231 else:
232 # Sample based training with rampup batch size.
233 iterations = 0
234 consumed_samples = 0
235 # Rampup phase.
236 while consumed_samples <= int(args.rampup_batch_size[2]):
237 update_num_microbatches(consumed_samples, consistency_check=False)
238 consumed_samples += get_current_global_batch_size()
239 iterations += 1
240 # Reset
241 update_num_microbatches(0, consistency_check=False)
242 # Constant phase
243 # Note that we throw away any partial last batch.
244 iterations += (args.train_samples - consumed_samples) // args.global_batch_size
245 args.train_iters = iterations
246
247 print_rank_0("setting training iterations to {}".format(args.train_iters))
248
249
250def get_model(model_provider_func):

Callers 1

Calls 3

update_num_microbatchesFunction · 0.90
print_rank_0Function · 0.90

Tested by

no test coverage detected