| 1748 | return |
| 1749 | |
| 1750 | class DecayAdagradOptimizer(Optimizer): |
| 1751 | def __init__( |
| 1752 | self, |
| 1753 | alpha=0.01, |
| 1754 | beta1=0.0, |
| 1755 | beta2=0.999, |
| 1756 | epsilon=0.1, |
| 1757 | weight_decay=0.0, |
| 1758 | ema_options=None, |
| 1759 | bias_correction_first=True, |
| 1760 | policy="fixed", |
| 1761 | engine="", |
| 1762 | **kwargs |
| 1763 | ): |
| 1764 | super().__init__() |
| 1765 | self.alpha = alpha |
| 1766 | self.beta1 = beta1 |
| 1767 | self.beta2 = beta2 |
| 1768 | self.epsilon = epsilon |
| 1769 | self.weight_decay = weight_decay |
| 1770 | self.bias_correction_first = bias_correction_first |
| 1771 | self.policy = policy |
| 1772 | self.engine = engine |
| 1773 | self.init_kwargs = kwargs |
| 1774 | self._process_ema_options(ema_options) |
| 1775 | |
| 1776 | def set_mapping_for_param2ema_teacher_param(self, param_mapping: Dict[str, Any]) -> None: |
| 1777 | self.param2ema_teacher_param = param_mapping |
| 1778 | |
| 1779 | def _process_ema_options(self, ema_options): |
| 1780 | self.ema_enabled = True if ema_options and "ema_alpha" in ema_options else False |
| 1781 | self.ema_teacher_enabled = True if ema_options and "ema_teacher_alpha" in ema_options else False |
| 1782 | self.param2ema_teacher_param = {} |
| 1783 | if self.ema_enabled or self.ema_teacher_enabled: |
| 1784 | self.ema_start = ema_options.get("ema_start", None) |
| 1785 | self.ema_end = ema_options.get("ema_end", None) |
| 1786 | self.ema_step = ema_options.get("ema_step", None) |
| 1787 | self.ema_alpha = ema_options.get("ema_alpha", None) |
| 1788 | self.ema_teacher_alpha = ema_options.get("ema_alpha", None) |
| 1789 | self.ema_teacher_module_name = ema_options.get( |
| 1790 | "ema_teacher_module_name", "ema_teacher_arch" |
| 1791 | ) |
| 1792 | |
| 1793 | def _run(self, net, param_init_net, param_info): |
| 1794 | param = param_info.blob |
| 1795 | grad = param_info.grad |
| 1796 | |
| 1797 | if self.alpha <= 0: |
| 1798 | return |
| 1799 | |
| 1800 | lr, iteration = self.build_lr( |
| 1801 | net, |
| 1802 | param_init_net, |
| 1803 | base_learning_rate=self.alpha, |
| 1804 | policy=self.policy, |
| 1805 | **(self.init_kwargs) |
| 1806 | ) |
| 1807 |
no outgoing calls
no test coverage detected
searching dependent graphs…