MCPcopy Create free account
hub / github.com/OpenGVLab/EfficientQAT / __call__

Method __call__

utils.py:31–45  ·  view source on GitHub ↗
(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True,retain_graph=False)

Source from the content-addressed store, hash-verified

29 self._scaler = torch.cuda.amp.GradScaler()
30
31 def __call__(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True,retain_graph=False):
32 self._scaler.scale(loss).backward(create_graph=create_graph, retain_graph=retain_graph)
33 if update_grad:
34 if clip_grad is not None:
35 assert parameters is not None
36 self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place
37 norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad)
38 else:
39 self._scaler.unscale_(optimizer)
40 norm = ampscaler_get_grad_norm(parameters)
41 self._scaler.step(optimizer)
42 self._scaler.update()
43 else:
44 norm = None
45 return norm
46
47 def state_dict(self):
48 return self._scaler.state_dict()

Callers

nothing calls this directly

Calls 2

ampscaler_get_grad_normFunction · 0.85
backwardMethod · 0.80

Tested by

no test coverage detected