| 339 | |
| 340 | |
| 341 | class NativeScalerWithGradNormCount: |
| 342 | state_dict_key = "amp_scaler" |
| 343 | |
| 344 | def __init__(self): |
| 345 | self._scaler = torch.cuda.amp.GradScaler() |
| 346 | |
| 347 | def __call__(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True): |
| 348 | self._scaler.scale(loss).backward(create_graph=create_graph) |
| 349 | if update_grad: |
| 350 | if clip_grad is not None: |
| 351 | assert parameters is not None |
| 352 | self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place |
| 353 | norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad) |
| 354 | else: |
| 355 | self._scaler.unscale_(optimizer) |
| 356 | norm = get_grad_norm_(parameters) |
| 357 | self._scaler.step(optimizer) |
| 358 | self._scaler.update() |
| 359 | else: |
| 360 | norm = None |
| 361 | return norm |
| 362 | |
| 363 | def state_dict(self): |
| 364 | return self._scaler.state_dict() |
| 365 | |
| 366 | def load_state_dict(self, state_dict): |
| 367 | self._scaler.load_state_dict(state_dict) |
| 368 | |
| 369 | |
| 370 | def get_grad_norm_(parameters, norm_type: float = 2.0) -> torch.Tensor: |
nothing calls this directly
no outgoing calls
no test coverage detected