| 147 | BENCHMARK(BM_Adagrad)->Arg(128 << 10)->Arg(256 << 10); |
| 148 | |
| 149 | static void SparseAdagrad(int32 m, int32 n, Graph** init_g, Graph** train_g) { |
| 150 | { |
| 151 | Graph* g = new Graph(OpRegistry::Global()); |
| 152 | auto var = Var(g, m, n); |
| 153 | auto accum = Var(g, m, n); |
| 154 | auto zero = Zeros(g, m, n); |
| 155 | test::graph::Assign(g, var, zero); |
| 156 | test::graph::Assign(g, accum, zero); |
| 157 | *init_g = g; |
| 158 | } |
| 159 | { |
| 160 | Graph* g = new Graph(OpRegistry::Global()); |
| 161 | auto var = Var(g, m, n); |
| 162 | auto accum = Var(g, m, n); |
| 163 | auto lr = Scalar(g, 0.01); |
| 164 | auto grad = Random(g, m, n); |
| 165 | auto indices = Iota(g, m); |
| 166 | test::graph::Multi(g, "SparseApplyAdagrad", |
| 167 | {var, accum, lr, grad, indices}); |
| 168 | *train_g = g; |
| 169 | } |
| 170 | } |
| 171 | static void BM_SparseAdagrad(int iters, int m, int n) { |
| 172 | const int64 tot = static_cast<int64>(iters) * m * n; |
| 173 | testing::UseRealTime(); |