MCPcopy Create free account
hub / github.com/alinlab/SelfPatch / step

Method step

utils.py:545–571  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

543
544 @torch.no_grad()
545 def step(self):
546 for g in self.param_groups:
547 for p in g['params']:
548 dp = p.grad
549
550 if dp is None:
551 continue
552
553 if p.ndim != 1:
554 dp = dp.add(p, alpha=g['weight_decay'])
555
556 if p.ndim != 1:
557 param_norm = torch.norm(p)
558 update_norm = torch.norm(dp)
559 one = torch.ones_like(param_norm)
560 q = torch.where(param_norm > 0.,
561 torch.where(update_norm > 0,
562 (g['eta'] * param_norm / update_norm), one), one)
563 dp = dp.mul(q)
564
565 param_state = self.state[p]
566 if 'mu' not in param_state:
567 param_state['mu'] = torch.zeros_like(p)
568 mu = param_state['mu']
569 mu.mul_(g['momentum']).add_(dp)
570
571 p.add_(mu, alpha=-g['lr'])
572
573
574class MultiCropWrapper(nn.Module):

Callers 1

train_one_epochFunction · 0.80

Calls

no outgoing calls

Tested by

no test coverage detected