Almost copy-paste from https://github.com/facebookresearch/barlowtwins/blob/main/main.py
| 531 | |
| 532 | |
| 533 | class LARS(torch.optim.Optimizer): |
| 534 | """ |
| 535 | Almost copy-paste from https://github.com/facebookresearch/barlowtwins/blob/main/main.py |
| 536 | """ |
| 537 | def __init__(self, params, lr=0, weight_decay=0, momentum=0.9, eta=0.001, |
| 538 | weight_decay_filter=None, lars_adaptation_filter=None): |
| 539 | defaults = dict(lr=lr, weight_decay=weight_decay, momentum=momentum, |
| 540 | eta=eta, weight_decay_filter=weight_decay_filter, |
| 541 | lars_adaptation_filter=lars_adaptation_filter) |
| 542 | super().__init__(params, defaults) |
| 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 | |
| 574 | class MultiCropWrapper(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected