| 115 | BENCHMARK(BM_SGD)->Arg(128 << 10)->Arg(256 << 10); |
| 116 | |
| 117 | static void Adagrad(int32 n, Graph** init_g, Graph** train_g) { |
| 118 | { |
| 119 | Graph* g = new Graph(OpRegistry::Global()); |
| 120 | auto var = Var(g, n); |
| 121 | auto accum = Var(g, n); |
| 122 | auto zero = Zeros(g, n); |
| 123 | test::graph::Assign(g, var, zero); |
| 124 | test::graph::Assign(g, accum, zero); |
| 125 | *init_g = g; |
| 126 | } |
| 127 | { |
| 128 | Graph* g = new Graph(OpRegistry::Global()); |
| 129 | auto var = Var(g, n); |
| 130 | auto accum = Var(g, n); |
| 131 | auto lr = Scalar(g, 0.01); |
| 132 | auto grad = Random(g, n); |
| 133 | test::graph::Multi(g, "ApplyAdagrad", {var, accum, lr, grad}); |
| 134 | *train_g = g; |
| 135 | } |
| 136 | } |
| 137 | |
| 138 | static void BM_Adagrad(int iters, int params) { |
| 139 | const int64 tot = static_cast<int64>(iters) * params; |