@@ -19,8 +19,33 @@ int main(int argc, char* argv[]) {
1919 model_name = argv[++i];
2020 } else if (std::string (argv[i]) == " --onednn" ) {
2121 options.backend = Backend::kOneDnn ;
22- } else if (std::string (argv[i]) == " --parallel" ) {
23- options.parallel = true ;
22+ if (options.isParallel ()) {
23+ std::cout << " Warning: oneDNN backend is not compatible with parallel "
24+ " execution. Disabling parallelism."
25+ << ' \n ' ;
26+ options.setParallelBackend (ParBackend::kSeq );
27+ }
28+ } else if (std::string (argv[i]) == " --parallel" && i + 1 < argc) {
29+ if (options.backend == Backend::kOneDnn ) {
30+ std::cout << " Warning: Parallel execution is not compatible with "
31+ " oneDNN backend. Ignoring --parallel option."
32+ << ' \n ' ;
33+ i++;
34+ continue ;
35+ }
36+
37+ std::string backend_str = argv[++i];
38+ if (backend_str == " tbb" ) {
39+ options.setParallelBackend (ParBackend::kTbb );
40+ } else if (backend_str == " threads" || backend_str == " stl" ) {
41+ options.setParallelBackend (ParBackend::kThreads );
42+ } else if (backend_str == " omp" ) {
43+ options.setParallelBackend (ParBackend::kOmp );
44+ } else {
45+ std::cerr << " Unknown parallel backend: " << backend_str
46+ << " . Using default (Threads)." << ' \n ' ;
47+ options.setParallelBackend (ParBackend::kThreads );
48+ }
2449 } else if (std::string (argv[i]) == " --threads" && i + 1 < argc) {
2550 options.threads = std::stoi (argv[++i]);
2651 }
0 commit comments