| 346 | BENCHMARK(BM_AddSign)->Arg(128 << 10)->Arg(256 << 10); |
| 347 | |
| 348 | static void PowerSign(int32 n, Graph** init_g, Graph** train_g) { |
| 349 | TensorShape shape({n}); |
| 350 | { |
| 351 | Graph* g = new Graph(OpRegistry::Global()); |
| 352 | auto var = Var(g, n); |
| 353 | auto m = Var(g, n); |
| 354 | auto zero = Zeros(g, n); |
| 355 | test::graph::Assign(g, var, zero); |
| 356 | test::graph::Assign(g, m, zero); |
| 357 | *init_g = g; |
| 358 | } |
| 359 | { |
| 360 | Graph* g = new Graph(OpRegistry::Global()); |
| 361 | auto var = Var(g, n); |
| 362 | auto m = Var(g, n); |
| 363 | auto lr = Scalar(g, 0.01); |
| 364 | auto logbase = Scalar(g, 2); |
| 365 | auto sign_decay = Scalar(g, 0.9); |
| 366 | auto beta = Scalar(g, 0.8); |
| 367 | auto grad = Random(g, n); |
| 368 | test::graph::Multi(g, "ApplyPowerSign", |
| 369 | {var, m, lr, logbase, sign_decay, beta, grad}); |
| 370 | *train_g = g; |
| 371 | } |
| 372 | } |
| 373 | |
| 374 | static void BM_PowerSign(int iters, int params) { |
| 375 | const int64 tot = static_cast<int64>(iters) * params; |