(target, dev, np_func, opt_type, *args, **kwargs)
| 66 | |
| 67 | @tvm.testing.parametrize_targets("llvm") |
| 68 | def _test_optimizer(target, dev, np_func, opt_type, *args, **kwargs): |
| 69 | x = relax.Var("x", R.Tensor((3, 3), "float32")) |
| 70 | y = relax.Var("y", R.Tensor((3,), "float32")) |
| 71 | opt = opt_type(*args, **kwargs).init([x, y]) |
| 72 | mod = IRModule.from_expr(opt.get_function().with_attr("global_symbol", "main")) |
| 73 | tvm_func = _legalize_and_build(mod, target, dev)["main"] |
| 74 | |
| 75 | param_arr = [np.random.rand(3, 3).astype(np.float32), np.random.rand(3).astype(np.float32)] |
| 76 | grad_arr = [np.random.rand(3, 3).astype(np.float32), np.random.rand(3).astype(np.float32)] |
| 77 | state_arr = _tvm_to_numpy(opt.state) |
| 78 | |
| 79 | _assert_run_result_same(tvm_func, np_func, [param_arr, grad_arr, state_arr]) |
| 80 | |
| 81 | |
| 82 | lr, weight_decay = tvm.testing.parameters( |
no test coverage detected
searching dependent graphs…