(self, loss, optimizer, model, clip_grad=None, parameters=None, create_graph=False, update_grad=True)
| 314 | self._scaler = ShardedGradScaler(enabled=args.precision in ["fp16"]) |
| 315 | |
| 316 | def __call__(self, loss, optimizer, model, clip_grad=None, parameters=None, create_graph=False, update_grad=True): |
| 317 | if update_grad: |
| 318 | self._scaler.scale(loss).backward(create_graph=create_graph) |
| 319 | if clip_grad is not None: |
| 320 | assert parameters is not None |
| 321 | self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place |
| 322 | # norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad) |
| 323 | norm = model.clip_grad_norm_(clip_grad) |
| 324 | else: |
| 325 | raise NotImplementedError("please set clip_grad to a very large value if you do not want to clip.") |
| 326 | self._scaler.unscale_(optimizer) |
| 327 | norm = get_grad_norm_(parameters) |
| 328 | self._scaler.step(optimizer) |
| 329 | self._scaler.update() |
| 330 | else: |
| 331 | with model.no_sync(): |
| 332 | self._scaler.scale(loss).backward(create_graph=create_graph) |
| 333 | norm = None |
| 334 | return norm |
| 335 | |
| 336 | def state_dict(self): |
| 337 | return self._scaler.state_dict() |
nothing calls this directly
no test coverage detected