(self, net, param_init_net, param_info)
| 1983 | self.init_kwargs = kwargs |
| 1984 | |
| 1985 | def _run(self, net, param_init_net, param_info): |
| 1986 | param = param_info.blob |
| 1987 | grad = param_info.grad |
| 1988 | |
| 1989 | assert self.alpha > 0 |
| 1990 | assert not isinstance( |
| 1991 | grad, core.GradientSlice |
| 1992 | ), "RmsPropOptimizer doesn't support sparse gradients" |
| 1993 | |
| 1994 | dev = scope.CurrentDeviceScope() |
| 1995 | if dev is None: |
| 1996 | dev = core.DeviceOption(caffe2_pb2.CPU) |
| 1997 | |
| 1998 | ONE = param_init_net.ConstantFill( |
| 1999 | [], "ONE_{}_{}".format(dev.device_type, dev.device_id), shape=[1], value=1.0 |
| 2000 | ) |
| 2001 | |
| 2002 | lr, _ = self.build_lr( |
| 2003 | net, |
| 2004 | param_init_net, |
| 2005 | base_learning_rate=-self.alpha, |
| 2006 | policy=self.policy, |
| 2007 | **(self.init_kwargs) |
| 2008 | ) |
| 2009 | |
| 2010 | grad_o = param_init_net.ConstantFill( |
| 2011 | [param], str(param) + "_grad_o", values=0.0 |
| 2012 | ) |
| 2013 | |
| 2014 | ms = param_init_net.ConstantFill( |
| 2015 | [param], str(param) + "_mean_squares", values=0.0 |
| 2016 | ) |
| 2017 | |
| 2018 | mom = param_init_net.ConstantFill([param], str(param) + "_momentum", values=0.0) |
| 2019 | |
| 2020 | self._aux_params.local.append(ms) |
| 2021 | self._aux_params.local.append(mom) |
| 2022 | |
| 2023 | net.RmsProp( |
| 2024 | [grad, ms, mom, ONE], |
| 2025 | [grad_o, ms, mom], |
| 2026 | decay=self.decay, |
| 2027 | momentum=self.momentum, |
| 2028 | epsilon=self.epsilon, |
| 2029 | engine=self.engine, |
| 2030 | ) |
| 2031 | |
| 2032 | net.MomentumSGDUpdate([grad_o, mom, lr, param], [grad_o, mom, param]) |
| 2033 | |
| 2034 | def scale_learning_rate(self, scale): |
| 2035 | self.alpha *= scale |
nothing calls this directly
no test coverage detected