Backward step.
(optimizer, model, lm_loss, args, timers)
| 268 | |
| 269 | |
| 270 | def 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 | |
| 310 | def see_memory_usage(message, force=False): |
no test coverage detected