From a1d20ec8aa0cb3a92e173640c0b0bd60985204f2 Mon Sep 17 00:00:00 2001 From: Jeff Bloomfield Date: Sun, 6 Aug 2023 09:49:39 -0700 Subject: [PATCH] Fix DML regression from allocator refactor and enable unrounded weight allocations through ORT API --- .../inc/DmlExecutionProvider.h | 1 - .../src/BucketizedBufferAllocator.cpp | 7 ++++++- .../src/BucketizedBufferAllocator.h | 8 +++++++- .../src/ExecutionProvider.cpp | 18 +++++------------- .../src/ExecutionProvider.h | 8 -------- .../core/providers/dml/dml_provider_factory.cc | 12 ------------ winml/adapter/winml_adapter_apis.h | 1 - winml/adapter/winml_adapter_c_api.cpp | 1 - winml/adapter/winml_adapter_c_api.h | 8 -------- winml/adapter/winml_adapter_dml.cpp | 17 ----------------- .../Api.Ort/OnnxruntimeDmlSessionBuilder.cpp | 3 --- winml/test/adapter/AdapterDmlEpTest.cpp | 9 --------- winml/test/adapter/AdapterDmlEpTest.h | 2 -- 13 files changed, 18 insertions(+), 77 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h index fd96bd812d..52018500b1 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h @@ -32,7 +32,6 @@ namespace Dml ID3D12Resource* GetD3D12ResourceFromAllocation(onnxruntime::IAllocator* allocator, void* ptr); void FlushContext(onnxruntime::IExecutionProvider* provider); - void SetDefaultRoundingMode(onnxruntime::IExecutionProvider* provider, AllocatorRoundingMode roundingMode); void ReleaseCompletedReferences(onnxruntime::IExecutionProvider* provider); onnxruntime::common::Status CopyTensor( diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.cpp index 81550714a6..5dbea41901 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.cpp @@ -81,6 +81,11 @@ namespace Dml } void* BucketizedBufferAllocator::Alloc(size_t size) + { + return Alloc(size, m_defaultRoundingMode); + } + + void* BucketizedBufferAllocator::Alloc(size_t size, AllocatorRoundingMode roundingMode) { // For some reason lotus likes requesting 0 bytes of memory size = std::max(1, size); @@ -90,7 +95,7 @@ namespace Dml uint64_t bucketSize = 0; // Use a pooled resource if the size (post rounding, if requested) matches a bucket size - if (m_defaultRoundingMode == AllocatorRoundingMode::Enabled || size == GetBucketSizeFromIndex(GetBucketIndexFromSize(size))) + if (roundingMode == AllocatorRoundingMode::Enabled || size == GetBucketSizeFromIndex(GetBucketIndexFromSize(size))) { Bucket* bucket = nullptr; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.h index 26090e492f..ab47fdc8ec 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/BucketizedBufferAllocator.h @@ -107,6 +107,7 @@ namespace Dml void SetDefaultRoundingMode(AllocatorRoundingMode roundingMode); public: // onnxruntime::IAllocator + void* Alloc(size_t size, AllocatorRoundingMode roundingMode); void* Alloc(size_t size) final; void Free(void* p) final; @@ -142,7 +143,12 @@ namespace Dml std::vector m_pool; size_t m_currentAllocationId = 0; uint64_t m_currentResourceId = 0; - AllocatorRoundingMode m_defaultRoundingMode = AllocatorRoundingMode::Enabled; + + // Unless specifically requested, allocation sizes are not rounded to enable pooling + // until SetDefaultRoundingMode is called. This should be done at completion of session + // initialization. + AllocatorRoundingMode m_defaultRoundingMode = AllocatorRoundingMode::Disabled; + std::shared_ptr m_context; std::unique_ptr m_subAllocator; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp index d9338dbf94..75da803817 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp @@ -116,7 +116,7 @@ namespace Dml ORT_TRY { ComPtr allocation; - allocation.Attach(static_cast(m_allocator->Alloc(size))); + allocation.Attach(static_cast(m_allocator->Alloc(size, roundingMode))); const auto* allocInfo = m_allocator->DecodeDataHandle(allocation.Get()); @@ -204,7 +204,6 @@ namespace Dml D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS, D3D12_RESOURCE_STATE_UNORDERED_ACCESS, std::make_unique(m_d3d12Device.Get())); - m_allocator->SetDefaultRoundingMode(m_defaultRoundingMode); // TODO(leca): REVIEW: the original code is able to set roudingMode multiple times during alloc's life time. Double check this case is not happening m_context->SetAllocator(m_allocator); // CPU Allocator used to create buffers for the MemcpyFromHost, Shape and Size operators. m_cpuInputAllocator = std::make_shared(OrtMemType::OrtMemTypeCPUInput); @@ -970,11 +969,6 @@ namespace Dml m_context->ReleaseCompletedReferences(); } - void ExecutionProviderImpl::SetDefaultRoundingMode(AllocatorRoundingMode roundingMode) - { - m_defaultRoundingMode = roundingMode; - } - void ExecutionProviderImpl::ReleaseCompletedReferences() { m_context->ReleaseCompletedReferences(); @@ -1132,6 +1126,10 @@ namespace Dml m_context->ReleaseCompletedReferences(); m_uploadHeap->Trim(); + // Allocations after this point are potentially transient and their sizes are + // rounded to enable pooling. + m_allocator->SetDefaultRoundingMode(AllocatorRoundingMode::Enabled); + return onnxruntime::common::Status::OK(); } @@ -1155,12 +1153,6 @@ namespace Dml dmlexecutionprovider->Flush(); } - void SetDefaultRoundingMode(onnxruntime::IExecutionProvider* provider, AllocatorRoundingMode roundingMode) - { - ExecutionProvider* dmlexecutionprovider = static_cast(provider); - dmlexecutionprovider->SetDefaultRoundingMode(roundingMode); - } - void ReleaseCompletedReferences(onnxruntime::IExecutionProvider * provider) { ExecutionProvider* dmlexecutionprovider = static_cast(provider); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h index 1dfdf5cf2d..a21d996daa 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h @@ -127,8 +127,6 @@ namespace Dml STDMETHOD_(void, Flush)() const override; STDMETHOD_(void, FlushAndSync)() const override; - void SetDefaultRoundingMode(AllocatorRoundingMode roundingMode); - // Waits for flushed work, discards unflushed work, and discards associated references to // prevent circular references. Must be the last call on the object before destruction. void Close() override; @@ -199,7 +197,6 @@ namespace Dml bool m_closed = false; mutable std::chrono::time_point m_lastUploadFlushTime; static constexpr std::chrono::milliseconds m_batchFlushInterval = std::chrono::milliseconds(10); - AllocatorRoundingMode m_defaultRoundingMode = AllocatorRoundingMode::Enabled; }; class DataTransfer : public onnxruntime::IDataTransfer @@ -288,11 +285,6 @@ namespace Dml return m_impl->Flush(); } - void SetDefaultRoundingMode(AllocatorRoundingMode roundingMode) - { - return m_impl->SetDefaultRoundingMode(roundingMode); - } - void ReleaseCompletedReferences() { return m_impl->ReleaseCompletedReferences(); diff --git a/onnxruntime/core/providers/dml/dml_provider_factory.cc b/onnxruntime/core/providers/dml/dml_provider_factory.cc index 8652c10430..6a2740e436 100644 --- a/onnxruntime/core/providers/dml/dml_provider_factory.cc +++ b/onnxruntime/core/providers/dml/dml_provider_factory.cc @@ -32,27 +32,20 @@ struct DMLProviderFactory : IExecutionProviderFactory { ~DMLProviderFactory() override {} std::unique_ptr CreateProvider() override; - void SetDefaultRoundingMode(AllocatorRoundingMode rounding_mode); void SetMetacommandsEnabled(bool metacommands_enabled); private: ComPtr dml_device_{}; ComPtr cmd_queue_{}; - AllocatorRoundingMode rounding_mode_ = AllocatorRoundingMode::Enabled; bool metacommands_enabled_ = true; }; std::unique_ptr DMLProviderFactory::CreateProvider() { auto provider = Dml::CreateExecutionProvider(dml_device_.Get(), cmd_queue_.Get(), metacommands_enabled_); - Dml::SetDefaultRoundingMode(provider.get(), rounding_mode_); return provider; } -void DMLProviderFactory::SetDefaultRoundingMode(AllocatorRoundingMode rounding_mode) { - rounding_mode_ = rounding_mode; -} - void DMLProviderFactory::SetMetacommandsEnabled(bool metacommands_enabled) { metacommands_enabled_ = metacommands_enabled; } @@ -80,11 +73,6 @@ std::shared_ptr CreateExecutionProviderFactory_DML(ID return std::make_shared(dml_device, cmd_queue); } -void DmlConfigureProviderFactoryDefaultRoundingMode(IExecutionProviderFactory* factory, AllocatorRoundingMode rounding_mode) { - auto dml_provider_factory = static_cast(factory); - dml_provider_factory->SetDefaultRoundingMode(rounding_mode); -} - void DmlConfigureProviderFactoryMetacommandsEnabled(IExecutionProviderFactory* factory, bool metacommandsEnabled) { auto dml_provider_factory = static_cast(factory); dml_provider_factory->SetMetacommandsEnabled(metacommandsEnabled); diff --git a/winml/adapter/winml_adapter_apis.h b/winml/adapter/winml_adapter_apis.h index 7d1f7941f9..3bd94c3279 100644 --- a/winml/adapter/winml_adapter_apis.h +++ b/winml/adapter/winml_adapter_apis.h @@ -62,7 +62,6 @@ ORT_API_STATUS(SessionGetIntraOpThreadSpinning, _In_ OrtSession* session, _Out_ ORT_API_STATUS(SessionGetNamedDimensionsOverrides, _In_ OrtSession* session, _Out_ winrt::Windows::Foundation::Collections::IMapView& overrides); // Dml methods (TODO need to figure out how these need to move to session somehow...) -ORT_API_STATUS(DmlExecutionProviderSetDefaultRoundingMode, _In_ OrtExecutionProvider* dml_provider, _In_ bool is_enabled); ORT_API_STATUS(DmlExecutionProviderFlushContext, _In_ OrtExecutionProvider* dml_provider); ORT_API_STATUS(DmlExecutionProviderReleaseCompletedReferences, _In_ OrtExecutionProvider* dml_provider); diff --git a/winml/adapter/winml_adapter_c_api.cpp b/winml/adapter/winml_adapter_c_api.cpp index a3e597fb88..754bcf1c81 100644 --- a/winml/adapter/winml_adapter_c_api.cpp +++ b/winml/adapter/winml_adapter_c_api.cpp @@ -62,7 +62,6 @@ static constexpr WinmlAdapterApi winml_adapter_api_1 = { &winmla::SessionGetNamedDimensionsOverrides, // Dml methods (TODO need to figure out how these need to move to session somehow...) - &winmla::DmlExecutionProviderSetDefaultRoundingMode, &winmla::DmlExecutionProviderFlushContext, &winmla::DmlExecutionProviderReleaseCompletedReferences, &winmla::DmlCopyTensor, diff --git a/winml/adapter/winml_adapter_c_api.h b/winml/adapter/winml_adapter_c_api.h index 2817467818..62a92b86b0 100644 --- a/winml/adapter/winml_adapter_c_api.h +++ b/winml/adapter/winml_adapter_c_api.h @@ -354,14 +354,6 @@ struct WinmlAdapterApi { */ OrtStatus*(ORT_API_CALL* SessionGetNamedDimensionsOverrides)(_In_ OrtSession* session, _Out_ winrt::Windows::Foundation::Collections::IMapView& overrides)NO_EXCEPTION; - /** - * DmlExecutionProviderSetDefaultRoundingMode - * This api is used to configure the DML EP to turn on/off rounding. - * - * WinML uses this to disable rounding during session initialization and then enables it again post initialization. - */ - OrtStatus*(ORT_API_CALL* DmlExecutionProviderSetDefaultRoundingMode)(_In_ OrtExecutionProvider* dml_provider, _In_ bool is_enabled)NO_EXCEPTION; - /** * DmlExecutionProviderFlushContext * This api is used to flush the DML EP. diff --git a/winml/adapter/winml_adapter_dml.cpp b/winml/adapter/winml_adapter_dml.cpp index acda0d332b..3481116d3d 100644 --- a/winml/adapter/winml_adapter_dml.cpp +++ b/winml/adapter/winml_adapter_dml.cpp @@ -68,7 +68,6 @@ Microsoft::WRL::ComPtr CreateDmlDevice(ID3D12Device* d3d12Device) { } namespace onnxruntime { -void DmlConfigureProviderFactoryDefaultRoundingMode(onnxruntime::IExecutionProviderFactory* factory, AllocatorRoundingMode rounding_mode); void DmlConfigureProviderFactoryMetacommandsEnabled(IExecutionProviderFactory* factory, bool metacommandsEnabled); } @@ -83,12 +82,6 @@ ORT_API_STATUS_IMPL(winmla::OrtSessionOptionsAppendExecutionProviderEx_DML, _In_ return status; } auto factory = options->provider_factories.back().get(); - - // OnnxRuntime uses the default rounding mode when calling the session's allocator. - // During initialization, OnnxRuntime allocates weights, which are permanent across session - // lifetime and can be large, so shouldn't be rounded. - // So we create the provider with rounding disabled, and expect the caller to enable it after. - onnxruntime::DmlConfigureProviderFactoryDefaultRoundingMode(factory, AllocatorRoundingMode::Disabled); onnxruntime::DmlConfigureProviderFactoryMetacommandsEnabled(factory, metacommands_enabled); #endif // USE_DML @@ -96,16 +89,6 @@ ORT_API_STATUS_IMPL(winmla::OrtSessionOptionsAppendExecutionProviderEx_DML, _In_ API_IMPL_END } -ORT_API_STATUS_IMPL(winmla::DmlExecutionProviderSetDefaultRoundingMode, _In_ OrtExecutionProvider* dml_provider, _In_ bool is_enabled) { - API_IMPL_BEGIN -#ifdef USE_DML - auto dml_provider_internal = reinterpret_cast<::onnxruntime::IExecutionProvider*>(dml_provider); - Dml::SetDefaultRoundingMode(dml_provider_internal, is_enabled ? AllocatorRoundingMode::Enabled : AllocatorRoundingMode::Disabled); -#endif - return nullptr; - API_IMPL_END -} - ORT_API_STATUS_IMPL(winmla::DmlExecutionProviderFlushContext, _In_ OrtExecutionProvider* dml_provider) { API_IMPL_BEGIN #ifdef USE_DML diff --git a/winml/lib/Api.Ort/OnnxruntimeDmlSessionBuilder.cpp b/winml/lib/Api.Ort/OnnxruntimeDmlSessionBuilder.cpp index a02c6a0431..95de7a4e52 100644 --- a/winml/lib/Api.Ort/OnnxruntimeDmlSessionBuilder.cpp +++ b/winml/lib/Api.Ort/OnnxruntimeDmlSessionBuilder.cpp @@ -95,9 +95,6 @@ HRESULT OnnxruntimeDmlSessionBuilder::Initialize( RETURN_HR_IF_NOT_OK_MSG(winml_adapter_api->SessionGetExecutionProvider(session, 0, &ort_provider), engine_factory_->UseOrtApi()); - RETURN_HR_IF_NOT_OK_MSG(winml_adapter_api->DmlExecutionProviderSetDefaultRoundingMode(ort_provider, true), - engine_factory_->UseOrtApi()); - // Flush the D3D12 work from the DML execution provider RETURN_HR_IF_NOT_OK_MSG(winml_adapter_api->DmlExecutionProviderFlushContext(ort_provider), engine_factory_->UseOrtApi()); diff --git a/winml/test/adapter/AdapterDmlEpTest.cpp b/winml/test/adapter/AdapterDmlEpTest.cpp index 3069c618b1..b5364a8447 100644 --- a/winml/test/adapter/AdapterDmlEpTest.cpp +++ b/winml/test/adapter/AdapterDmlEpTest.cpp @@ -80,14 +80,6 @@ UniqueOrtSession CreateCpuSession() { return CreateUniqueOrtSession(FileHelpers::GetModulePath() + L"fns-candy.onnx", session_options); } -void DmlExecutionProviderSetDefaultRoundingMode() { - GPUTEST; - auto session = CreateDmlSession(); - OrtExecutionProvider* ort_provider; - THROW_IF_NOT_OK_MSG(winml_adapter_api->SessionGetExecutionProvider(session.get(), 0, &ort_provider), ort_api); - THROW_IF_NOT_OK_MSG(winml_adapter_api->DmlExecutionProviderSetDefaultRoundingMode(ort_provider, false), ort_api); -} - void DmlExecutionProviderFlushContext() { GPUTEST; auto session = CreateDmlSession(); @@ -299,7 +291,6 @@ const AdapterDmlEpTestApi& getapi() { { AdapterDmlEpTestSetup, AdapterDmlEpTestTeardown, - DmlExecutionProviderSetDefaultRoundingMode, DmlExecutionProviderFlushContext, DmlExecutionProviderReleaseCompletedReferences, DmlCreateAndFreeGPUAllocationFromD3DResource, diff --git a/winml/test/adapter/AdapterDmlEpTest.h b/winml/test/adapter/AdapterDmlEpTest.h index f08bcdc25a..4743cb0977 100644 --- a/winml/test/adapter/AdapterDmlEpTest.h +++ b/winml/test/adapter/AdapterDmlEpTest.h @@ -6,7 +6,6 @@ struct AdapterDmlEpTestApi { SetupTest AdapterDmlEpTestSetup; TeardownClass AdapterDmlEpTestTeardown; - VoidTest DmlExecutionProviderSetDefaultRoundingMode; VoidTest DmlExecutionProviderFlushContext; VoidTest DmlExecutionProviderReleaseCompletedReferences; VoidTest DmlCreateGPUAllocationFromD3DResource; @@ -24,7 +23,6 @@ WINML_TEST_CLASS_BEGIN(AdapterDmlEpTest) WINML_TEST_CLASS_SETUP_METHOD(AdapterDmlEpTestSetup) WINML_TEST_CLASS_TEARDOWN_METHOD(AdapterDmlEpTestTeardown) WINML_TEST_CLASS_BEGIN_TESTS -WINML_TEST(AdapterDmlEpTest, DmlExecutionProviderSetDefaultRoundingMode) WINML_TEST(AdapterDmlEpTest, DmlExecutionProviderFlushContext) WINML_TEST(AdapterDmlEpTest, DmlExecutionProviderReleaseCompletedReferences) WINML_TEST(AdapterDmlEpTest, DmlCreateGPUAllocationFromD3DResource)