diff --git a/cpp_thread_test/dgemm_thread_safety_mixed.cpp b/cpp_thread_test/dgemm_thread_safety_mixed.cpp index 1ad021bb3..62d6b8383 100644 --- a/cpp_thread_test/dgemm_thread_safety_mixed.cpp +++ b/cpp_thread_test/dgemm_thread_safety_mixed.cpp @@ -14,6 +14,21 @@ #endif #include "cpp_thread_safety_common.h" +std::atomic callbackInvocations(0); + +void thread_callback(int sync, openblas_dojob_callback doJob, int numJobs, + size_t jobDataElementSize, void* jobData, int doJobData){ + (void)sync; + callbackInvocations.fetch_add(1, std::memory_order_relaxed); + std::vector workers; + workers.reserve(numJobs); + char* jobs = static_cast(jobData); + for(int i=0; i& transA, std::vector& noTransA, std::vector& B, double* firstOutput, double* secondOutput, const blasint randomMatSize, const bool sameVariant){ cblas_dgemm(CblasRowMajor, CblasTrans, CblasNoTrans, randomMatSize, 2, 2, 1.0, &transA[0], randomMatSize, &B[0], 2, 0.0, firstOutput, 2); if (sameVariant) @@ -48,26 +63,31 @@ int main(int argc, char* argv[]){ uint32_t numTestRounds = 200; uint32_t maxHwThreads = GetMaxHwThreads(); bool sameVariant = false; + bool useCallback = false; if (maxHwThreads < numConcurrentThreads) numConcurrentThreads = maxHwThreads; - if (argc != 1 && argc != 4 && argc != 5){ - std::cout<<"ERROR: expected zero arguments, or: [sameVariant]"< positionalArgs; + for (int i = 1; i < argc; i++){ + std::cout< [sameVariant]] [--callback]"< cliArgs; - for (int i = 1; i < argc; i++){ - cliArgs.push_back(argv[i]); - std::cout<(matrixElements) * 2 * 8 + static_cast(outputElements) * (2 + 2 * numConcurrentThreads) * 8)/static_cast(1024*1024)<<" MiB of RAM\n"<