| 23 | #include "tensorflow/core/platform/test_benchmark.h" |
| 24 | |
| 25 | GTEST_API_ int main(int real_argc, char** real_argv) { |
| 26 | std::vector<tensorflow::Flag> flag_list; |
| 27 | tensorflow::AppendMarkForCompilationPassFlags(&flag_list); |
| 28 | auto usage = tensorflow::Flags::Usage(real_argv[0], flag_list); |
| 29 | |
| 30 | std::vector<char*> args; |
| 31 | |
| 32 | args.reserve(real_argc + 1); |
| 33 | for (int i = 0; i < real_argc; i++) { |
| 34 | args.push_back(real_argv[i]); |
| 35 | } |
| 36 | |
| 37 | struct FreeDeleter { |
| 38 | void operator()(char* ptr) { free(ptr); } |
| 39 | }; |
| 40 | |
| 41 | std::unique_ptr<char, FreeDeleter> enable_global_jit_arg( |
| 42 | strdup("--tf_xla_cpu_global_jit=true")); |
| 43 | args.push_back(enable_global_jit_arg.get()); |
| 44 | |
| 45 | std::unique_ptr<char, FreeDeleter> reduce_min_cluster_size_arg( |
| 46 | strdup("--tf_xla_min_cluster_size=2")); |
| 47 | args.push_back(reduce_min_cluster_size_arg.get()); |
| 48 | |
| 49 | int argc = args.size(); |
| 50 | |
| 51 | if (!tensorflow::Flags::Parse(&argc, &args.front(), flag_list)) { |
| 52 | LOG(ERROR) << "\n" << usage; |
| 53 | return 2; |
| 54 | } |
| 55 | |
| 56 | testing::InitGoogleTest(&argc, &args.front()); |
| 57 | return RUN_ALL_TESTS(); |
| 58 | } |