Check HAS_OPENMP
This commit is contained in:
@@ -86,8 +86,10 @@ void AsyncLearning<
|
||||
tasks.push(i);
|
||||
|
||||
// Synchronization for get/put task.
|
||||
omp_lock_t tasksLock;
|
||||
omp_init_lock(&tasksLock);
|
||||
#ifdef HAS_OPENMP
|
||||
omp_lock_t tasksLock;
|
||||
omp_init_lock(&tasksLock);
|
||||
#endif
|
||||
|
||||
/**
|
||||
* Compute the number of threads for the for-loop. In general, we should use
|
||||
@@ -95,15 +97,21 @@ void AsyncLearning<
|
||||
* compiler. We can switch to OpenMP task once MSVC supports OpenMP 3.0.
|
||||
*/
|
||||
size_t numThreads = 0;
|
||||
#pragma omp parallel reduction(+:numThreads)
|
||||
#ifdef HAS_OPENMP
|
||||
#pragma omp parallel reduction(+:numThreads)
|
||||
#endif
|
||||
numThreads++;
|
||||
Log::Debug << numThreads << " threads will be used in total." << std::endl;
|
||||
|
||||
#pragma omp parallel for shared(stop, workers, tasks, learningNetwork, \
|
||||
targetNetwork, totalSteps, policy)
|
||||
#ifdef HAS_OPENMP
|
||||
#pragma omp parallel for shared(stop, workers, tasks, learningNetwork, \
|
||||
targetNetwork, totalSteps, policy)
|
||||
#endif
|
||||
for (omp_size_t i = 0; i < numThreads; ++i)
|
||||
{
|
||||
#pragma omp critical
|
||||
#ifdef HAS_OPENMP
|
||||
#pragma omp critical
|
||||
#endif
|
||||
{
|
||||
Log::Debug << "Thread " << omp_get_thread_num() <<
|
||||
" started." << std::endl;
|
||||
@@ -111,14 +119,19 @@ void AsyncLearning<
|
||||
while (!stop)
|
||||
{
|
||||
// Assign task to current thread from queue.
|
||||
omp_set_lock(&tasksLock);
|
||||
if (tasks.empty()) {
|
||||
omp_unset_lock(&tasksLock);
|
||||
continue;
|
||||
}
|
||||
#ifdef HAS_OPENMP
|
||||
omp_set_lock(&tasksLock);
|
||||
if (tasks.empty())
|
||||
{
|
||||
omp_unset_lock(&tasksLock);
|
||||
continue;
|
||||
}
|
||||
#endif
|
||||
size_t task = tasks.front();
|
||||
tasks.pop();
|
||||
omp_unset_lock(&tasksLock);
|
||||
#ifdef HAS_OPENMP
|
||||
omp_unset_lock(&tasksLock);
|
||||
#endif
|
||||
|
||||
// Get corresponding worker.
|
||||
WorkerType& worker = workers[task];
|
||||
@@ -130,9 +143,13 @@ void AsyncLearning<
|
||||
}
|
||||
|
||||
// Put task back to queue.
|
||||
omp_set_lock(&tasksLock);
|
||||
#ifdef HAS_OPENMP
|
||||
omp_set_lock(&tasksLock);
|
||||
#endif
|
||||
tasks.push(task);
|
||||
omp_unset_lock(&tasksLock);
|
||||
#ifdef HAS_OPENMP
|
||||
omp_unset_lock(&tasksLock);
|
||||
#endif
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -64,7 +64,7 @@ class OneStepQLearningWorker
|
||||
void Initialize(NetworkType& learningNetwork)
|
||||
{
|
||||
updater.Initialize(learningNetwork.Parameters().n_rows,
|
||||
learningNetwork.Parameters().n_cols);
|
||||
learningNetwork.Parameters().n_cols);
|
||||
// Build local network.
|
||||
network = learningNetwork;
|
||||
}
|
||||
@@ -112,7 +112,9 @@ class OneStepQLearningWorker
|
||||
return false;
|
||||
}
|
||||
|
||||
#pragma omp atomic
|
||||
#ifdef HAS_OPENMP
|
||||
#pragma omp atomic
|
||||
#endif
|
||||
totalSteps++;
|
||||
|
||||
pending[pendingIndex] = std::make_tuple(state, action, reward, nextState);
|
||||
@@ -129,7 +131,9 @@ class OneStepQLearningWorker
|
||||
|
||||
// Compute the target state-action value.
|
||||
arma::colvec actionValue;
|
||||
#pragma omp critical
|
||||
#ifdef HAS_OPENMP
|
||||
#pragma omp critical
|
||||
#endif
|
||||
{
|
||||
targetNetwork.Predict(
|
||||
std::get<3>(transition).Encode(), actionValue);
|
||||
@@ -171,7 +175,9 @@ class OneStepQLearningWorker
|
||||
// Update global target network.
|
||||
if (totalSteps % config.TargetNetworkSyncInterval() == 0)
|
||||
{
|
||||
#pragma omp critical
|
||||
#ifdef HAS_OPENMP
|
||||
#pragma omp critical
|
||||
#endif
|
||||
{ targetNetwork = learningNetwork; }
|
||||
}
|
||||
|
||||
|
||||
@@ -38,7 +38,9 @@ BOOST_AUTO_TEST_CASE(OneStepQLearningTest)
|
||||
* This is for the Travis CI server, in your own machine you shuold use more
|
||||
* threads.
|
||||
*/
|
||||
omp_set_num_threads(1);
|
||||
#ifdef HAS_OPENMP
|
||||
omp_set_num_threads(1);
|
||||
#endif
|
||||
|
||||
// Set up the network.
|
||||
FFN<MeanSquaredError<>, GaussianInitialization> model(MeanSquaredError<>(),
|
||||
|
||||
Reference in New Issue
Block a user