MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / update

Function update

dnn/src/naive/lamb/opr_impl.cpp:15–64  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

13
14template <typename T, typename T_ACC = float>
15void 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

Calls 3

powFunction · 0.50
sqrtFunction · 0.50
total_nr_elemsMethod · 0.45

Tested by

no test coverage detected