(self, net, param_init_net, param_info)
| 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 | |
| 1808 | if isinstance(grad, core.GradientSlice): |
| 1809 | # hack for position weighted. |
| 1810 | param_squared_sum = param_init_net.ConstantFill([param], param + "_squared_sum", value=0.0) |
| 1811 | self._aux_params.local.append(param_squared_sum) |
| 1812 | output_blobs = [param, param_squared_sum] |
| 1813 | net.SparseAdagrad( |
| 1814 | [param, param_squared_sum, grad.indices, grad.values, lr], |
| 1815 | output_blobs, |
| 1816 | epsilon=self.epsilon, |
| 1817 | ) |
| 1818 | else: |
| 1819 | m1 = param_init_net.ConstantFill([param], param + "_first_mo1ment", value=0.0) |
| 1820 | m2 = param_init_net.ConstantFill([param], param + "_second_moment", value=0.0) |
| 1821 | self._aux_params.shared.append(iteration) |
| 1822 | self._aux_params.local.append(m1) |
| 1823 | self._aux_params.local.append(m2) |
| 1824 | output_blobs = [param, m1, m2] |
| 1825 | net.DecayAdagrad( |
| 1826 | [param, m1, m2, grad, lr, iteration], |
| 1827 | output_blobs, |
| 1828 | beta1=self.beta1, |
| 1829 | beta2=self.beta2, |
| 1830 | epsilon=self.epsilon, |
| 1831 | weight_decay=self.weight_decay, |
| 1832 | bias_correction_first=self.bias_correction_first, |
| 1833 | ) |
| 1834 | |
| 1835 | if self.ema_enabled: |
| 1836 | param_ema = str(param) + "_ema" |
| 1837 | if not param_init_net.BlobIsDefined(param_ema): |
| 1838 | param_init_net.ConstantFill([param], param_ema, value=0.0) |
| 1839 | self._aux_params.local.append(param_ema) |
| 1840 | |
| 1841 | net.EMA( |
| 1842 | [param, param_ema, iteration], |
| 1843 | [param, param_ema], |
| 1844 | ema_start=self.ema_start, |
| 1845 | ema_end=self.ema_end, |
| 1846 | ema_step=self.ema_step, |
| 1847 | ema_alpha=self.ema_alpha, |
| 1848 | ) |
| 1849 | |
| 1850 | if self.ema_teacher_enabled: |
nothing calls this directly
no test coverage detected