(target, dev, lr, betas, eps, weight_decay)
| 142 | |
| 143 | @tvm.testing.parametrize_targets("llvm") |
| 144 | def test_adam(target, dev, lr, betas, eps, weight_decay): |
| 145 | def np_func(param_tuple, grad_tuple, state_tuple): |
| 146 | num_steps = state_tuple[0] |
| 147 | num_steps_new = num_steps + 1 |
| 148 | |
| 149 | param_tuple_new = [] |
| 150 | state_tuple_new = [None] * len(state_tuple) # type: ignore |
| 151 | state_tuple_new[0] = num_steps_new |
| 152 | state_tuple_new[1] = state_tuple[1] * betas[0] |
| 153 | state_tuple_new[2] = state_tuple[2] * betas[1] |
| 154 | |
| 155 | for i in range(len(param_tuple)): |
| 156 | param = param_tuple[i] |
| 157 | grad = grad_tuple[i] |
| 158 | m = state_tuple[i + 3] |
| 159 | v = state_tuple[i + 3 + len(param_tuple)] |
| 160 | grad = grad + weight_decay * param |
| 161 | m = betas[0] * m + (1 - betas[0]) * grad |
| 162 | v = betas[1] * v + (1 - betas[1]) * grad * grad |
| 163 | m_hat = m / (1 - betas[0] ** num_steps_new) |
| 164 | v_hat = v / (1 - betas[1] ** num_steps_new) |
| 165 | param = param - lr * m_hat / (np.sqrt(v_hat) + eps) |
| 166 | param_tuple_new.append(param) |
| 167 | state_tuple_new[i + 3] = m |
| 168 | state_tuple_new[i + 3 + len(param_tuple)] = v |
| 169 | |
| 170 | return param_tuple_new, state_tuple_new |
| 171 | |
| 172 | _test_optimizer(target, dev, np_func, Adam, lr, betas, eps, weight_decay) |
| 173 | |
| 174 | |
| 175 | if __name__ == "__main__": |
nothing calls this directly
no test coverage detected
searching dependent graphs…