(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True)
| 251 | self._scaler = torch.cuda.amp.GradScaler() |
| 252 | |
| 253 | def __call__(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True): |
| 254 | self._scaler.scale(loss).backward(create_graph=create_graph) |
| 255 | if update_grad: |
| 256 | if clip_grad is not None: |
| 257 | assert parameters is not None |
| 258 | self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place |
| 259 | norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad) |
| 260 | else: |
| 261 | self._scaler.unscale_(optimizer) |
| 262 | norm = get_grad_norm_(parameters) |
| 263 | self._scaler.step(optimizer) |
| 264 | self._scaler.update() |
| 265 | else: |
| 266 | norm = None |
| 267 | return norm |
| 268 | |
| 269 | def state_dict(self): |
| 270 | return self._scaler.state_dict() |
nothing calls this directly
no test coverage detected