diff --git a/src/mlpack/methods/reinforcement_learning/async_learning_impl.hpp b/src/mlpack/methods/reinforcement_learning/async_learning_impl.hpp index 2cbb1c5c40..5cd82f06a8 100644 --- a/src/mlpack/methods/reinforcement_learning/async_learning_impl.hpp +++ b/src/mlpack/methods/reinforcement_learning/async_learning_impl.hpp @@ -97,24 +97,20 @@ void AsyncLearning< * compiler. We can switch to OpenMP task once MSVC supports OpenMP 3.0. */ size_t numThreads = 0; - #ifdef HAS_OPENMP - #pragma omp parallel reduction(+:numThreads) - #endif + #pragma omp parallel reduction(+:numThreads) numThreads++; Log::Debug << numThreads << " threads will be used in total." << std::endl; - #ifdef HAS_OPENMP - #pragma omp parallel for shared(stop, workers, tasks, learningNetwork, \ - targetNetwork, totalSteps, policy) - #endif + #pragma omp parallel for shared(stop, workers, tasks, learningNetwork, \ + targetNetwork, totalSteps, policy) for (omp_size_t i = 0; i < numThreads; ++i) { - #ifdef HAS_OPENMP - #pragma omp critical - #endif + #pragma omp critical { - Log::Debug << "Thread " << omp_get_thread_num() << - " started." << std::endl; + #ifdef HAS_OPENMP + Log::Debug << "Thread " << omp_get_thread_num() << + " started." << std::endl; + #endif } while (!stop) { diff --git a/src/mlpack/methods/reinforcement_learning/worker/one_step_q_learning_worker.hpp b/src/mlpack/methods/reinforcement_learning/worker/one_step_q_learning_worker.hpp index 6749fe6d12..561891c707 100644 --- a/src/mlpack/methods/reinforcement_learning/worker/one_step_q_learning_worker.hpp +++ b/src/mlpack/methods/reinforcement_learning/worker/one_step_q_learning_worker.hpp @@ -112,9 +112,7 @@ class OneStepQLearningWorker return false; } - #ifdef HAS_OPENMP - #pragma omp atomic - #endif + #pragma omp atomic totalSteps++; pending[pendingIndex] = std::make_tuple(state, action, reward, nextState); @@ -131,9 +129,7 @@ class OneStepQLearningWorker // Compute the target state-action value. arma::colvec actionValue; - #ifdef HAS_OPENMP - #pragma omp critical - #endif + #pragma omp critical { targetNetwork.Predict( std::get<3>(transition).Encode(), actionValue); @@ -175,9 +171,7 @@ class OneStepQLearningWorker // Update global target network. if (totalSteps % config.TargetNetworkSyncInterval() == 0) { - #ifdef HAS_OPENMP - #pragma omp critical - #endif + #pragma omp critical { targetNetwork = learningNetwork; } }