diff --git a/winml/lib/Api/LearningModelDevice.cpp b/winml/lib/Api/LearningModelDevice.cpp index ef0bf8e430..011ca06fdf 100644 --- a/winml/lib/Api/LearningModelDevice.cpp +++ b/winml/lib/Api/LearningModelDevice.cpp @@ -53,16 +53,12 @@ Windows::AI::MachineLearning::LearningModelDevice LearningModelDevice::CreateFro WINML_CATCH_ALL std::shared_ptr<::Windows::AI::MachineLearning::ConverterResourceStore> LearningModelDevice::TensorizerStore() { - if (m_tensorizerStore == nullptr) { - m_tensorizerStore = ::Windows::AI::MachineLearning::ConverterResourceStore::Create(5); - } + std::call_once(m_tensorizerStoreInitialized, [this](){ m_tensorizerStore = ::Windows::AI::MachineLearning::ConverterResourceStore::Create(5); }); return m_tensorizerStore; } std::shared_ptr<::Windows::AI::MachineLearning::ConverterResourceStore> LearningModelDevice::DetensorizerStore() { - if (m_detensorizerStore == nullptr) { - m_detensorizerStore = ::Windows::AI::MachineLearning::ConverterResourceStore::Create(5); - } + std::call_once(m_detensorizerStoreInitialized, [this](){ m_detensorizerStore = ::Windows::AI::MachineLearning::ConverterResourceStore::Create(5); }); return m_detensorizerStore; } diff --git a/winml/lib/Api/LearningModelDevice.h b/winml/lib/Api/LearningModelDevice.h index ce20a4b765..26639d7c7e 100644 --- a/winml/lib/Api/LearningModelDevice.h +++ b/winml/lib/Api/LearningModelDevice.h @@ -84,7 +84,9 @@ struct LearningModelDevice : LearningModelDeviceT m_detensorizerStore; + std::once_flag m_detensorizerStoreInitialized; std::shared_ptr m_tensorizerStore; + std::once_flag m_tensorizerStoreInitialized; std::unique_ptr m_deviceCache; }; diff --git a/winml/test/concurrency/ConcurrencyTests.cpp b/winml/test/concurrency/ConcurrencyTests.cpp index 1df5875a08..9200f8776c 100644 --- a/winml/test/concurrency/ConcurrencyTests.cpp +++ b/winml/test/concurrency/ConcurrencyTests.cpp @@ -250,9 +250,11 @@ void MultiThreadMultiSessionOnDevice(const LearningModelDevice& device) { void MultiThreadMultiSession() { MultiThreadMultiSessionOnDevice(LearningModelDeviceKind::Cpu); - if (GPUTEST_ENABLED) { - MultiThreadMultiSessionOnDevice(LearningModelDeviceKind::DirectX); - } +} + +void MultiThreadMultiSessionGpu() { + GPUTEST + MultiThreadMultiSessionOnDevice(LearningModelDeviceKind::DirectX); } // Create different sessions for each thread, and evaluate @@ -318,9 +320,11 @@ void MultiThreadSingleSessionOnDevice(const LearningModelDevice& device) { void MultiThreadSingleSession() { MultiThreadSingleSessionOnDevice(LearningModelDeviceKind::Cpu); - if (GPUTEST_ENABLED) { - MultiThreadSingleSessionOnDevice(LearningModelDeviceKind::DirectX); - } +} + +void MultiThreadSingleSessionGpu() { + GPUTEST + MultiThreadSingleSessionOnDevice(LearningModelDeviceKind::DirectX); } } @@ -330,7 +334,9 @@ const ConcurrencyTestsApi& getapi() { LoadBindEvalSqueezenetRealDataWithValidationConcurrently, MultiThreadLoadModel, MultiThreadMultiSession, + MultiThreadMultiSessionGpu, MultiThreadSingleSession, + MultiThreadSingleSessionGpu, EvalAsyncDifferentModels, EvalAsyncDifferentSessions, EvalAsyncDifferentBindings diff --git a/winml/test/concurrency/ConcurrencyTests.h b/winml/test/concurrency/ConcurrencyTests.h index 45e0518870..2bfb8a4896 100644 --- a/winml/test/concurrency/ConcurrencyTests.h +++ b/winml/test/concurrency/ConcurrencyTests.h @@ -10,7 +10,9 @@ struct ConcurrencyTestsApi VoidTest LoadBindEvalSqueezenetRealDataWithValidationConcurrently; VoidTest MultiThreadLoadModel; VoidTest MultiThreadMultiSession; + VoidTest MultiThreadMultiSessionGpu; VoidTest MultiThreadSingleSession; + VoidTest MultiThreadSingleSessionGpu; VoidTest EvalAsyncDifferentModels; VoidTest EvalAsyncDifferentSessions; VoidTest EvalAsyncDifferentBindings; @@ -21,7 +23,9 @@ WINML_TEST_CLASS_BEGIN_WITH_SETUP(ConcurrencyTests, ConcurrencyTestsApiSetup) WINML_TEST(ConcurrencyTests, LoadBindEvalSqueezenetRealDataWithValidationConcurrently) WINML_TEST(ConcurrencyTests, MultiThreadLoadModel) WINML_TEST(ConcurrencyTests, MultiThreadMultiSession) +WINML_TEST(ConcurrencyTests, MultiThreadMultiSessionGpu) WINML_TEST(ConcurrencyTests, MultiThreadSingleSession) +WINML_TEST(ConcurrencyTests, MultiThreadSingleSessionGpu) WINML_TEST(ConcurrencyTests, EvalAsyncDifferentModels) WINML_TEST(ConcurrencyTests, EvalAsyncDifferentSessions) WINML_TEST(ConcurrencyTests, EvalAsyncDifferentBindings)