| 87 | } |
| 88 | |
| 89 | static void SGD(int32 n, Graph** init_g, Graph** train_g) { |
| 90 | { |
| 91 | Graph* g = new Graph(OpRegistry::Global()); |
| 92 | auto var = Var(g, n); |
| 93 | test::graph::Assign(g, var, Zeros(g, n)); |
| 94 | *init_g = g; |
| 95 | } |
| 96 | { |
| 97 | Graph* g = new Graph(OpRegistry::Global()); |
| 98 | auto var = Var(g, n); |
| 99 | auto lr = Scalar(g, 0.01); |
| 100 | auto grad = Random(g, n); |
| 101 | test::graph::Multi(g, "ApplyGradientDescent", {var, lr, grad}); |
| 102 | *train_g = g; |
| 103 | } |
| 104 | } |
| 105 | |
| 106 | static void BM_SGD(int iters, int params) { |
| 107 | const int64 tot = static_cast<int64>(iters) * params; |