| 13 | |
| 14 | template <typename T, typename T_ACC = float> |
| 15 | void update( |
| 16 | _megdnn_tensor_in m_t_1, _megdnn_tensor_in v_t_1, _megdnn_tensor_in lamb_param, |
| 17 | _megdnn_tensor_in grad, _megdnn_tensor_out m_t, _megdnn_tensor_out v_t, |
| 18 | _megdnn_tensor_out new_param, const Param& param) { |
| 19 | float beta_1 = param.beta_1; |
| 20 | float beta_2 = param.beta_2; |
| 21 | float step = param.step; |
| 22 | float lr = param.lr; |
| 23 | float weight_decay = param.weight_decay; |
| 24 | float eps = param.eps; |
| 25 | bool bias_correction = param.bias_correction; |
| 26 | bool always_adapt = param.always_adapt; |
| 27 | |
| 28 | size_t total_elem = lamb_param.layout.total_nr_elems(); |
| 29 | T_ACC mt, vt, bc_1, bc_2, rt, d_norm = 0; |
| 30 | bc_1 = bias_correction ? 1 - pow(beta_1, step) : 1; |
| 31 | bc_2 = bias_correction ? 1 - pow(beta_2, step) : 1; |
| 32 | |
| 33 | for (size_t i = 0; i < total_elem; i++) { |
| 34 | mt = m_t.ptr<T_ACC>()[i] = beta_1 * m_t_1.ptr<T_ACC>()[i] + |
| 35 | (1 - beta_1) * static_cast<T_ACC>(grad.ptr<T>()[i]); |
| 36 | vt = v_t.ptr<T_ACC>()[i] = |
| 37 | beta_2 * v_t_1.ptr<T_ACC>()[i] + |
| 38 | (1 - beta_2) * std::pow(static_cast<T_ACC>(grad.ptr<T>()[i]), 2); |
| 39 | rt = (mt / bc_1) / (sqrt(vt / bc_2) + eps); |
| 40 | if (weight_decay != 0) { |
| 41 | rt += lamb_param.ptr<T_ACC>()[i] * weight_decay; |
| 42 | } |
| 43 | d_norm += rt * rt; |
| 44 | } |
| 45 | d_norm = sqrt(d_norm); |
| 46 | auto get_norm = [=](_megdnn_tensor_in norm) -> T_ACC { |
| 47 | return sqrt(std::accumulate( |
| 48 | norm.ptr<T_ACC>(), norm.ptr<T_ACC>() + total_elem, 0, |
| 49 | [](T_ACC t1, T_ACC t2) -> T_ACC { return t1 + t2 * t2; })); |
| 50 | }; |
| 51 | T_ACC p_norm = get_norm(lamb_param), trust_ratio = 1; |
| 52 | if ((always_adapt || weight_decay > 0) && p_norm > 0 && d_norm > 0) { |
| 53 | trust_ratio = p_norm / d_norm; |
| 54 | } |
| 55 | for (size_t i = 0; i < total_elem; i++) { |
| 56 | mt = m_t.ptr<T_ACC>()[i]; |
| 57 | vt = v_t.ptr<T_ACC>()[i]; |
| 58 | rt = (mt / bc_1) / (sqrt(vt / bc_2) + eps); |
| 59 | if (weight_decay != 0) { |
| 60 | rt += lamb_param.ptr<T_ACC>()[i] * weight_decay; |
| 61 | } |
| 62 | new_param.ptr<T_ACC>()[i] = lamb_param.ptr<T_ACC>()[i] - lr * trust_ratio * rt; |
| 63 | } |
| 64 | } |
| 65 | |
| 66 | } // namespace |
| 67 |
no test coverage detected