| 1581 | |
| 1582 | |
| 1583 | class AdamOptimizer(Optimizer): |
| 1584 | def __init__( |
| 1585 | self, |
| 1586 | alpha=0.001, |
| 1587 | beta1=0.9, |
| 1588 | beta2=0.999, |
| 1589 | epsilon=1e-8, |
| 1590 | policy="fixed", |
| 1591 | use_lr_adaption=False, |
| 1592 | lr_alpha=0.01, |
| 1593 | normalized_lr_adaption=True, |
| 1594 | sparse_dedup_aggregator=None, |
| 1595 | rowWise=False, |
| 1596 | engine="", |
| 1597 | enableRAdam=False, |
| 1598 | use_smart_decay=False, # See https://fburl.com/2jdiwrhy for context. |
| 1599 | **kwargs |
| 1600 | ): |
| 1601 | super().__init__() |
| 1602 | self.alpha = alpha |
| 1603 | self.beta1 = beta1 |
| 1604 | self.beta2 = beta2 |
| 1605 | self.epsilon = epsilon |
| 1606 | self.policy = policy |
| 1607 | self.use_lr_adaption = use_lr_adaption |
| 1608 | self.lr_alpha = lr_alpha |
| 1609 | self.normalized_lr_adaption = normalized_lr_adaption |
| 1610 | self.sparse_dedup_aggregator = sparse_dedup_aggregator |
| 1611 | self.rowWise = rowWise |
| 1612 | self.engine = engine |
| 1613 | self.enableRAdam = enableRAdam |
| 1614 | if use_smart_decay: |
| 1615 | if rowWise: |
| 1616 | raise NotImplementedError(('Smart decay is not implemented for rowWise Adam. ' |
| 1617 | 'Set rowWise or use_smart_decay to False.')) |
| 1618 | if enableRAdam: |
| 1619 | raise NotImplementedError(('Smart decay is not implemented for RAdam. ' |
| 1620 | 'Set enableRAdam or use_smart_decay to False.')) |
| 1621 | if use_lr_adaption: |
| 1622 | raise NotImplementedError(('Smart decay is not implemented with lr_adaption. ' |
| 1623 | 'Set use_lr_adaption or use_smart_decay to False.')) |
| 1624 | |
| 1625 | self.use_smart_decay = use_smart_decay |
| 1626 | self.init_kwargs = kwargs |
| 1627 | |
| 1628 | def _run(self, net, param_init_net, param_info): |
| 1629 | param = param_info.blob |
| 1630 | grad = param_info.grad |
| 1631 | |
| 1632 | if self.alpha <= 0: |
| 1633 | return |
| 1634 | |
| 1635 | lr, iteration = self.build_lr( |
| 1636 | net, |
| 1637 | param_init_net, |
| 1638 | base_learning_rate=self.alpha, |
| 1639 | policy=self.policy, |
| 1640 | **(self.init_kwargs) |
no outgoing calls
searching dependent graphs…