From b28a3a92e6781147ae0af877c36a0471ac1cf2fa Mon Sep 17 00:00:00 2001 From: Ryan Lai Date: Tue, 17 Dec 2019 14:58:46 -0800 Subject: [PATCH] Delete Ort Allocator in LearningModelBinding (#2653) * Delete OrtAllocator in LearningModelBinding * PR comments to make Ort::Allocator a smart pointer * Small comment change * PR feedback to clean up code * PR feedback on move semantics * Clean up std::move --- winml/adapter/WinMLAdapter.h | 58 +++++++++++++- winml/lib/Api/ImageFeatureValue.cpp | 14 +--- winml/lib/Api/ImageFeatureValue.h | 5 +- winml/lib/Api/LearningModelBinding.cpp | 75 ++++++++++++------- winml/lib/Api/LearningModelBinding.h | 10 ++- winml/lib/Api/impl/MapBase.h | 3 +- winml/lib/Api/impl/SequenceBase.h | 4 +- winml/lib/Api/impl/TensorBase.h | 3 +- .../lib/Api/inc/ILotusValueProviderPrivate.h | 2 +- 9 files changed, 125 insertions(+), 49 deletions(-) diff --git a/winml/adapter/WinMLAdapter.h b/winml/adapter/WinMLAdapter.h index 6df268244d..0c5637f5c7 100644 --- a/winml/adapter/WinMLAdapter.h +++ b/winml/adapter/WinMLAdapter.h @@ -146,4 +146,60 @@ private: std::shared_ptr session_; }; -} // namespace Windows::AI::MachineLearning::Adapter \ No newline at end of file +} // namespace Windows::AI::MachineLearning::Adapter + +namespace Ort { +// Ort::Allocator is not in the C ABI yet so it will have to be in the WinMLAdapter for now. +// This struct was copied using the Base struct from onnxruntime_cxx_api.h for reference +// Ort::Allocator struct is used as a smart pointer to OrtAllocator. +struct Allocator { + Allocator() { + m_ort_allocator = nullptr; + m_adapter = nullptr; + } + Allocator(winmla::IWinMLAdapter* adapter, OrtAllocator* ort_allocator) : + m_adapter(adapter), m_ort_allocator(ort_allocator) {} + + ~Allocator() { + if (m_adapter != nullptr && m_ort_allocator != nullptr) { + m_adapter->FreeProviderAllocator(m_ort_allocator); + } + } + + operator OrtAllocator*() { return m_ort_allocator; } + operator const OrtAllocator*() const { return m_ort_allocator; } + + OrtAllocator* release() { + OrtAllocator* p = m_ort_allocator; + m_ort_allocator = nullptr; + m_adapter = nullptr; + return p; + } + + OrtAllocator** put() noexcept { + assert(m_ort_allocator == nullptr); + return &m_ort_allocator; + } + + Allocator(const Allocator&) = delete; + Allocator& operator=(const Allocator&) = delete; + Allocator(Allocator&& v) noexcept : + m_adapter{v.m_adapter}, m_ort_allocator{v.m_ort_allocator} { + v.m_adapter = nullptr; + v.m_ort_allocator = nullptr; + } + void operator=(Allocator&& v) noexcept { + if (m_ort_allocator != nullptr && m_adapter != nullptr) { + m_adapter->FreeProviderAllocator(m_ort_allocator); + } + m_adapter = v.m_adapter; + m_ort_allocator = v.m_ort_allocator; + v.m_adapter = nullptr; + v.m_ort_allocator = nullptr; + } + + private: + winmla::IWinMLAdapter* m_adapter; + OrtAllocator* m_ort_allocator; +}; +} // namespace Ort \ No newline at end of file diff --git a/winml/lib/Api/ImageFeatureValue.cpp b/winml/lib/Api/ImageFeatureValue.cpp index 5cc0e1b4ce..f24d414891 100644 --- a/winml/lib/Api/ImageFeatureValue.cpp +++ b/winml/lib/Api/ImageFeatureValue.cpp @@ -173,12 +173,6 @@ ImageFeatureValue::ImageFeatureValue(IVectorView con Initialize(); } -ImageFeatureValue::~ImageFeatureValue() { - for (auto allocator : m_tensorAllocators) { - m_adapter->FreeProviderAllocator(allocator); - } -} - static std::optional GetBitmapPixelFormatFromMetadata(const IPropertySet& properties) { if (properties != nullptr && properties.HasKey(L"BitmapPixelFormat")) { if (auto pixelFormatInspectable = properties.Lookup(L"BitmapPixelFormat")) { @@ -496,7 +490,7 @@ std::optional ImageFeatureValue::GetIn return ImageResourceMetadata{bounds, imageTensorDescriptor}; } -HRESULT ImageFeatureValue::GetOrtValue(WinML::BindingContext& context, OrtValue** ort_value) try { +HRESULT ImageFeatureValue::GetOrtValue(WinML::BindingContext& context, OrtValue** ort_value, OrtAllocator** ort_allocator) try { FAIL_FAST_IF(!(std::all_of(m_widths.begin(), m_widths.end(), [](int i) { return i != 0; }))); FAIL_FAST_IF(!(std::all_of(m_heights.begin(), m_heights.end(), [](int i) { return i != 0; }))); @@ -516,8 +510,8 @@ HRESULT ImageFeatureValue::GetOrtValue(WinML::BindingContext& context, OrtValue* } // create the OrtValue - OrtAllocator* dml_allocator; - WINML_THROW_IF_FAILED(m_adapter->GetProviderAllocator(provider, &dml_allocator)); + Ort::Allocator dml_allocator(m_adapter.get(), nullptr); + WINML_THROW_IF_FAILED(m_adapter->GetProviderAllocator(provider, dml_allocator.put())); // create the OrtValue as a tensor letting ort know that we own the data buffer Ort::Value ort_tensor = Ort::Value::CreateTensor( @@ -525,7 +519,6 @@ HRESULT ImageFeatureValue::GetOrtValue(WinML::BindingContext& context, OrtValue* &(resourceMetadata.TensorDescriptor.sizes[0]), sizeof(resourceMetadata.TensorDescriptor.sizes) / sizeof(resourceMetadata.TensorDescriptor.sizes[0]), (resourceMetadata.TensorDescriptor.dataType == kImageTensorDataTypeFloat32) ? ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT : ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16); - m_tensorAllocators.emplace_back(dml_allocator); // Get the tensor raw data void* pAllocatedResource = nullptr; @@ -545,6 +538,7 @@ HRESULT ImageFeatureValue::GetOrtValue(WinML::BindingContext& context, OrtValue* } *ort_value = ort_tensor.release(); + *ort_allocator = dml_allocator.release(); return S_OK; } WINML_CATCH_ALL_COM diff --git a/winml/lib/Api/ImageFeatureValue.h b/winml/lib/Api/ImageFeatureValue.h index 31e19b7e5d..d826c12e23 100644 --- a/winml/lib/Api/ImageFeatureValue.h +++ b/winml/lib/Api/ImageFeatureValue.h @@ -14,7 +14,6 @@ struct ImageFeatureValue : ImageFeatureValueT const& images); ImageFeatureValue(winrt::Windows::Foundation::Collections::IVectorView const& images); @@ -34,7 +33,7 @@ struct ImageFeatureValue : ImageFeatureValueT Widths() { return m_widths; } std::vector Heights() { return m_heights; } bool IsBatch() { return m_batchSize > 1; } - private: com_ptr m_adapter; winrt::Windows::Foundation::Collections::IVector m_videoFrames; std::vector m_widths = {}; std::vector m_heights = {}; - std::vector m_tensorAllocators; uint32_t m_batchSize = 1; // Crop the image with desired aspect ratio. // This function does not crop image to desried width and height, but crops to center for desired ratio diff --git a/winml/lib/Api/LearningModelBinding.cpp b/winml/lib/Api/LearningModelBinding.cpp index 540f3dee1b..8076f2fb00 100644 --- a/winml/lib/Api/LearningModelBinding.cpp +++ b/winml/lib/Api/LearningModelBinding.cpp @@ -39,6 +39,10 @@ static Windows::AI::MachineLearning::ILearningModelFeatureDescriptor FindValidBi return nullptr; } +LearningModelBinding::~LearningModelBinding() { + Clear(); +} + using NullableBindingPort = std::optional>; static NullableBindingPort FindValidBinding( @@ -59,7 +63,7 @@ void LearningModelBinding::CacheProvider( m_providers[name] = providerInfo; } -std::tuple LearningModelBinding::CreateBinding( +std::tuple LearningModelBinding::CreateBinding( const std::string& name, const Windows::Foundation::IInspectable& inspectable, Windows::Foundation::Collections::IPropertySet const& properties) { @@ -99,6 +103,7 @@ std::tuple LearningModelBinding::CreateBind // Get the bound tensor Ort::Value value(nullptr); + Ort::Allocator ort_allocator(adapter_.get(), nullptr); // Get the native ORT interface for the given bind value auto spLotusValueProvider = featureValue.as(); @@ -121,7 +126,7 @@ std::tuple LearningModelBinding::CreateBind if (!isPlaceHolder || shouldAlwaysTensorize) { // If not a placeholder, attempt to get the underlying resource WINML_THROW_IF_FAILED_MSG( - spLotusValueProvider->GetOrtValue(context, value.put()), + spLotusValueProvider->GetOrtValue(context, value.put(), ort_allocator.put()), "The model variable %s failed tensorization.", name.c_str()); } else { @@ -136,7 +141,7 @@ std::tuple LearningModelBinding::CreateBind auto providerInfo = ProviderInfo{inspectable, spLotusValueProvider, context}; CacheProvider(name, providerInfo); - return std::make_tuple(name, value.release(), bindingType); + return std::make_tuple(name, value.release(), bindingType, ort_allocator.release()); } void LearningModelBinding::Bind( @@ -155,16 +160,23 @@ void LearningModelBinding::Bind( BindingType bindingType; std::string bindingName; OrtValue* binding_value = nullptr; - + OrtAllocator* ort_allocator = nullptr; auto featureName = WinML::Strings::UTF8FromHString(name); - std::tie(bindingName, binding_value, bindingType) = CreateBinding(featureName, value, properties); + std::tie(bindingName, binding_value, bindingType, ort_allocator) = CreateBinding(featureName, value, properties); Ort::Value ortValue = binding_value ? Ort::Value(binding_value) : Ort::Value(nullptr); + Ort::Allocator ortAllocator(adapter_.get(), ort_allocator); switch (bindingType) { case BindingType::kInput: - WINML_THROW_IF_FAILED(BindInput(bindingName, ortValue)); + WINML_THROW_IF_FAILED(BindInput( + bindingName, + std::move(ortValue), + std::move(ortAllocator))); break; case BindingType::kOutput: - WINML_THROW_IF_FAILED(BindOutput(bindingName, ortValue)); + WINML_THROW_IF_FAILED(BindOutput( + bindingName, + std::move(ortValue), + std::move(ortAllocator))); break; default: FAIL_FAST(); @@ -179,6 +191,8 @@ void LearningModelBinding::Clear() try { outputs_.clear(); output_names_.clear(); m_providers.clear(); + input_allocators_.clear(); + output_allocators_.clear(); } WINML_CATCH_ALL @@ -219,7 +233,7 @@ bool LearningModelBinding::HasKey(hstring const& key) { void LearningModelBinding::Split( Windows::Foundation::Collections::IMapView& first, Windows::Foundation::Collections::IMapView& second) { - // the winrt api guide states: + // the winrt api guide states: // If the IMapView instance cannot be split, then both the first and second parameters are null when the method returns. first = nullptr; second = nullptr; @@ -490,21 +504,28 @@ STDMETHODIMP LearningModelBinding::Bind( BindingType bindingType; std::string bindingName; OrtValue* binding_value_ptr = nullptr; - + OrtAllocator* ort_allocator = nullptr; winrt::Windows::Foundation::IInspectable to; RETURN_IF_FAILED(value->QueryInterface( winrt::guid_of(), reinterpret_cast(winrt::put_abi(to)))); auto featureName = WinML::Strings::UTF8FromUnicode(name, cchName); - std::tie(bindingName, binding_value_ptr, bindingType) = CreateBinding(featureName, to, nullptr); + std::tie(bindingName, binding_value_ptr, bindingType, ort_allocator) = CreateBinding(featureName, to, nullptr); Ort::Value ortValue = binding_value_ptr ? Ort::Value(binding_value_ptr) : Ort::Value(nullptr); + Ort::Allocator ortAllocator(adapter_.get(), ort_allocator); switch (bindingType) { case BindingType::kInput: - WINML_THROW_IF_FAILED(BindInput(bindingName, ortValue)); + WINML_THROW_IF_FAILED(BindInput( + bindingName, + std::move(ortValue), + std::move(ortAllocator))); break; case BindingType::kOutput: - WINML_THROW_IF_FAILED(BindOutput(bindingName, ortValue)); + WINML_THROW_IF_FAILED(BindOutput( + bindingName, + std::move(ortValue), + std::move(ortAllocator))); break; default: FAIL_FAST(); @@ -523,40 +544,43 @@ static std::pair Contains(const std::vector& names, c } // This method releases control of memory of ml_value from caller of BindInput -HRESULT LearningModelBinding::BindInput(const std::string& name, Ort::Value& ml_value) { +HRESULT LearningModelBinding::BindInput(const std::string& name, Ort::Value&& ml_value, Ort::Allocator&& ort_allocator) { auto rc = Contains(input_names_, name); - auto add_or_replace = [this, &name](const bool exists, size_t index, Ort::Value& value) { + auto add_or_replace = [this, &name](const bool exists, size_t index, Ort::Value&& value, Ort::Allocator&& ort_allocator) { if (exists) { - inputs_[index] = Ort::Value(value.release()); + inputs_[index] = std::move(value); + input_allocators_[index] = std::move(ort_allocator); } else { input_names_.push_back(name); - inputs_.push_back(Ort::Value(value.release())); + inputs_.push_back(std::move(value)); + input_allocators_.push_back(std::move(ort_allocator)); } }; if (ml_value.IsTensor()) { - Ort::Value new_mlvalue = Ort::Value(nullptr); + OrtValue* new_mlvalue; WINML_THROW_IF_FAILED(m_session.as() ->GetIInferenceSession() - ->CopyOneInputAcrossDevices(name.c_str(), ml_value, new_mlvalue.put())); - add_or_replace(rc.first, rc.second, new_mlvalue); + ->CopyOneInputAcrossDevices(name.c_str(), ml_value, &new_mlvalue)); + add_or_replace(rc.first, rc.second, Ort::Value(new_mlvalue), std::move(ort_allocator)); } else { - add_or_replace(rc.first, rc.second, ml_value); + add_or_replace(rc.first, rc.second, Ort::Value(ml_value.release()), std::move(ort_allocator)); } return S_OK; } // This method releases control of memory of ml_value from caller of BindInput -HRESULT LearningModelBinding::BindOutput(const std::string& name, Ort::Value& ml_value) { +HRESULT LearningModelBinding::BindOutput(const std::string& name, Ort::Value&& ml_value, Ort::Allocator&& ort_allocator) { auto rc = Contains(output_names_, name); - OrtValue* ml_value_data = ml_value.release(); if (rc.first) { - outputs_[rc.second] = ml_value_data ? Ort::Value(ml_value_data) : Ort::Value(nullptr); + outputs_[rc.second] = std::move(ml_value); + output_allocators_[rc.second] = std::move(ort_allocator); return S_OK; } output_names_.push_back(name); - outputs_.push_back(ml_value_data ? Ort::Value(ml_value_data) : Ort::Value(nullptr)); + outputs_.push_back(std::move(ml_value)); + output_allocators_.push_back(std::move(ort_allocator)); return S_OK; } @@ -610,8 +634,7 @@ void LearningModelBinding::BindUnboundOutputs() { // Add all unbound outputs to binding collection for (const auto& unbound_output : unbound_output_names) { - Ort::Value out(nullptr); - WINML_THROW_IF_FAILED(BindOutput(unbound_output, out)); + WINML_THROW_IF_FAILED(BindOutput(unbound_output, Ort::Value(nullptr), Ort::Allocator())); } } diff --git a/winml/lib/Api/LearningModelBinding.h b/winml/lib/Api/LearningModelBinding.h index d4d08f7e2f..fa45ce8e1d 100644 --- a/winml/lib/Api/LearningModelBinding.h +++ b/winml/lib/Api/LearningModelBinding.h @@ -22,7 +22,7 @@ struct LearningModelBinding : LearningModelBindingT; LearningModelBinding() = delete; - + ~LearningModelBinding(); LearningModelBinding(Windows::AI::MachineLearning::LearningModelSession const& session); void Bind(hstring const& name, Windows::Foundation::IInspectable const& value); @@ -36,7 +36,7 @@ struct LearningModelBinding : LearningModelBindingT& first, Windows::Foundation::Collections::IMapView& second); - std::tuple CreateBinding( + std::tuple CreateBinding( const std::string& name, const Windows::Foundation::IInspectable& value, Windows::Foundation::Collections::IPropertySet const& properties); @@ -55,7 +55,7 @@ struct LearningModelBinding : LearningModelBindingT& LearningModelBinding::GetOutputs(); const std::vector& LearningModelBinding::GetInputNames() const; const std::vector& LearningModelBinding::GetInputs() const; - HRESULT BindOutput(const std::string& name, Ort::Value& ml_value); + HRESULT BindOutput(const std::string& name, Ort::Value&& ml_value, Ort::Allocator&& ort_allocator); void BindUnboundOutputs(); private: @@ -67,7 +67,7 @@ struct LearningModelBinding : LearningModelBindingT adapter_; std::vector input_names_; std::vector inputs_; + std::vector input_allocators_; std::vector output_names_; std::vector outputs_; + std::vector output_allocators_; }; } // namespace winrt::Windows::AI::MachineLearning::implementation diff --git a/winml/lib/Api/impl/MapBase.h b/winml/lib/Api/impl/MapBase.h index 632ef1e677..b3490bbe1a 100644 --- a/winml/lib/Api/impl/MapBase.h +++ b/winml/lib/Api/impl/MapBase.h @@ -158,7 +158,8 @@ struct MapBase : winrt::implements< } STDMETHOD(GetOrtValue) - (WinML::BindingContext& context, OrtValue** ort_value) { + (WinML::BindingContext& context, OrtValue** ort_value, OrtAllocator** ort_allocator) { + ORT_UNUSED_PARAMETER(ort_allocator); ORT_UNUSED_PARAMETER(context); // TODO: Tensorized data should be cached so multiple bindings work more efficiently diff --git a/winml/lib/Api/impl/SequenceBase.h b/winml/lib/Api/impl/SequenceBase.h index a611a1da3b..f3917d74aa 100644 --- a/winml/lib/Api/impl/SequenceBase.h +++ b/winml/lib/Api/impl/SequenceBase.h @@ -195,7 +195,9 @@ struct SequenceBase : public winrt::implements< STDMETHOD(GetOrtValue)( WinML::BindingContext& context, - OrtValue** ort_value) { + OrtValue** ort_value, + OrtAllocator** ort_allocator) { + ORT_UNUSED_PARAMETER(ort_allocator); // TODO: Tensorized data should be cached so multiple bindings work more efficiently // TODO : we need to handle inputs. for now only handle outputs and don't pre allocate anything diff --git a/winml/lib/Api/impl/TensorBase.h b/winml/lib/Api/impl/TensorBase.h index 40ef0e625d..c706ecbd4f 100644 --- a/winml/lib/Api/impl/TensorBase.h +++ b/winml/lib/Api/impl/TensorBase.h @@ -208,7 +208,8 @@ struct TensorBase : TBase { // ILotusValueProviderPrivate::GetOrtValue STDMETHOD(GetOrtValue) - (WinML::BindingContext& context, OrtValue** ort_value) { + (WinML::BindingContext& context, OrtValue** ort_value, OrtAllocator** ort_allocator) { + ORT_UNUSED_PARAMETER(ort_allocator); RETURN_HR_IF_NULL_MSG( WINML_ERR_INVALID_BINDING, m_resources, diff --git a/winml/lib/Api/inc/ILotusValueProviderPrivate.h b/winml/lib/Api/inc/ILotusValueProviderPrivate.h index b38920bfb8..5ae5adc902 100644 --- a/winml/lib/Api/inc/ILotusValueProviderPrivate.h +++ b/winml/lib/Api/inc/ILotusValueProviderPrivate.h @@ -24,7 +24,7 @@ struct BindingContext { }; struct __declspec(uuid("27e2f437-0112-4693-849e-e04323a620fb")) __declspec(novtable) ILotusValueProviderPrivate : IUnknown { - virtual HRESULT __stdcall GetOrtValue(BindingContext& binding_context, OrtValue ** ort_value) = 0; + virtual HRESULT __stdcall GetOrtValue(BindingContext& binding_context, OrtValue** ort_value, OrtAllocator** ort_allocator) = 0; virtual HRESULT __stdcall IsPlaceholder(bool* is_placeholder) = 0; virtual HRESULT __stdcall UpdateSourceResourceData(BindingContext& binding_context, OrtValue* ort_value) = 0; virtual HRESULT __stdcall AbiRepresentation(winrt::Windows::Foundation::IInspectable& abi_representation) = 0;