| 52 | } |
| 53 | |
| 54 | Graph* TruncatedNormal(int64 n) { |
| 55 | Graph* g = new Graph(OpRegistry::Global()); |
| 56 | test::graph::TruncatedNormal(g, test::graph::Constant(g, VecShape(n)), |
| 57 | DT_FLOAT); |
| 58 | return g; |
| 59 | } |
| 60 | |
| 61 | #define BM_RNG(DEVICE, RNG) \ |
| 62 | void BM_##DEVICE##_##RNG(int iters, int arg) { \ |