MCPcopy Create free account
hub / github.com/apache/tvm / test_adam

Function test_adam

tests/python/relax/test_training_optimizer_numeric.py:144–172  ·  view source on GitHub ↗
(target, dev, lr, betas, eps, weight_decay)

Source from the content-addressed store, hash-verified

142
143@tvm.testing.parametrize_targets("llvm")
144def 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
175if __name__ == "__main__":

Callers

nothing calls this directly

Calls 1

_test_optimizerFunction · 0.85

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…