MCPcopy Create free account
hub / github.com/coperception/star / __call__

Method __call__

star/utils/misc.py:257–271  ·  view source on GitHub ↗
(self, loss, optimizer, clip_grad=None, parameters=None, create_graph=False, update_grad=True)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 3

get_grad_norm_Function · 0.85
stepMethod · 0.80
updateMethod · 0.45

Tested by

no test coverage detected