MCPcopy Create free account
hub / github.com/LTH14/rcg / NativeScalerWithGradNormCount

Class NativeScalerWithGradNormCount

util/misc.py:254–280  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

252
253
254class NativeScalerWithGradNormCount:
255 state_dict_key = "amp_scaler"
256
257 def __init__(self):
258 self._scaler = torch.cuda.amp.GradScaler()
259
260 def __call__(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True):
261 self._scaler.scale(loss).backward(create_graph=create_graph)
262 if update_grad:
263 if clip_grad is not None:
264 assert parameters is not None
265 self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place
266 norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad)
267 else:
268 self._scaler.unscale_(optimizer)
269 norm = get_grad_norm_(parameters)
270 self._scaler.step(optimizer)
271 self._scaler.update()
272 else:
273 norm = None
274 return norm
275
276 def state_dict(self):
277 return self._scaler.state_dict()
278
279 def load_state_dict(self, state_dict):
280 self._scaler.load_state_dict(state_dict)
281
282
283def 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