MCPcopy Create free account
hub / github.com/csuhan/OneLLM / __call__

Method __call__

util/misc.py:316–334  ·  view source on GitHub ↗
(self, loss, optimizer, model, clip_grad=None, parameters=None, create_graph=False, update_grad=True)

Source from the content-addressed store, hash-verified

314 self._scaler = ShardedGradScaler(enabled=args.precision in ["fp16"])
315
316 def __call__(self, loss, optimizer, model, clip_grad=None, parameters=None, create_graph=False, update_grad=True):
317 if update_grad:
318 self._scaler.scale(loss).backward(create_graph=create_graph)
319 if clip_grad is not None:
320 assert parameters is not None
321 self._scaler.unscale_(optimizer) # unscale the gradients of optimizer's assigned params in-place
322 # norm = torch.nn.utils.clip_grad_norm_(parameters, clip_grad)
323 norm = model.clip_grad_norm_(clip_grad)
324 else:
325 raise NotImplementedError("please set clip_grad to a very large value if you do not want to clip.")
326 self._scaler.unscale_(optimizer)
327 norm = get_grad_norm_(parameters)
328 self._scaler.step(optimizer)
329 self._scaler.update()
330 else:
331 with model.no_sync():
332 self._scaler.scale(loss).backward(create_graph=create_graph)
333 norm = None
334 return norm
335
336 def state_dict(self):
337 return self._scaler.state_dict()

Callers

nothing calls this directly

Calls 3

get_grad_norm_Function · 0.85
backwardMethod · 0.45
updateMethod · 0.45

Tested by

no test coverage detected