MCPcopy Create free account
hub / github.com/MCG-NJU/VideoMAE / NativeScalerWithGradNormCount

Class NativeScalerWithGradNormCount

utils.py:341–367  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

339
340
341class 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
370def get_grad_norm_(parameters, norm_type: float = 2.0) -> torch.Tensor:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected