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

Function backward_step

train_utils.py:270–307  ·  view source on GitHub ↗

Backward step.

(optimizer, model, lm_loss, args, timers)

Source from the content-addressed store, hash-verified

268
269
270def backward_step(optimizer, model, lm_loss, args, timers):
271 """Backward step."""
272
273 # Total loss.
274 loss = lm_loss
275
276 # Backward pass.
277 if args.deepspeed:
278 model.backward(loss)
279 else:
280 # optimizer.zero_grad()
281 if args.fp16:
282 optimizer.backward(loss, update_master_grads=False)
283 else:
284 loss.backward()
285
286 if args.deepspeed or args.DDP_impl == 'torch':
287 # DeepSpeed backward propagation already addressed all reduce communication.
288 # Reset the timer to avoid breaking timer logs below.
289 timers('allreduce').reset()
290 else:
291 timers('allreduce').start()
292 model.allreduce_params(reduce_after=False, fp32_allreduce=args.fp32_allreduce)
293 timers('allreduce').stop()
294
295 # Update master gradients.
296 if not args.deepspeed:
297 if args.fp16:
298 optimizer.update_master_grads()
299
300 # Clipping gradients helps prevent the exploding gradient.
301 if args.clip_grad > 0:
302 if not args.fp16:
303 mpu.clip_grad_norm(model.parameters(), args.clip_grad)
304 else:
305 optimizer.clip_master_grads(args.clip_grad)
306
307 return lm_loss
308
309
310def see_memory_usage(message, force=False):

Callers 1

train_stepFunction · 0.85

Calls 7

startMethod · 0.80
allreduce_paramsMethod · 0.80
stopMethod · 0.80
update_master_gradsMethod · 0.80
clip_master_gradsMethod · 0.80
backwardMethod · 0.45
resetMethod · 0.45

Tested by

no test coverage detected