MCPcopy Create free account
hub / github.com/THUDM/GLM / train_step

Function train_step

train_utils.py:324–387  ·  view source on GitHub ↗

Single training step.

(data_iterator, model, optimizer, lr_scheduler, args, timers, forward_step_func, mems=None,
               single_step=False)

Source from the content-addressed store, hash-verified

322
323
324def train_step(data_iterator, model, optimizer, lr_scheduler, args, timers, forward_step_func, mems=None,
325 single_step=False):
326 """Single training step."""
327 lm_loss_total, count = 0.0, 0
328 mems = [] if mems is None else mems
329 if not args.deepspeed:
330 optimizer.zero_grad()
331 while True:
332 skipped_iter, complete = 0, False
333 # Forward model for one step.
334 timers('forward').start()
335 lm_loss, mems, _ = forward_step_func(data_iterator, model, args, timers, mems)
336 timers('forward').stop()
337 # print_rank_0("Forward step")
338 if not args.deepspeed:
339 lm_loss /= args.gradient_accumulation_steps
340
341 reduced_loss = lm_loss.detach().clone().view(1)
342 torch.distributed.all_reduce(reduced_loss.data, group=mpu.get_data_parallel_group())
343 reduced_loss.data = reduced_loss.data / (args.world_size / args.model_parallel_size)
344
345 if not DynamicLossScaler._has_inf_or_nan(reduced_loss):
346 lm_loss_total += reduced_loss
347 count += 1
348
349 # Calculate gradients, reduce across processes, and clip.
350 timers('backward').start()
351 backward_step(optimizer, model, lm_loss, args, timers)
352 timers('backward').stop()
353 # print_rank_0("Backward step")
354 # Update parameters.
355 timers('optimizer').start()
356 if args.deepspeed:
357 if model.is_gradient_accumulation_boundary():
358 model.step()
359 complete = True
360 if not (args.fp16 and optimizer.overflow):
361 lr_scheduler.step()
362 else:
363 skipped_iter = 1
364 else:
365 model.step()
366 else:
367 if count == args.gradient_accumulation_steps:
368 optimizer.step()
369 complete = True
370 # Update learning rate.
371 if not (args.fp16 and optimizer.overflow):
372 lr_scheduler.step()
373 else:
374 skipped_iter = 1
375 # print_rank_0("Optimizer step")
376 timers('optimizer').stop()
377 if complete:
378 break
379 else:
380 print_rank_0("Found NaN loss, skip backward")
381 del lm_loss, reduced_loss

Callers 2

_trainFunction · 0.90
trainFunction · 0.90

Calls 7

print_rank_0Function · 0.90
backward_stepFunction · 0.85
zero_gradMethod · 0.80
startMethod · 0.80
stopMethod · 0.80
_has_inf_or_nanMethod · 0.45
stepMethod · 0.45

Tested by

no test coverage detected