diff --git a/cmake/onnxruntime_unittests.cmake b/cmake/onnxruntime_unittests.cmake index 945a04d774..19f0f1addd 100644 --- a/cmake/onnxruntime_unittests.cmake +++ b/cmake/onnxruntime_unittests.cmake @@ -1255,6 +1255,9 @@ if (NOT onnxruntime_ENABLE_TRAINING_TORCH_INTEROP) if (onnxruntime_USE_TENSORRT) list(APPEND onnxruntime_shared_lib_test_LIBS ${TENSORRT_LIBRARY_INFER}) endif() + if (onnxruntime_USE_DML) + list(APPEND onnxruntime_shared_lib_test_LIBS d3d12.lib) + endif() if (CMAKE_SYSTEM_NAME STREQUAL "Android") list(APPEND onnxruntime_shared_lib_test_LIBS ${android_shared_libs}) endif() diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h index e376167d34..8de2a29fe4 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/DmlExecutionProvider.h @@ -33,7 +33,7 @@ namespace Dml IDMLDevice* dmlDevice, Dml::ExecutionContext* execution_context, bool enableMetacommands, - bool enableDynamicGraphFusion, + bool enableGraphCapture, bool enableCpuSyncSpinning); ID3D12Resource* GetD3D12ResourceFromAllocation(onnxruntime::IAllocator* allocator, void* ptr); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.cpp index 27168bc8e9..90942f4eef 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.cpp @@ -3,6 +3,8 @@ #include "DmlGraphFusionHelper.h" #include "DmlRuntimeFusedGraphKernel.h" +using namespace Windows::AI::MachineLearning::Adapter; + namespace Dml { namespace DmlGraphFusionHelper @@ -94,13 +96,13 @@ namespace DmlGraphFusionHelper ID3D12Resource** resource, uint64_t* allocId) { - IUnknown* allocationUnk = static_cast(const_cast(tensor->DataRaw())); - Microsoft::WRL::ComPtr resourceUnk; - winmlProvider->GetABIDataInterface(false, allocationUnk, &resourceUnk); + IUnknown* allocationUnknown = static_cast(const_cast(tensor->DataRaw())); + Microsoft::WRL::ComPtr resourceUnknown; + winmlProvider->GetABIDataInterface(false, allocationUnknown, &resourceUnknown); - *allocId = winmlProvider->TryGetPooledAllocationId(allocationUnk, 0); + *allocId = winmlProvider->TryGetPooledAllocationId(allocationUnknown, 0); - ORT_THROW_IF_FAILED(resourceUnk->QueryInterface(resource)); + ORT_THROW_IF_FAILED(resourceUnknown->QueryInterface(resource)); } std::tuple, std::vector, std::byte*, size_t> UnpackInitializer( @@ -363,7 +365,7 @@ namespace DmlGraphFusionHelper oldNodeIndexToNewNodeIndexMap[index] = static_cast(dmlGraphNodes.size()); auto& constantData = std::get(constantNodeVariant); - + DML_CONSTANT_DATA_GRAPH_NODE_DESC* constantNode = allocator.template Allocate(); constantNode->Name = node.Name.data(); constantNode->DataSize = constantData.dataSize; @@ -598,16 +600,16 @@ namespace DmlGraphFusionHelper { if (graphSerializationEnabled) { - + const std::wstring modelName = GetModelName(graph.ModelPath()); auto buffer = SerializeDmlGraph(graphDesc); - + const std::wstring partitionName = L"Partition_" + std::to_wstring(partitionIndex) + L".bin"; WriteToFile(modelName, partitionName, buffer.data(), buffer.size()); - + std::vector> rawData; DmlSerializedGraphDesc deserializedGraphDesc = DeserializeDmlGraph(buffer.data(), rawData); GraphDescBuilder::GraphDesc deserializedDmlGraphDesc = {}; @@ -619,7 +621,7 @@ namespace DmlGraphFusionHelper deserializedDmlGraphDesc.OutputEdges = std::move(deserializedGraphDesc.OutputEdges); deserializedDmlGraphDesc.reuseCommandList = graphDesc.reuseCommandList; deserializedDmlGraphDesc.outputShapes = graphDesc.outputShapes; - + compiledExecutionPlanOperator = DmlGraphFusionHelper::TryCreateCompiledOperator( deserializedDmlGraphDesc, indexedSubGraph, @@ -858,5 +860,212 @@ namespace DmlGraphFusionHelper graph.FinalizeFuseSubGraph(*indexedSubGraph, fusedNode); } + + std::unique_ptr BuildReusableCommandList( + IExecutionProvider* provider, + IDMLCompiledOperator* compiledExecutionPlanOperator, + ID3D12Resource* persistentResource, + std::optional persistentResourceBinding) + { + auto commandListState = std::make_unique(); + + ComPtr device; + ORT_THROW_IF_FAILED(provider->GetDmlDevice(device.GetAddressOf())); + + DML_BINDING_PROPERTIES execBindingProps = compiledExecutionPlanOperator->GetBindingProperties(); + + D3D12_DESCRIPTOR_HEAP_DESC desc = {}; + desc.Flags = D3D12_DESCRIPTOR_HEAP_FLAG_SHADER_VISIBLE; + desc.NumDescriptors = execBindingProps.RequiredDescriptorCount; + desc.Type = D3D12_DESCRIPTOR_HEAP_TYPE_CBV_SRV_UAV; + + ComPtr d3dDevice; + ORT_THROW_IF_FAILED(provider->GetD3DDevice(d3dDevice.GetAddressOf())); + + ORT_THROW_IF_FAILED(d3dDevice->CreateDescriptorHeap(&desc, IID_GRAPHICS_PPV_ARGS(commandListState->heap.ReleaseAndGetAddressOf()))); + + // Create a binding table for execution. + DML_BINDING_TABLE_DESC bindingTableDesc = {}; + bindingTableDesc.Dispatchable = compiledExecutionPlanOperator; + bindingTableDesc.CPUDescriptorHandle = commandListState->heap->GetCPUDescriptorHandleForHeapStart(); + bindingTableDesc.GPUDescriptorHandle = commandListState->heap->GetGPUDescriptorHandleForHeapStart(); + bindingTableDesc.SizeInDescriptors = execBindingProps.RequiredDescriptorCount; + + ORT_THROW_IF_FAILED(device->CreateBindingTable(&bindingTableDesc, IID_PPV_ARGS(&commandListState->bindingTable))); + + ORT_THROW_IF_FAILED(d3dDevice->CreateCommandAllocator( + provider->GetCommandListTypeForQueue(), + IID_GRAPHICS_PPV_ARGS(commandListState->commandAllocator.ReleaseAndGetAddressOf()))); + + ORT_THROW_IF_FAILED(d3dDevice->CreateCommandList( + 0, + provider->GetCommandListTypeForQueue(), + commandListState->commandAllocator.Get(), + nullptr, + IID_GRAPHICS_PPV_ARGS(commandListState->graphicsCommandList.ReleaseAndGetAddressOf()))); + + if (persistentResource) + { + DML_BINDING_DESC persistentResourceBindingDesc = + { DML_BINDING_TYPE_BUFFER, persistentResourceBinding ? &*persistentResourceBinding : nullptr }; + commandListState->bindingTable->BindPersistentResource(&persistentResourceBindingDesc); + } + + ID3D12DescriptorHeap* descriptorHeaps[] = { commandListState->heap.Get() }; + commandListState->graphicsCommandList->SetDescriptorHeaps(ARRAYSIZE(descriptorHeaps), descriptorHeaps); + + ComPtr recorder; + ORT_THROW_IF_FAILED(device->CreateCommandRecorder(IID_PPV_ARGS(recorder.GetAddressOf()))); + + recorder->RecordDispatch(commandListState->graphicsCommandList.Get(), compiledExecutionPlanOperator, commandListState->bindingTable.Get()); + + ORT_THROW_IF_FAILED(commandListState->graphicsCommandList->Close()); + + return commandListState; + } + + void ExecuteReusableCommandList( + onnxruntime::OpKernelContext* kernelContext, + DmlReusedCommandListState& commandListState, + IDMLCompiledOperator* compiledExecutionPlanOperator, + const onnxruntime::OpKernelInfo& kernelInfo, + gsl::span isInputsUploadedByDmlEP, + const std::vector& inputsUsed, + gsl::span> nonOwnedGraphInputsFromInitializers, + const Windows::AI::MachineLearning::Adapter::EdgeShapes& outputShapes, + IWinmlExecutionProvider* winmlProvider, + IExecutionProvider* provider, + IUnknown* persistentResourceAllocatorUnknown) + { + DML_BINDING_PROPERTIES execBindingProps = compiledExecutionPlanOperator->GetBindingProperties(); + + std::vector inputBindings(kernelContext->InputCount()); + std::vector inputBindingDescs(kernelContext->InputCount()); + + OpKernelContextWrapper contextWrapper( + kernelContext, + kernelInfo.GetExecutionProvider(), + true, + nullptr); + + // Populate input bindings, excluding those which were specified as owned by DML and provided + // at initialization instead. + commandListState.inputBindingAllocIds.resize(inputBindings.size()); + bool inputBindingsChanged = false; + + for (uint32_t i = 0; i < inputBindings.size(); ++i) + { + if (!isInputsUploadedByDmlEP[i] && inputsUsed[i]) + { + if (nonOwnedGraphInputsFromInitializers[i]) + { + inputBindings[i].Buffer = nonOwnedGraphInputsFromInitializers[i].Get(); + inputBindings[i].SizeInBytes = nonOwnedGraphInputsFromInitializers[i]->GetDesc().Width; + inputBindingDescs[i] = {DML_BINDING_TYPE_BUFFER, &inputBindings[i]}; + } + else + { + assert(kernelContext->InputType(gsl::narrow_cast(i))->IsTensorType()); + const onnxruntime::Tensor* tensor = kernelContext->Input(gsl::narrow_cast(i)); + + uint64_t allocId; + DmlGraphFusionHelper::UnwrapTensor(winmlProvider, tensor, &inputBindings[i].Buffer, &allocId); + inputBindingsChanged = inputBindingsChanged || (!allocId || commandListState.inputBindingAllocIds[i] != allocId); + inputBindings[i].Buffer->Release(); // Avoid holding an additional reference + inputBindings[i].SizeInBytes = DmlGraphFusionHelper::AlignToPow2(tensor->SizeInBytes(), 4); + inputBindingDescs[i] = {DML_BINDING_TYPE_BUFFER, &inputBindings[i]}; + commandListState.inputBindingAllocIds[i] = allocId; + } + } + } + + if (inputBindingsChanged) + { + commandListState.bindingTable->BindInputs(gsl::narrow_cast(inputBindingDescs.size()), inputBindingDescs.data()); + } + + // Populate Output bindings + std::vector outputBindings(kernelContext->OutputCount()); + std::vector outputBindingDescs(kernelContext->OutputCount()); + + commandListState.outputBindingAllocIds.resize(outputBindings.size()); + bool outputBindingsChanged = false; + + for (uint32_t i = 0; i < outputBindings.size(); ++i) + { + std::vector outputDims; + outputDims.reserve(outputShapes.GetShape(i).size()); + for (uint32_t dimSize : outputShapes.GetShape(i)) + { + outputDims.push_back(dimSize); + } + + onnxruntime::Tensor* tensor = kernelContext->Output( + static_cast(i), + onnxruntime::TensorShape::FromExistingBuffer(outputDims) + ); + + uint64_t allocId; + DmlGraphFusionHelper::UnwrapTensor(winmlProvider, tensor, &outputBindings[i].Buffer, &allocId); + outputBindingsChanged = outputBindingsChanged || (!allocId || commandListState.outputBindingAllocIds[i] != allocId); + outputBindings[i].Buffer->Release(); // Avoid holding an additional reference + outputBindings[i].SizeInBytes = DmlGraphFusionHelper::AlignToPow2(tensor->SizeInBytes(), 4); + outputBindingDescs[i] = {DML_BINDING_TYPE_BUFFER, &outputBindings[i]}; + commandListState.outputBindingAllocIds[i] = allocId; + } + + if (outputBindingsChanged) + { + commandListState.bindingTable->BindOutputs(gsl::narrow_cast(outputBindingDescs.size()), outputBindingDescs.data()); + } + + if (execBindingProps.TemporaryResourceSize > 0) + { + // Allocate temporary data which will automatically be freed when the GPU work + // which is scheduled up to the point that this method returns has completed. + ComPtr tempAlloc; + uint64_t tempAllocId = 0; + ORT_THROW_IF_FAILED(contextWrapper.AllocateTemporaryData(static_cast(execBindingProps.TemporaryResourceSize), tempAlloc.GetAddressOf(), &tempAllocId)); + + ComPtr tempResourceUnknown; + winmlProvider->GetABIDataInterface(false, tempAlloc.Get(), &tempResourceUnknown); + + // Bind the temporary resource. + ComPtr tempResource; + ORT_THROW_IF_FAILED(tempResourceUnknown->QueryInterface(tempResource.GetAddressOf())); + DML_BUFFER_BINDING tempBufferBinding = {tempResource.Get(), 0, execBindingProps.TemporaryResourceSize}; + DML_BINDING_DESC tempBindingDesc = { DML_BINDING_TYPE_BUFFER, &tempBufferBinding }; + + if (!tempAllocId || commandListState.tempBindingAllocId != tempAllocId) + { + commandListState.bindingTable->BindTemporaryResource(&tempBindingDesc); + } + + commandListState.tempBindingAllocId = tempAllocId; + } + + // Execute the command list and if it succeeds, update the fence value at which this command may be + // re-used. + ComPtr fence; + uint64_t completionValue; + HRESULT hr = provider->ExecuteCommandList(commandListState.graphicsCommandList.Get(), fence.GetAddressOf(), &completionValue); + + if (hr == DXGI_ERROR_DEVICE_REMOVED) + { + ComPtr device; + ORT_THROW_IF_FAILED(provider->GetD3DDevice(&device)); + ORT_THROW_IF_FAILED(device->GetDeviceRemovedReason()); + } + + ORT_THROW_IF_FAILED(hr); + commandListState.fence = fence; + commandListState.completionValue = completionValue; + + // Queue references to objects which must be kept alive until resulting GPU work completes + winmlProvider->QueueReference(WRAP_GRAPHICS_UNKNOWN(commandListState.graphicsCommandList).Get()); + winmlProvider->QueueReference(WRAP_GRAPHICS_UNKNOWN(commandListState.heap).Get()); + winmlProvider->QueueReference(commandListState.bindingTable.Get()); + winmlProvider->QueueReference(persistentResourceAllocatorUnknown); + } } } diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.h index f1e9654021..a0e2456334 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.h @@ -6,7 +6,9 @@ #include "GraphPartitioner.h" #include "FusedGraphKernel.h" #include "MLOperatorAuthorImpl.h" +#include "DmlReusedCommandListState.h" +using Windows::AI::MachineLearning::Adapter::IWinmlExecutionProvider; namespace Dml { @@ -100,5 +102,24 @@ namespace DmlGraphFusionHelper const std::unordered_set& dynamicCpuInputMap, std::shared_ptr indexedSubGraph, std::unordered_map>&& isInitializerTransferable); + + std::unique_ptr BuildReusableCommandList( + IExecutionProvider* provider, + IDMLCompiledOperator* compiledExecutionPlanOperator, + ID3D12Resource* persistentResource, + std::optional persistentResourceBinding); + + void ExecuteReusableCommandList( + onnxruntime::OpKernelContext* kernelContext, + DmlReusedCommandListState& commandListState, + IDMLCompiledOperator* compiledExecutionPlanOperator, + const onnxruntime::OpKernelInfo& kernelInfo, + gsl::span isInputsUploadedByDmlEP, + const std::vector& inputsUsed, + gsl::span> nonOwnedGraphInputsFromInitializers, + const Windows::AI::MachineLearning::Adapter::EdgeShapes& outputShapes, + IWinmlExecutionProvider* winmlProvider, + IExecutionProvider* provider, + IUnknown* persistentResourceAllocatorUnknown); } } diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlReusedCommandListState.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlReusedCommandListState.h new file mode 100644 index 0000000000..5c46021a21 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlReusedCommandListState.h @@ -0,0 +1,30 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include + +namespace Dml +{ + struct DmlReusedCommandListState + { + // Re-usable command list, supporting descriptor heap, and DML binding table to update that heap. + Microsoft::WRL::ComPtr graphicsCommandList; + Microsoft::WRL::ComPtr commandAllocator; + Microsoft::WRL::ComPtr heap; + Microsoft::WRL::ComPtr bindingTable; + Microsoft::WRL::ComPtr persistentResource; + Microsoft::WRL::ComPtr persistentResourceAllocatorUnknown; + + // Bindings from previous executions of a re-used command list + mutable std::vector inputBindingAllocIds; + mutable std::vector outputBindingAllocIds; + mutable uint64_t tempBindingAllocId = 0; + + // Fence tracking the status of the command list's last execution, and whether its descriptor heap + // can safely be updated. + mutable Microsoft::WRL::ComPtr fence; + mutable uint64_t completionValue = 0; + }; +} diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlRuntimeFusedGraphKernel.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlRuntimeFusedGraphKernel.cpp index 0f0d445a95..f2c5ae8418 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlRuntimeFusedGraphKernel.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlRuntimeFusedGraphKernel.cpp @@ -6,6 +6,7 @@ #include "core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.h" #include "core/providers/dml/DmlExecutionProvider/src/DmlRuntimeFusedGraphKernel.h" #include "core/providers/dml/DmlExecutionProvider/src/DmlGraphFusionHelper.h" +#include "core/providers/dml/DmlExecutionProvider/src/DmlReusedCommandListState.h" using namespace Windows::AI::MachineLearning::Adapter; @@ -74,7 +75,7 @@ namespace Dml static_cast(persistentResourceSize), AllocatorRoundingMode::Disabled, m_persistentResource.ReleaseAndGetAddressOf(), - m_persistentResourceAllocatorUnk.ReleaseAndGetAddressOf())); + m_persistentResourceAllocatorUnknown.ReleaseAndGetAddressOf())); m_persistentResourceBinding = DML_BUFFER_BINDING { m_persistentResource.Get(), 0, persistentResourceSize }; } @@ -209,7 +210,7 @@ namespace Dml m_outputShapes = graphDesc.outputShapes; // Walk through each graph edge and mark used inputs - m_inputsUsed.resize(fusedNodeInputCount, false); + m_inputsUsed = std::vector(fusedNodeInputCount); for (auto it = serializedGraphInputIndexToSubgraphInputIndex.begin(); it != serializedGraphInputIndexToSubgraphInputIndex.end(); it++) { m_inputsUsed[it->second] = true; } @@ -217,6 +218,10 @@ namespace Dml m_inputsUsed[it->second] = true; } + m_isInputsUploadedByDmlEP.resize(fusedNodeInputCount, 0); + m_nonOwnedGraphInputsFromInitializers.resize(fusedNodeInputCount); + graphDesc.reuseCommandList = true; + // Compile the operator m_compiledExecutionPlanOperator = DmlGraphFusionHelper::TryCreateCompiledOperator( graphDesc, @@ -232,108 +237,77 @@ namespace Dml Info(), initializeResourceRefs, initInputBindings); + + std::vector inputBindings(kernelContext->InputCount()); + std::vector inputBindingDescs(kernelContext->InputCount()); + + m_reusedCommandLists.clear(); } - // Wrap tensors as required by Dml::IExecutionProvider::ExecuteOperator - OpKernelContextWrapper contextWrapper( - kernelContext, - Info().GetExecutionProvider(), - true, - nullptr); + auto providerImpl = static_cast(m_provider.Get()); - ORT_THROW_IF_FAILED(m_provider->AddUAVBarrier()); - - // Get input resources for execution, excluding those which were specified as owned by DML and provided - // at initialization instead. - std::vector> inputTensors(kernelContext->InputCount()); - std::vector inputPtrs(kernelContext->InputCount()); - - for (int i = 0; i < kernelContext->InputCount(); ++i) + // When we are capturing a graph, we don't pool the command list and instead transfer it to the execution provider. Captured graph + // have the same bindings for their entire lifetime. + if (providerImpl->GraphCaptureEnabled() && providerImpl->GetCurrentGraphAnnotationId() != -1 && !providerImpl->GraphCaptured(providerImpl->GetCurrentGraphAnnotationId())) { - if (!m_inputsUsed[i]) + auto reusableCommandList = DmlGraphFusionHelper::BuildReusableCommandList( + m_provider.Get(), + m_compiledExecutionPlanOperator.Get(), + m_persistentResource.Get(), + m_persistentResourceBinding); + + reusableCommandList->persistentResource = m_persistentResource; + reusableCommandList->persistentResourceAllocatorUnknown = m_persistentResourceAllocatorUnknown; + + DmlGraphFusionHelper::ExecuteReusableCommandList( + kernelContext, + *reusableCommandList, + m_compiledExecutionPlanOperator.Get(), + Info(), + m_isInputsUploadedByDmlEP, + m_inputsUsed, + m_nonOwnedGraphInputsFromInitializers, + m_outputShapes, + m_winmlProvider.Get(), + m_provider.Get(), + m_persistentResourceAllocatorUnknown.Get()); + + providerImpl->AppendCapturedGraph(providerImpl->GetCurrentGraphAnnotationId(), std::move(reusableCommandList)); + } + else + { + if (m_reusedCommandLists.empty() || + m_reusedCommandLists.front()->fence && m_reusedCommandLists.front()->fence->GetCompletedValue() < m_reusedCommandLists.front()->completionValue) { - continue; + auto reusableCommandList = DmlGraphFusionHelper::BuildReusableCommandList( + m_provider.Get(), + m_compiledExecutionPlanOperator.Get(), + m_persistentResource.Get(), + m_persistentResourceBinding); + + m_reusedCommandLists.push_front(std::move(reusableCommandList)); } - ORT_THROW_IF_FAILED(contextWrapper.GetInputTensor(i, inputTensors[i].GetAddressOf())); - inputPtrs[i] = m_provider->DecodeResource(MLOperatorTensor(inputTensors[i].Get()).GetDataInterface().Get()); + DmlGraphFusionHelper::ExecuteReusableCommandList( + kernelContext, + *m_reusedCommandLists.front(), + m_compiledExecutionPlanOperator.Get(), + Info(), + m_isInputsUploadedByDmlEP, + m_inputsUsed, + m_nonOwnedGraphInputsFromInitializers, + m_outputShapes, + m_winmlProvider.Get(), + m_provider.Get(), + m_persistentResourceAllocatorUnknown.Get()); + + m_reusedCommandLists.push_back(std::move(m_reusedCommandLists.front())); + m_reusedCommandLists.pop_front(); } - auto outputTensors = contextWrapper.GetOutputTensors(m_outputShapes); - ExecuteOperator( - m_compiledExecutionPlanOperator.Get(), - m_persistentResourceBinding ? &*m_persistentResourceBinding : nullptr, - inputPtrs, - outputTensors); - - ORT_THROW_IF_FAILED(m_provider->AddUAVBarrier()); - return onnxruntime::Status::OK(); } - void ExecuteOperator( - IDMLCompiledOperator* op, - _In_opt_ const DML_BUFFER_BINDING* persistentResourceBinding, - gsl::span inputTensors, - gsl::span outputTensors) const - { - auto FillBindingsFromTensors = [this](auto& bufferBindings, auto& bindingDescs, gsl::span& tensors) - { - for (IMLOperatorTensor* tensor : tensors) - { - if (tensor) - { - assert(tensor->IsDataInterface()); - ID3D12Resource* resource = m_provider->DecodeResource(MLOperatorTensor(tensor).GetDataInterface().Get()); - D3D12_RESOURCE_DESC resourceDesc = resource->GetDesc(); - bufferBindings.push_back({ resource, 0, resourceDesc.Width }); - bindingDescs.push_back({ DML_BINDING_TYPE_BUFFER, &bufferBindings.back() }); - } - else - { - bufferBindings.push_back({ nullptr, 0, 0 }); - bindingDescs.push_back({ DML_BINDING_TYPE_NONE, nullptr }); - } - } - }; - - auto FillBindingsFromBuffers = [](auto& bufferBindings, auto& bindingDescs, gsl::span& resources) - { - for (ID3D12Resource* resource : resources) - { - if (resource) - { - D3D12_RESOURCE_DESC resourceDesc = resource->GetDesc(); - bufferBindings.push_back({ resource, 0, resourceDesc.Width }); - bindingDescs.push_back({ DML_BINDING_TYPE_BUFFER, &bufferBindings.back() }); - } - else - { - bufferBindings.push_back({ nullptr, 0, 0 }); - bindingDescs.push_back({ DML_BINDING_TYPE_NONE, nullptr }); - } - } - }; - - std::vector inputBufferBindings; - inputBufferBindings.reserve(inputTensors.size()); - std::vector inputBindings; - inputBindings.reserve(inputTensors.size()); - FillBindingsFromBuffers(inputBufferBindings, inputBindings, inputTensors); - - std::vector outputBufferBindings; - outputBufferBindings.reserve(outputTensors.size()); - std::vector outputBindings; - outputBindings.reserve(outputTensors.size()); - FillBindingsFromTensors(outputBufferBindings, outputBindings, outputTensors); - - ORT_THROW_IF_FAILED(m_provider->ExecuteOperator( - op, - persistentResourceBinding, - inputBindings, - outputBindings)); - } - private: ComPtr m_winmlProvider; ComPtr m_provider; @@ -356,9 +330,12 @@ namespace Dml mutable ComPtr m_compiledExecutionPlanOperator; mutable std::vector m_inputsUsed; mutable ComPtr m_persistentResource; - mutable ComPtr m_persistentResourceAllocatorUnk; // Controls when the persistent resource is returned to the allocator + mutable ComPtr m_persistentResourceAllocatorUnknown; // Controls when the persistent resource is returned to the allocator mutable Windows::AI::MachineLearning::Adapter::EdgeShapes m_outputShapes; mutable std::unordered_map m_inferredInputShapes; + mutable std::deque> m_reusedCommandLists; + mutable std::vector m_isInputsUploadedByDmlEP; + mutable std::vector> m_nonOwnedGraphInputsFromInitializers; }; onnxruntime::OpKernel* CreateRuntimeFusedGraphKernel( diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp index b57013c008..949cd216d5 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.cpp @@ -71,7 +71,7 @@ namespace Dml IDMLDevice* dmlDevice, Dml::ExecutionContext* executionContext, bool enableMetacommands, - bool enableDynamicGraphFusion, + bool enableGraphCapture, bool enableSyncSpinning) : IExecutionProvider(onnxruntime::kDmlExecutionProvider, OrtDevice(OrtDevice::GPU, OrtDevice::MemType::DEFAULT, 0)) { @@ -85,7 +85,7 @@ namespace Dml ComPtr device; GRAPHICS_THROW_IF_FAILED(dmlDevice->GetParentDevice(IID_GRAPHICS_PPV_ARGS(device.GetAddressOf()))); - m_impl = wil::MakeOrThrow(dmlDevice, device.Get(), executionContext, enableMetacommands, enableDynamicGraphFusion, enableSyncSpinning); + m_impl = wil::MakeOrThrow(dmlDevice, device.Get(), executionContext, enableMetacommands, enableGraphCapture, enableSyncSpinning); } std::vector> @@ -102,6 +102,9 @@ namespace Dml void ExecutionProviderImpl::Close() { + // Release the cached command list references before closing the context + m_capturedGraphs.clear(); + m_context->Close(); } @@ -146,11 +149,11 @@ namespace Dml } } - ExecutionProviderImpl::ExecutionProviderImpl(IDMLDevice* dmlDevice, ID3D12Device* d3d12Device, ExecutionContext* executionContext, bool enableMetacommands, bool enableDynamicGraphFusion, bool enableCpuSyncSpinning) + ExecutionProviderImpl::ExecutionProviderImpl(IDMLDevice* dmlDevice, ID3D12Device* d3d12Device, ExecutionContext* executionContext, bool enableMetacommands, bool enableGraphCapture, bool enableCpuSyncSpinning) : m_d3d12Device(d3d12Device), m_dmlDevice(dmlDevice), m_areMetacommandsEnabled(enableMetacommands), - m_dynamicGraphFusionEnabled(enableDynamicGraphFusion), + m_graphCaptureEnabled(enableGraphCapture), m_cpuSyncSpinningEnabled(enableCpuSyncSpinning), m_context(executionContext) { @@ -1142,11 +1145,16 @@ namespace Dml return m_cpuSyncSpinningEnabled; } - bool ExecutionProviderImpl::DynamicGraphFusionEnabled() const noexcept + bool ExecutionProviderImpl::GraphCaptureEnabled() const noexcept { - return m_dynamicGraphFusionEnabled; + return m_graphCaptureEnabled; } + bool ExecutionProviderImpl::GraphCaptured(int graph_annotation_id) const + { + return m_graphCapturingDone.find(graph_annotation_id) != m_graphCapturingDone.end(); + }; + std::shared_ptr ExecutionProviderImpl::GetInternalRegistrationInfoMap() const { @@ -1181,14 +1189,62 @@ namespace Dml return onnxruntime::common::Status::OK(); } + void ExecutionProviderImpl::AppendCapturedGraph(int annotationId, std::unique_ptr capturedGraph) + { + m_capturedGraphs[annotationId].push_back(std::move(capturedGraph)); + } + + onnxruntime::common::Status ExecutionProviderImpl::ReplayGraph(int graph_annotation_id) + { + for (auto& capturedGraph : m_capturedGraphs[graph_annotation_id]) + { + ExecuteCommandList(capturedGraph->graphicsCommandList.Get(), &capturedGraph->fence, &capturedGraph->completionValue); + } + + return onnxruntime::common::Status::OK(); + } + + Status ExecutionProviderImpl::OnRunStart(const onnxruntime::RunOptions& run_options) + { + if (GraphCaptureEnabled()) + { + auto graphAnnotationStr = run_options.config_options.GetConfigEntry(kOrtRunOptionsConfigCudaGraphAnnotation); + // If graph annotation is not provided, fall back to the one dml graph per session behavior + int dmlGraphAnnotationId = 0; + if (graphAnnotationStr.has_value()) + { + ORT_ENFORCE(onnxruntime::TryParseStringWithClassicLocale(*graphAnnotationStr, dmlGraphAnnotationId), + "Failed to parse the dml graph annotation id: ", + *graphAnnotationStr); + } + + m_currentGraphAnnotationId = dmlGraphAnnotationId; + } + + return onnxruntime::common::Status::OK(); + } + + Status ExecutionProviderImpl::OnRunEnd() + { + if (GraphCaptureEnabled() && m_currentGraphAnnotationId != -1) + { + m_graphCapturingDone.insert(m_currentGraphAnnotationId); + } + + // Flush any pending work to the GPU, but don't block for completion, permitting it + // to overlap other work. + Flush(); + return onnxruntime::common::Status::OK(); + } + std::unique_ptr CreateExecutionProvider( IDMLDevice* dmlDevice, Dml::ExecutionContext* executionContext, bool enableMetacommands, - bool enableDynamicGraphFusion, + bool enableGraphCapture, bool enableCpuSyncSpinning) { - return std::make_unique(dmlDevice, executionContext, enableMetacommands, enableDynamicGraphFusion, enableCpuSyncSpinning); + return std::make_unique(dmlDevice, executionContext, enableMetacommands, enableGraphCapture, enableCpuSyncSpinning); } ID3D12Resource* GetD3D12ResourceFromAllocation(onnxruntime::IAllocator* allocator, void* ptr) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h index cdf7487b2f..70bc99184b 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ExecutionProvider.h @@ -6,6 +6,7 @@ #include "GraphTransformer.h" #include "core/providers/dml/DmlExecutionProvider/inc/IWinmlExecutionProvider.h" #include "core/providers/dml/DmlExecutionProvider/src/IExecutionProvider.h" +#include "core/providers/dml/DmlExecutionProvider/src/DmlReusedCommandListState.h" #include #include @@ -36,7 +37,7 @@ namespace Dml ID3D12Device* d3d12Device, Dml::ExecutionContext* executionContext, bool enableMetacommands, - bool enableDynamicGraphFusion, + bool enableGraphCapture, bool enableCpuSyncSpinning); void ReleaseCompletedReferences(); @@ -154,8 +155,14 @@ namespace Dml STDMETHOD_(bool, CustomHeapsSupported)() const noexcept final; STDMETHOD_(bool, MetacommandsEnabled)() const noexcept final; + bool GraphCaptureEnabled() const noexcept; + bool GraphCaptured(int graph_annotation_id) const; + Status ReplayGraph(int graph_annotation_id); + Status OnRunStart(const onnxruntime::RunOptions& run_options); + Status OnRunEnd(); + int GetCurrentGraphAnnotationId() const { return m_currentGraphAnnotationId; } + void AppendCapturedGraph(int annotationId, std::unique_ptr capturedGraph); bool CpuSyncSpinningEnabled() const noexcept; - bool DynamicGraphFusionEnabled() const noexcept; std::shared_ptr GetGpuAllocator(); std::shared_ptr GetCpuInputAllocator(); @@ -191,8 +198,13 @@ namespace Dml bool m_isMcdmDevice = false; bool m_areCustomHeapsSupported = false; bool m_areMetacommandsEnabled = true; - bool m_dynamicGraphFusionEnabled = false; + int m_currentGraphAnnotationId = 0; bool m_native16BitShaderOpsSupported = false; + bool m_graphCaptured = false; + bool m_graphCaptureEnabled = false; + + std::unordered_map>> m_capturedGraphs; + std::unordered_set m_graphCapturingDone; bool m_sessionInitialized = false; bool m_cpuSyncSpinningEnabled = false; ComPtr m_context; @@ -247,7 +259,7 @@ namespace Dml IDMLDevice* dmlDevice, Dml::ExecutionContext* executionContext, bool enableMetacommands, - bool enableDynamicGraphFusion, + bool enableGraphCapture, bool enableSyncSpinning ); @@ -283,12 +295,14 @@ namespace Dml return Status::OK(); } - onnxruntime::Status OnRunEnd(bool /*sync_stream*/, const onnxruntime::RunOptions& /*run_options*/) final override + Status OnRunStart(const onnxruntime::RunOptions& run_options) final { - // Flush any pending work to the GPU, but don't block for completion, permitting it - // to overlap other work. - m_impl->Flush(); - return Status::OK(); + return m_impl->OnRunStart(run_options); + } + + Status OnRunEnd(bool /*sync_stream*/, const onnxruntime::RunOptions& /*run_options*/) final + { + return m_impl->OnRunEnd(); } void Flush() @@ -311,16 +325,26 @@ namespace Dml return m_impl.Get(); } - bool DynamicGraphFusionEnabled() const - { - return m_impl->DynamicGraphFusionEnabled(); - } - virtual std::vector CreatePreferredAllocators() override { return m_impl->CreatePreferredAllocators(); } + bool IsGraphCaptureEnabled() const override + { + return m_impl->GraphCaptureEnabled(); + } + + bool IsGraphCaptured(int graph_annotation_id) const override + { + return m_impl->GraphCaptured(graph_annotation_id); + } + + Status ReplayGraph(int graph_annotation_id) override + { + return m_impl->ReplayGraph(graph_annotation_id); + } + private: ComPtr m_impl; }; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/FusedGraphKernel.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/FusedGraphKernel.cpp index 430ccec3cd..0cbf817a82 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/FusedGraphKernel.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/FusedGraphKernel.cpp @@ -13,26 +13,6 @@ namespace Dml { class FusedGraphKernel : public onnxruntime::OpKernel { - private: - struct ReusedCommandListState - { - // Re-usable command list, supporting descriptor heap, and DML binding table to update that heap. - ComPtr graphicsCommandList; - ComPtr commandAllocator; - ComPtr heap; - ComPtr bindingTable; - - // Bindings from previous executions of a re-used command list - mutable std::vector inputBindingAllocIds; - mutable std::vector outputBindingAllocIds; - mutable uint64_t tempBindingAllocId = 0; - - // Fence tracking the status of the command list's last execution, and whether its descriptor heap - // can safely be updated. - mutable ComPtr fence; - mutable uint64_t completionValue = 0; - }; - public: FusedGraphKernel() = delete; @@ -87,7 +67,7 @@ namespace Dml static_cast(persistentResourceSize), AllocatorRoundingMode::Disabled, m_persistentResource.GetAddressOf(), - m_persistentResourceAllocatorUnk.GetAddressOf())); + m_persistentResourceAllocatorUnknown.GetAddressOf())); m_persistentResourceBinding = DML_BUFFER_BINDING { m_persistentResource.Get(), 0, persistentResourceSize }; } @@ -99,7 +79,7 @@ namespace Dml // Queue references to objects which must be kept alive until resulting GPU work completes m_winmlProvider->QueueReference(m_compiledExecutionPlanOperator.Get()); - m_winmlProvider->QueueReference(m_persistentResourceAllocatorUnk.Get()); + m_winmlProvider->QueueReference(m_persistentResourceAllocatorUnknown.Get()); std::for_each( initializeResourceRefs.begin(), @@ -109,7 +89,13 @@ namespace Dml if (reuseCommandList) { - m_reusedCommandLists.push_back(BuildReusableCommandList()); + auto reusableCommandList = DmlGraphFusionHelper::BuildReusableCommandList( + m_provider.Get(), + m_compiledExecutionPlanOperator.Get(), + m_persistentResource.Get(), + m_persistentResourceBinding); + + m_reusedCommandLists.push_back(std::move(reusableCommandList)); } } @@ -162,17 +148,35 @@ namespace Dml // Queue references to objects which must be kept alive until resulting GPU work completes m_winmlProvider->QueueReference(m_compiledExecutionPlanOperator.Get()); - m_winmlProvider->QueueReference(m_persistentResourceAllocatorUnk.Get()); + m_winmlProvider->QueueReference(m_persistentResourceAllocatorUnknown.Get()); } else { if (m_reusedCommandLists.front()->fence && m_reusedCommandLists.front()->fence->GetCompletedValue() < m_reusedCommandLists.front()->completionValue) { - m_reusedCommandLists.push_front(BuildReusableCommandList()); + auto reusableCommandList = DmlGraphFusionHelper::BuildReusableCommandList( + m_provider.Get(), + m_compiledExecutionPlanOperator.Get(), + m_persistentResource.Get(), + m_persistentResourceBinding); + + m_reusedCommandLists.push_front(std::move(reusableCommandList)); } - - ExecuteReusableCommandList(kernelContext, *m_reusedCommandLists.front()); + + DmlGraphFusionHelper::ExecuteReusableCommandList( + kernelContext, + *m_reusedCommandLists.front(), + m_compiledExecutionPlanOperator.Get(), + Info(), + m_isInputsUploadedByDmlEP, + m_inputsUsed, + m_nonOwnedGraphInputsFromInitializers, + m_outputShapes, + m_winmlProvider.Get(), + m_provider.Get(), + m_persistentResourceAllocatorUnknown.Get()); + m_reusedCommandLists.push_back(std::move(m_reusedCommandLists.front())); m_reusedCommandLists.pop_front(); } @@ -244,198 +248,6 @@ namespace Dml } private: - std::unique_ptr BuildReusableCommandList() const - { - auto commandListState = std::make_unique(); - - ComPtr device; - ORT_THROW_IF_FAILED(m_provider->GetDmlDevice(device.GetAddressOf())); - - DML_BINDING_PROPERTIES execBindingProps = m_compiledExecutionPlanOperator->GetBindingProperties(); - - D3D12_DESCRIPTOR_HEAP_DESC desc = {}; - desc.Flags = D3D12_DESCRIPTOR_HEAP_FLAG_SHADER_VISIBLE; - desc.NumDescriptors = execBindingProps.RequiredDescriptorCount; - desc.Type = D3D12_DESCRIPTOR_HEAP_TYPE_CBV_SRV_UAV; - - ComPtr d3dDevice; - ORT_THROW_IF_FAILED(m_provider->GetD3DDevice(d3dDevice.GetAddressOf())); - - ORT_THROW_IF_FAILED(d3dDevice->CreateDescriptorHeap(&desc, IID_GRAPHICS_PPV_ARGS(commandListState->heap.ReleaseAndGetAddressOf()))); - - // Create a binding table for execution. - DML_BINDING_TABLE_DESC bindingTableDesc = {}; - bindingTableDesc.Dispatchable = m_compiledExecutionPlanOperator.Get(); - bindingTableDesc.CPUDescriptorHandle = commandListState->heap->GetCPUDescriptorHandleForHeapStart(); - bindingTableDesc.GPUDescriptorHandle = commandListState->heap->GetGPUDescriptorHandleForHeapStart(); - bindingTableDesc.SizeInDescriptors = execBindingProps.RequiredDescriptorCount; - - ORT_THROW_IF_FAILED(device->CreateBindingTable(&bindingTableDesc, IID_PPV_ARGS(&commandListState->bindingTable))); - - ORT_THROW_IF_FAILED(d3dDevice->CreateCommandAllocator( - m_provider->GetCommandListTypeForQueue(), - IID_GRAPHICS_PPV_ARGS(commandListState->commandAllocator.ReleaseAndGetAddressOf()))); - - ORT_THROW_IF_FAILED(d3dDevice->CreateCommandList( - 0, - m_provider->GetCommandListTypeForQueue(), - commandListState->commandAllocator.Get(), - nullptr, - IID_GRAPHICS_PPV_ARGS(commandListState->graphicsCommandList.ReleaseAndGetAddressOf()))); - - if (m_persistentResource) - { - DML_BINDING_DESC persistentResourceBindingDesc = - { DML_BINDING_TYPE_BUFFER, m_persistentResourceBinding ? &*m_persistentResourceBinding : nullptr }; - commandListState->bindingTable->BindPersistentResource(&persistentResourceBindingDesc); - } - - ID3D12DescriptorHeap* descriptorHeaps[] = { commandListState->heap.Get() }; - commandListState->graphicsCommandList->SetDescriptorHeaps(ARRAYSIZE(descriptorHeaps), descriptorHeaps); - - ComPtr recorder; - ORT_THROW_IF_FAILED(device->CreateCommandRecorder(IID_PPV_ARGS(recorder.GetAddressOf()))); - - recorder->RecordDispatch(commandListState->graphicsCommandList.Get(), m_compiledExecutionPlanOperator.Get(), commandListState->bindingTable.Get()); - - ORT_THROW_IF_FAILED(commandListState->graphicsCommandList->Close()); - - return commandListState; - } - - void ExecuteReusableCommandList(onnxruntime::OpKernelContext* kernelContext, ReusedCommandListState& commandListState) const - { - DML_BINDING_PROPERTIES execBindingProps = m_compiledExecutionPlanOperator->GetBindingProperties(); - - std::vector inputBindings(kernelContext->InputCount()); - std::vector inputBindingDescs(kernelContext->InputCount()); - - OpKernelContextWrapper contextWrapper( - kernelContext, - Info().GetExecutionProvider(), - true, - nullptr); - - // Populate input bindings, excluding those which were specified as owned by DML and provided - // at initialization instead. - commandListState.inputBindingAllocIds.resize(inputBindings.size()); - bool inputBindingsChanged = false; - - for (uint32_t i = 0; i < inputBindings.size(); ++i) - { - if (!m_isInputsUploadedByDmlEP[i] && m_inputsUsed[i]) - { - if (m_nonOwnedGraphInputsFromInitializers[i]) - { - inputBindings[i].Buffer = m_nonOwnedGraphInputsFromInitializers[i].Get(); - inputBindings[i].SizeInBytes = m_nonOwnedGraphInputsFromInitializers[i]->GetDesc().Width; - inputBindingDescs[i] = {DML_BINDING_TYPE_BUFFER, &inputBindings[i]}; - } - else - { - assert(kernelContext->InputType(gsl::narrow_cast(i))->IsTensorType()); - const onnxruntime::Tensor* tensor = kernelContext->Input(gsl::narrow_cast(i)); - - uint64_t allocId; - DmlGraphFusionHelper::UnwrapTensor(m_winmlProvider.Get(), tensor, &inputBindings[i].Buffer, &allocId); - inputBindingsChanged = inputBindingsChanged || (!allocId || commandListState.inputBindingAllocIds[i] != allocId); - inputBindings[i].Buffer->Release(); // Avoid holding an additional reference - inputBindings[i].SizeInBytes = DmlGraphFusionHelper::AlignToPow2(tensor->SizeInBytes(), 4); - inputBindingDescs[i] = {DML_BINDING_TYPE_BUFFER, &inputBindings[i]}; - commandListState.inputBindingAllocIds[i] = allocId; - } - } - } - - if (inputBindingsChanged) - { - commandListState.bindingTable->BindInputs(gsl::narrow_cast(inputBindingDescs.size()), inputBindingDescs.data()); - } - - // Populate Output bindings - std::vector outputBindings(kernelContext->OutputCount()); - std::vector outputBindingDescs(kernelContext->OutputCount()); - - commandListState.outputBindingAllocIds.resize(outputBindings.size()); - bool outputBindingsChanged = false; - - for (uint32_t i = 0; i < outputBindings.size(); ++i) - { - std::vector outputDims; - outputDims.reserve(m_outputShapes.GetShape(i).size()); - for (uint32_t dimSize : m_outputShapes.GetShape(i)) - { - outputDims.push_back(dimSize); - } - - onnxruntime::Tensor* tensor = kernelContext->Output( - static_cast(i), - onnxruntime::TensorShape::FromExistingBuffer(outputDims) - ); - - uint64_t allocId; - DmlGraphFusionHelper::UnwrapTensor(m_winmlProvider.Get(), tensor, &outputBindings[i].Buffer, &allocId); - outputBindingsChanged = outputBindingsChanged || (!allocId || commandListState.outputBindingAllocIds[i] != allocId); - outputBindings[i].Buffer->Release(); // Avoid holding an additional reference - outputBindings[i].SizeInBytes = DmlGraphFusionHelper::AlignToPow2(tensor->SizeInBytes(), 4); - outputBindingDescs[i] = {DML_BINDING_TYPE_BUFFER, &outputBindings[i]}; - commandListState.outputBindingAllocIds[i] = allocId; - } - - if (outputBindingsChanged) - { - commandListState.bindingTable->BindOutputs(gsl::narrow_cast(outputBindingDescs.size()), outputBindingDescs.data()); - } - - if (execBindingProps.TemporaryResourceSize > 0) - { - // Allocate temporary data which will automatically be freed when the GPU work - // which is scheduled up to the point that this method returns has completed. - ComPtr tempAlloc; - uint64_t tempAllocId = 0; - ORT_THROW_IF_FAILED(contextWrapper.AllocateTemporaryData(static_cast(execBindingProps.TemporaryResourceSize), tempAlloc.GetAddressOf(), &tempAllocId)); - - ComPtr tempResourceUnk; - m_winmlProvider->GetABIDataInterface(false, tempAlloc.Get(), &tempResourceUnk); - - // Bind the temporary resource. - ComPtr tempResource; - ORT_THROW_IF_FAILED(tempResourceUnk->QueryInterface(tempResource.GetAddressOf())); - DML_BUFFER_BINDING tempBufferBinding = {tempResource.Get(), 0, execBindingProps.TemporaryResourceSize}; - DML_BINDING_DESC tempBindingDesc = { DML_BINDING_TYPE_BUFFER, &tempBufferBinding }; - - if (!tempAllocId || commandListState.tempBindingAllocId != tempAllocId) - { - commandListState.bindingTable->BindTemporaryResource(&tempBindingDesc); - } - - commandListState.tempBindingAllocId = tempAllocId; - } - - // Execute the command list and if it succeeds, update the fence value at which this command may be - // re-used. - ComPtr fence; - uint64_t completionValue; - HRESULT hr = m_provider->ExecuteCommandList(commandListState.graphicsCommandList.Get(), fence.GetAddressOf(), &completionValue); - - if (hr == DXGI_ERROR_DEVICE_REMOVED) - { - ComPtr device; - ORT_THROW_IF_FAILED(m_provider->GetD3DDevice(&device)); - ORT_THROW_IF_FAILED(device->GetDeviceRemovedReason()); - } - - ORT_THROW_IF_FAILED(hr); - commandListState.fence = fence; - commandListState.completionValue = completionValue; - - // Queue references to objects which must be kept alive until resulting GPU work completes - m_winmlProvider->QueueReference(WRAP_GRAPHICS_UNKNOWN(commandListState.graphicsCommandList).Get()); - m_winmlProvider->QueueReference(WRAP_GRAPHICS_UNKNOWN(commandListState.heap).Get()); - m_winmlProvider->QueueReference(commandListState.bindingTable.Get()); - m_winmlProvider->QueueReference(m_persistentResourceAllocatorUnk.Get()); - } - ComPtr m_compiledExecutionPlanOperator; std::vector m_inputsUsed; const void* m_executionHandle = nullptr; @@ -443,11 +255,11 @@ namespace Dml ComPtr m_provider; Windows::AI::MachineLearning::Adapter::EdgeShapes& m_outputShapes; - mutable std::deque> m_reusedCommandLists; + mutable std::deque> m_reusedCommandLists; std::optional m_persistentResourceBinding; ComPtr m_persistentResource; - ComPtr m_persistentResourceAllocatorUnk; // Controls when the persistent resource is returned to the allocator + ComPtr m_persistentResourceAllocatorUnknown; // Controls when the persistent resource is returned to the allocator std::vector m_isInputsUploadedByDmlEP; std::vector> m_nonOwnedGraphInputsFromInitializers; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h index 49abb89449..5ecc49f237 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h @@ -74,4 +74,4 @@ using Microsoft::WRL::ComPtr; #include "TensorDesc.h" #include "DescriptorPool.h" #include "IExecutionProvider.h" -#include "Utility.h" \ No newline at end of file +#include "Utility.h" diff --git a/onnxruntime/core/providers/dml/dml_provider_factory.cc b/onnxruntime/core/providers/dml/dml_provider_factory.cc index 796325f2c2..7a665ef82f 100644 --- a/onnxruntime/core/providers/dml/dml_provider_factory.cc +++ b/onnxruntime/core/providers/dml/dml_provider_factory.cc @@ -48,14 +48,13 @@ struct DMLProviderFactory : IExecutionProviderFactory { IDMLDevice* dml_device, ID3D12CommandQueue* cmd_queue, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api ) : dml_device_(dml_device), cmd_queue_(cmd_queue), metacommands_enabled_(!disable_metacommands), - dynamic_graph_fusion_enabled_(enable_dynamic_graph_fusion), python_api_(python_api) { + graph_capture_enabled_ = ConfigValueIsTrue(config_options.GetConfigOrDefault(kOrtSessionOptionsConfigEnableGraphCapture, "0")); cpu_sync_spinning_enabled_ = ConfigValueIsTrue(config_options.GetConfigOrDefault(kOrtSessionOptionsConfigEnableCpuSyncSpinning, "0")); } @@ -69,7 +68,7 @@ struct DMLProviderFactory : IExecutionProviderFactory { ComPtr dml_device_{}; ComPtr cmd_queue_{}; bool metacommands_enabled_ = true; - bool dynamic_graph_fusion_enabled_ = false; + bool graph_capture_enabled_ = false; bool cpu_sync_spinning_enabled_ = false; bool python_api_ = false; }; @@ -92,7 +91,7 @@ std::unique_ptr DMLProviderFactory::CreateProvider() { execution_context = wil::MakeOrThrow(d3d12_device.Get(), dml_device_.Get(), cmd_queue_.Get(), cpu_sync_spinning_enabled_, false); } - auto provider = Dml::CreateExecutionProvider(dml_device_.Get(), execution_context.Get(), metacommands_enabled_, dynamic_graph_fusion_enabled_, cpu_sync_spinning_enabled_); + auto provider = Dml::CreateExecutionProvider(dml_device_.Get(), execution_context.Get(), metacommands_enabled_, graph_capture_enabled_, cpu_sync_spinning_enabled_); return provider; } @@ -104,17 +103,7 @@ std::shared_ptr CreateExecutionProviderFactory_DML(co IDMLDevice* dml_device, ID3D12CommandQueue* cmd_queue, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api) { - ComPtr ownedDmlDevice; - if (dml_device == nullptr) { - ComPtr d3d12_device; - ORT_THROW_IF_FAILED(cmd_queue->GetDevice(IID_PPV_ARGS(&d3d12_device))); - - ownedDmlDevice = onnxruntime::DMLProviderFactoryCreator::CreateDMLDevice(d3d12_device.Get()); - dml_device = ownedDmlDevice.Get(); - } - #ifndef _GAMING_XBOX // Validate that the D3D12 devices match between DML and the command queue. This specifically asks for IUnknown in // order to be able to compare the pointers for COM object identity. @@ -133,7 +122,7 @@ std::shared_ptr CreateExecutionProviderFactory_DML(co const Env& env = Env::Default(); auto luid = d3d12_device->GetAdapterLuid(); env.GetTelemetryProvider().LogExecutionProviderEvent(&luid); - return std::make_shared(config_options, dml_device, cmd_queue, disable_metacommands, enable_dynamic_graph_fusion, python_api); + return std::make_shared(config_options, dml_device, cmd_queue, disable_metacommands, python_api); } void DmlConfigureProviderFactoryMetacommandsEnabled(IExecutionProviderFactory* factory, bool metacommandsEnabled) { @@ -305,7 +294,6 @@ std::shared_ptr DMLProviderFactoryCreator::CreateFrom const ConfigOptions& config_options, const OrtDmlDeviceOptions* device_options, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api) { auto default_device_options = OrtDmlDeviceOptions { Default, Gpu }; if (device_options == nullptr) { @@ -353,7 +341,7 @@ std::shared_ptr DMLProviderFactoryCreator::CreateFrom adapters.begin(), [](auto& a){ return a.Adapter; }); - return onnxruntime::DMLProviderFactoryCreator::CreateFromAdapterList(config_options, std::move(adapters), disable_metacommands, enable_dynamic_graph_fusion, python_api); + return onnxruntime::DMLProviderFactoryCreator::CreateFromAdapterList(config_options, std::move(adapters), disable_metacommands, python_api); } static std::optional ParsePerformancePreference(const ProviderOptions& provider_options) { @@ -445,13 +433,12 @@ std::shared_ptr DMLProviderFactoryCreator::CreateFrom bool python_api) { bool disable_metacommands = ParseBoolean(provider_options, "disable_metacommands"); - bool enable_dynamic_graph_fusion = ParseBoolean(provider_options, "enable_dynamic_graph_fusion"); bool skip_software_device_check = false; auto device_id = ParseDeviceId(provider_options); if (device_id.has_value()) { - return onnxruntime::DMLProviderFactoryCreator::Create(config_options, device_id.value(), skip_software_device_check, disable_metacommands, enable_dynamic_graph_fusion, python_api); + return onnxruntime::DMLProviderFactoryCreator::Create(config_options, device_id.value(), skip_software_device_check, disable_metacommands, python_api); } auto preference = ParsePerformancePreference(provider_options); @@ -459,7 +446,7 @@ std::shared_ptr DMLProviderFactoryCreator::CreateFrom // If no preference/filters are specified then create with default preference/filters. if (!preference.has_value() && !filter.has_value()) { - return onnxruntime::DMLProviderFactoryCreator::CreateFromDeviceOptions(config_options, nullptr, disable_metacommands, enable_dynamic_graph_fusion, python_api); + return onnxruntime::DMLProviderFactoryCreator::CreateFromDeviceOptions(config_options, nullptr, disable_metacommands, python_api); } if (!preference.has_value()) { @@ -473,7 +460,7 @@ std::shared_ptr DMLProviderFactoryCreator::CreateFrom OrtDmlDeviceOptions device_options; device_options.Preference = preference.value(); device_options.Filter = filter.value(); - return onnxruntime::DMLProviderFactoryCreator::CreateFromDeviceOptions(config_options, &device_options, disable_metacommands, enable_dynamic_graph_fusion, python_api); + return onnxruntime::DMLProviderFactoryCreator::CreateFromDeviceOptions(config_options, &device_options, disable_metacommands, python_api); } Microsoft::WRL::ComPtr DMLProviderFactoryCreator::CreateD3D12Device( @@ -568,7 +555,6 @@ std::shared_ptr CreateDMLDeviceAndProviderFactory( const ConfigOptions& config_options, ID3D12Device* d3d12_device, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api = false) { D3D12_COMMAND_QUEUE_DESC cmd_queue_desc = {}; cmd_queue_desc.Type = CalculateCommandListType(d3d12_device); @@ -589,7 +575,7 @@ std::shared_ptr CreateDMLDeviceAndProviderFactory( dml_device = onnxruntime::DMLProviderFactoryCreator::CreateDMLDevice(d3d12_device); } - return CreateExecutionProviderFactory_DML(config_options, dml_device.Get(), cmd_queue.Get(), disable_metacommands, enable_dynamic_graph_fusion, python_api); + return CreateExecutionProviderFactory_DML(config_options, dml_device.Get(), cmd_queue.Get(), disable_metacommands, python_api); } std::shared_ptr DMLProviderFactoryCreator::Create( @@ -597,17 +583,15 @@ std::shared_ptr DMLProviderFactoryCreator::Create( int device_id, bool skip_software_device_check, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api) { ComPtr d3d12_device = CreateD3D12Device(device_id, skip_software_device_check); - return CreateDMLDeviceAndProviderFactory(config_options, d3d12_device.Get(), disable_metacommands, enable_dynamic_graph_fusion, python_api); + return CreateDMLDeviceAndProviderFactory(config_options, d3d12_device.Get(), disable_metacommands, python_api); } std::shared_ptr DMLProviderFactoryCreator::CreateFromAdapterList( const ConfigOptions& config_options, std::vector>&& adapters, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api) { // Choose the first device from the list since it's the highest priority auto adapter = adapters[0]; @@ -629,7 +613,7 @@ std::shared_ptr DMLProviderFactoryCreator::CreateFrom ORT_THROW_IF_FAILED(D3D12CreateDevice(adapter.Get(), feature_level, IID_GRAPHICS_PPV_ARGS(d3d12_device.ReleaseAndGetAddressOf()))); } - return CreateDMLDeviceAndProviderFactory(config_options, d3d12_device.Get(), disable_metacommands, enable_dynamic_graph_fusion, python_api); + return CreateDMLDeviceAndProviderFactory(config_options, d3d12_device.Get(), disable_metacommands, python_api); } } // namespace onnxruntime @@ -654,7 +638,6 @@ API_IMPL_BEGIN dml_device, cmd_queue, false, - false, false)); API_IMPL_END return nullptr; diff --git a/onnxruntime/core/providers/dml/dml_provider_factory_creator.h b/onnxruntime/core/providers/dml/dml_provider_factory_creator.h index e0204aaef6..0a1bd81b15 100644 --- a/onnxruntime/core/providers/dml/dml_provider_factory_creator.h +++ b/onnxruntime/core/providers/dml/dml_provider_factory_creator.h @@ -26,7 +26,6 @@ struct DMLProviderFactoryCreator { int device_id, bool skip_software_device_check, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api = false); static std::shared_ptr CreateFromProviderOptions( @@ -38,14 +37,12 @@ struct DMLProviderFactoryCreator { const ConfigOptions& config_options, const OrtDmlDeviceOptions* device_options, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api = false); static std::shared_ptr CreateFromAdapterList( const ConfigOptions& config_options, std::vector>&& dxcore_devices, bool disable_metacommands, - bool enable_dynamic_graph_fusion, bool python_api = false); static Microsoft::WRL::ComPtr CreateD3D12Device(int device_id, bool skip_software_device_check); diff --git a/onnxruntime/core/providers/dml/dml_session_options_config_keys.h b/onnxruntime/core/providers/dml/dml_session_options_config_keys.h index 77b2da013c..6418366342 100644 --- a/onnxruntime/core/providers/dml/dml_session_options_config_keys.h +++ b/onnxruntime/core/providers/dml/dml_session_options_config_keys.h @@ -22,4 +22,5 @@ // The default value is "0" static const char* const kOrtSessionOptionsConfigDisableDmlGraphFusion = "ep.dml.disable_graph_fusion"; static const char* const kOrtSessionOptionsConfigEnableGraphSerialization = "ep.dml.enable_graph_serialization"; +static const char* const kOrtSessionOptionsConfigEnableGraphCapture = "ep.dml.enable_graph_capture"; static const char* const kOrtSessionOptionsConfigEnableCpuSyncSpinning = "ep.dml.enable_cpu_sync_spinning"; diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index 49d16cf6b1..c1cd21570a 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -148,8 +148,8 @@ static bool HasMemcpyNodes(const Graph& graph) { return false; } -static bool AreAllComputeNodesAssignedToCudaOrJsEp(const Graph& graph) { - bool nodes_on_cpu_and_cuda_and_js_eps_only = true; +static bool AreAllComputeNodesAssignedToCudaOrJsOrDmlEp(const Graph& graph) { + bool nodes_on_cpu_and_cuda_and_js_and_dml_eps_only = true; for (const auto& node : graph.Nodes()) { const auto& node_provider = node.GetExecutionProviderType(); @@ -158,9 +158,10 @@ static bool AreAllComputeNodesAssignedToCudaOrJsEp(const Graph& graph) { if (!node_provider.empty() && !(node_provider == kCudaExecutionProvider || node_provider == kRocmExecutionProvider || - node_provider == kJsExecutionProvider) && + node_provider == kJsExecutionProvider || + node_provider == kDmlExecutionProvider) && node_provider != kCpuExecutionProvider) { - nodes_on_cpu_and_cuda_and_js_eps_only = false; + nodes_on_cpu_and_cuda_and_js_and_dml_eps_only = false; break; } } @@ -171,7 +172,7 @@ static bool AreAllComputeNodesAssignedToCudaOrJsEp(const Graph& graph) { // We allow CPU EPs to show up in the EP list as long as thre is no Memcpy // involved as shape subgraphs will be forced onto CPU and these will not have // Memcpy nodes involved. - return nodes_on_cpu_and_cuda_and_js_eps_only && !HasMemcpyNodes(graph); + return nodes_on_cpu_and_cuda_and_js_and_dml_eps_only && !HasMemcpyNodes(graph); } static bool AreAllNodesInMainGraphAssignedToOneEp(const Graph& graph, ProviderType provider) { @@ -1755,7 +1756,14 @@ common::Status InferenceSession::Initialize() { [](char ch) { return std::tolower(ch); }); bool dml_graph_serialization_enabled = dml_graph_serialization_enabled_config_val == "true"; - if (dml_graph_fusion_enabled) { + if (static_cast(dmlExecutionProvider)->IsGraphCaptureEnabled()) { + std::unique_ptr dmlRuntimeGraphFusionTransformer = std::make_unique("DmlRuntimeGraphFusionTransformer", + dmlExecutionProvider); + if (dmlRuntimeGraphFusionTransformer == nullptr) { + return Status(common::ONNXRUNTIME, common::FAIL, "DmlRuntimeGraphFusionTransformer is nullptr"); + } + ORT_RETURN_IF_ERROR_SESSIONID_(graph_transformer_mgr_.Register(std::move(dmlRuntimeGraphFusionTransformer), onnxruntime::TransformerLevel::Level3)); + } else if (dml_graph_fusion_enabled) { std::unique_ptr dmlGraphFusionTransformer = std::make_unique("DmlGraphFusionTransformer", dmlExecutionProvider, dml_graph_serialization_enabled); @@ -1763,15 +1771,6 @@ common::Status InferenceSession::Initialize() { return Status(common::ONNXRUNTIME, common::FAIL, "DmlGraphFusionTransformer is nullptr"); } ORT_RETURN_IF_ERROR_SESSIONID_(graph_transformer_mgr_.Register(std::move(dmlGraphFusionTransformer), onnxruntime::TransformerLevel::Level3)); - - if (static_cast(dmlExecutionProvider)->DynamicGraphFusionEnabled()) { - std::unique_ptr dmlRuntimeGraphFusionTransformer = std::make_unique("DmlRuntimeGraphFusionTransformer", - dmlExecutionProvider); - if (dmlRuntimeGraphFusionTransformer == nullptr) { - return Status(common::ONNXRUNTIME, common::FAIL, "DmlRuntimeGraphFusionTransformer is nullptr"); - } - ORT_RETURN_IF_ERROR_SESSIONID_(graph_transformer_mgr_.Register(std::move(dmlRuntimeGraphFusionTransformer), onnxruntime::TransformerLevel::Level3)); - } } // This transformer applies DML-specific fusions that go beyond what ORT offers by default @@ -1831,7 +1830,8 @@ common::Status InferenceSession::Initialize() { onnxruntime::kTensorrtExecutionProvider, onnxruntime::kCudaExecutionProvider, onnxruntime::kRocmExecutionProvider, - onnxruntime::kJsExecutionProvider}; + onnxruntime::kJsExecutionProvider, + onnxruntime::kDmlExecutionProvider}; for (auto& it : graph_support_ep_list) { auto* target_ep = execution_providers_.Get(it); @@ -1852,12 +1852,13 @@ common::Status InferenceSession::Initialize() { if (strcmp(target_ep->Type().c_str(), onnxruntime::kCudaExecutionProvider) == 0 || strcmp(target_ep->Type().c_str(), onnxruntime::kRocmExecutionProvider) == 0 || - strcmp(target_ep->Type().c_str(), onnxruntime::kJsExecutionProvider) == 0) { + strcmp(target_ep->Type().c_str(), onnxruntime::kJsExecutionProvider) == 0 || + strcmp(target_ep->Type().c_str(), onnxruntime::kDmlExecutionProvider) == 0) { // Ensure that all nodes have been partitioned to CUDA/JS or CPU EP && there are no memcpy nodes // The reasoning behind this logic is that certain shape nodes will be forced onto CPU // and as long as there are no memcpy nodes this is confirmation that no compute nodes have been placed on the CPU EP // which is all we care about. - if (!AreAllComputeNodesAssignedToCudaOrJsEp(graph)) { + if (!AreAllComputeNodesAssignedToCudaOrJsOrDmlEp(graph)) { LOGS(*session_logger_, ERROR) << "This session cannot use the graph capture feature as requested by the user " << " as all compute graph nodes have not been partitioned to the " << target_ep->Type(); diff --git a/onnxruntime/test/perftest/command_args_parser.cc b/onnxruntime/test/perftest/command_args_parser.cc index b637ea1874..b05f58a4e7 100644 --- a/onnxruntime/test/perftest/command_args_parser.cc +++ b/onnxruntime/test/perftest/command_args_parser.cc @@ -69,7 +69,7 @@ namespace perftest { "\t [DML only] [performance_preference]: DML device performance preference, options: 'default', 'minimum_power', 'high_performance', \n" "\t [DML only] [device_filter]: DML device filter, options: 'any', 'gpu', 'npu', \n" "\t [DML only] [disable_metacommands]: Options: 'true', 'false', \n" - "\t [DML only] [enable_dynamic_graph_fusion]: Options: 'true', 'false', \n" + "\t [DML only] [enable_graph_capture]: Options: 'true', 'false', \n" "\t [DML only] [enable_graph_serialization]: Options: 'true', 'false', \n" "\n" "\t [OpenVINO only] [device_type]: Overrides the accelerator hardware type and precision with these values at runtime.\n" diff --git a/onnxruntime/test/perftest/ort_test_session.cc b/onnxruntime/test/perftest/ort_test_session.cc index b1dcf094a3..f93efba811 100644 --- a/onnxruntime/test/perftest/ort_test_session.cc +++ b/onnxruntime/test/perftest/ort_test_session.cc @@ -520,7 +520,7 @@ select from 'TF8', 'TF16', 'UINT8', 'FLOAT', 'ITENSOR'. \n)"); dml_options["performance_preference"] = "high_performance"; dml_options["device_filter"] = "gpu"; dml_options["disable_metacommands"] = "false"; - dml_options["enable_dynamic_graph_fusion"] = "false"; + dml_options["enable_graph_capture"] = "false"; #ifdef _MSC_VER std::string ov_string = ToUTF8String(performance_test_config.run_config.ep_runtime_config_string); #else @@ -564,16 +564,16 @@ select from 'TF8', 'TF16', 'UINT8', 'FLOAT', 'ITENSOR'. \n)"); dml_options[key] = value; } else { ORT_THROW( - "[ERROR] [DML] You have selcted wrong value for the key 'disable_metacommands'. " + "[ERROR] [DML] You have selected a wrong value for the key 'disable_metacommands'. " "Select from 'true' or 'false' \n"); } - } else if (key == "enable_dynamic_graph_fusion") { + } else if (key == "enable_graph_capture") { std::set ov_supported_values = {"true", "True", "false", "False"}; if (ov_supported_values.find(value) != ov_supported_values.end()) { dml_options[key] = value; } else { ORT_THROW( - "[ERROR] [DML] You have selcted wrong value for the key 'enable_dynamic_graph_fusion'. " + "[ERROR] [DML] You have selected a wrong value for the key 'enable_graph_capture'. " "Select from 'true' or 'false' \n"); } } else if (key == "enable_graph_serialization") { @@ -582,7 +582,7 @@ select from 'TF8', 'TF16', 'UINT8', 'FLOAT', 'ITENSOR'. \n)"); session_options.AddConfigEntry(kOrtSessionOptionsConfigEnableGraphSerialization, value.data()); } else { ORT_THROW( - "[ERROR] [DML] You have selcted wrong value for the key 'enable_graph_serialization'. " + "[ERROR] [DML] You have selected a wrong value for the key 'enable_graph_serialization'. " "Select from 'true' or 'false' \n"); } } diff --git a/onnxruntime/test/python/onnxruntime_test_python_dmlgraph.py b/onnxruntime/test/python/onnxruntime_test_python_dmlgraph.py new file mode 100644 index 0000000000..29292c2a77 --- /dev/null +++ b/onnxruntime/test/python/onnxruntime_test_python_dmlgraph.py @@ -0,0 +1,185 @@ +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. + +import unittest +from typing import Dict, List + +import numpy as np +from helper import get_name + +import onnxruntime as onnxrt + + +class DmlGraphHelper: + def __init__( + self, + ort_session: onnxrt.InferenceSession, + input_and_output_shape: Dict[str, List[int]], + device_id: int = 0, + ): + self.input_names = [input.name for input in ort_session.get_inputs()] + self.output_names = [output.name for output in ort_session.get_outputs()] + + self.input_and_output_shape = input_and_output_shape + self.io_numpy_type = self.get_io_numpy_type_map(ort_session) + self.io_binding = ort_session.io_binding() + self.io_ort_value = {} + + for name in self.input_names + self.output_names: + ort_value = onnxrt.OrtValue.ortvalue_from_shape_and_type( + input_and_output_shape[name], self.io_numpy_type[name], "dml", device_id + ) + self.io_ort_value[name] = ort_value + if name in self.input_names: + self.io_binding.bind_ortvalue_input(name, ort_value) + else: + self.io_binding.bind_ortvalue_output(name, ort_value) + + def get_io_numpy_type_map(self, ort_session: onnxrt.InferenceSession): + ort_type_to_numpy_type = { + "tensor(int64)": np.longlong, + "tensor(int32)": np.intc, + "tensor(float)": np.float32, + "tensor(float16)": np.float16, + } + + name_to_numpy_type = {} + for _input in ort_session.get_inputs(): + name_to_numpy_type[_input.name] = ort_type_to_numpy_type[_input.type] + + for output in ort_session.get_outputs(): + name_to_numpy_type[output.name] = ort_type_to_numpy_type[output.type] + + return name_to_numpy_type + + def update_inputs(self, inputs: Dict[str, np.ndarray]): + for input_name in self.input_names: + self.io_ort_value[input_name].update_inplace(inputs[input_name]) + + def get_output(self, output_name: str): + return self.io_ort_value[output_name].numpy() + + +class TestInferenceSessionWithDmlGraph(unittest.TestCase): + def test_ort_value_update_in_place(self): + x0 = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], dtype=np.float32) + ortvalue_cpu = onnxrt.OrtValue.ortvalue_from_numpy(x0) + np.testing.assert_allclose(x0, ortvalue_cpu.numpy()) + + x1 = np.array([[10.0, 20.0], [30.0, 40.0], [50.0, 60.0]], dtype=np.float32) + ortvalue_cpu.update_inplace(x1) + np.testing.assert_allclose(x1, ortvalue_cpu.numpy()) + + if "DmlExecutionProvider" in onnxrt.get_available_providers(): + ortvalue_gpu = onnxrt.OrtValue.ortvalue_from_numpy(x0, "dml", 0) + np.testing.assert_allclose(x0, ortvalue_gpu.numpy()) + + ortvalue_gpu.update_inplace(x1) + np.testing.assert_allclose(x1, ortvalue_gpu.numpy()) + + def test_select_ep_to_run_dml_graph(self): + if "DmlExecutionProvider" in onnxrt.get_available_providers(): + providers = ["DmlExecutionProvider"] + self.run_model_with_dml_graph(providers) + self.run_model_with_dml_graph_annotation(providers) + + def run_model_with_dml_graph(self, providers): + INPUT_SIZE = 1280 # noqa: N806 + x = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]] * INPUT_SIZE, dtype=np.float32) + y = np.array([[0.0], [0.0], [0.0]] * INPUT_SIZE, dtype=np.float32) + x_ortvalue = onnxrt.OrtValue.ortvalue_from_numpy(x, "dml", 0) + y_ortvalue = onnxrt.OrtValue.ortvalue_from_numpy(y, "dml", 0) + + onnxrt.set_default_logger_severity(0) + sess_options = onnxrt.SessionOptions() + sess_options.add_session_config_entry("ep.dml.enable_graph_capture", "1") + session = onnxrt.InferenceSession(get_name("matmul_2.onnx"), providers=providers, sess_options=sess_options) + io_binding = session.io_binding() + + # Bind the input and output + io_binding.bind_ortvalue_input("X", x_ortvalue) + io_binding.bind_ortvalue_output("Y", y_ortvalue) + + ro = onnxrt.RunOptions() + + # One regular run for the necessary memory allocation and dml graph capturing + session.run_with_iobinding(io_binding, ro) + expected_y = np.array([[5.0], [11.0], [17.0]] * INPUT_SIZE, dtype=np.float32) + np.testing.assert_allclose(expected_y, y_ortvalue.numpy(), rtol=1e-05, atol=1e-05) + + # After capturing, DML graph replay happens from this Run onwards + session.run_with_iobinding(io_binding, ro) + np.testing.assert_allclose(expected_y, y_ortvalue.numpy(), rtol=1e-05, atol=1e-05) + + # Update input and then replay DML graph + x_ortvalue.update_inplace( + np.array( + [[10.0, 20.0], [30.0, 40.0], [50.0, 60.0]] * INPUT_SIZE, + dtype=np.float32, + ) + ) + session.run_with_iobinding(io_binding, ro) + np.testing.assert_allclose( + np.array([[50.0], [110.0], [170.0]] * INPUT_SIZE, dtype=np.float32), + y_ortvalue.numpy(), + rtol=1e-05, + atol=1e-05, + ) + + def run_model_with_dml_graph_annotation(self, providers): + INPUT_SIZE = 1280 # noqa: N806 + + x_base = [[1.0, 2.0], [3.0, 4.0], [5.0, 6.0], [7.0, 8.0]] + y_base = [[0.0], [0.0], [0.0], [0.0]] + expected_y_base = [[5.0], [11.0], [17.0], [23.0]] + + x_base_mul_10 = [[10.0, 20.0], [30.0, 40.0], [50.0, 60.0], [70.0, 80.0]] + expected_y_base_mul_10 = [[50.0], [110.0], [170.0], [230.0]] + + test_num = 4 + + x_ortvalues = [] + y_ortvalues = [] + for i in range(test_num): + x = np.array(x_base[: i + 1][:] * INPUT_SIZE, dtype=np.float32) + y = np.array(y_base[: i + 1][:] * INPUT_SIZE, dtype=np.float32) + x_ortvalues.append(onnxrt.OrtValue.ortvalue_from_numpy(x, "dml", 0)) + y_ortvalues.append(onnxrt.OrtValue.ortvalue_from_numpy(y, "dml", 0)) + + onnxrt.set_default_logger_severity(0) + sess_options = onnxrt.SessionOptions() + sess_options.add_session_config_entry("ep.dml.enable_graph_capture", "1") + session = onnxrt.InferenceSession(get_name("matmul_2.onnx"), providers=providers, sess_options=sess_options) + io_bindings = [session.io_binding()] * test_num + ro = onnxrt.RunOptions() + + # Regular run to capture DML graph + for i in range(test_num): + io_bindings[i].bind_ortvalue_input("X", x_ortvalues[i]) + io_bindings[i].bind_ortvalue_output("Y", y_ortvalues[i]) + # TODO: Temporarily remove the default dml graph capture test for the first regular run + # because it fails on a training CI. Need to investigate the root cause. + ro.add_run_config_entry("gpu_graph_id", str(i + 1)) + io_bindings[i].synchronize_inputs() + session.run_with_iobinding(io_bindings[i], ro) + io_bindings[i].synchronize_outputs() + expected_y = np.array(expected_y_base[: i + 1][:] * INPUT_SIZE, dtype=np.float32) + np.testing.assert_allclose(expected_y, y_ortvalues[i].numpy(), rtol=1e-05, atol=1e-05) + + del ro + ro = onnxrt.RunOptions() + + # After capturing, DML graph replay happens from this Run onwards + for i in range(test_num): + # Update input and then replay DML graph + x_ortvalues[i].update_inplace(np.array(x_base_mul_10[: i + 1][:] * INPUT_SIZE, dtype=np.float32)) + ro.add_run_config_entry("gpu_graph_id", str(i + 1)) + io_bindings[i].synchronize_inputs() + session.run_with_iobinding(io_bindings[i], ro) + io_bindings[i].synchronize_outputs() + expected_y = np.array(expected_y_base_mul_10[: i + 1][:] * INPUT_SIZE, dtype=np.float32) + np.testing.assert_allclose(expected_y, y_ortvalues[i].numpy(), rtol=1e-05, atol=1e-05) + + +if __name__ == "__main__": + unittest.main() diff --git a/onnxruntime/test/shared_lib/test_inference.cc b/onnxruntime/test/shared_lib/test_inference.cc index a7ce8127a7..2ccd3c69ab 100644 --- a/onnxruntime/test/shared_lib/test_inference.cc +++ b/onnxruntime/test/shared_lib/test_inference.cc @@ -47,6 +47,15 @@ #include #endif +#ifdef USE_DML +#include +#include +#include +#include "core/providers/dml/DmlExecutionProvider/src/ErrorHandling.h" + +using Microsoft::WRL::ComPtr; +#endif + // Once we use C++17 this could be replaced with std::size template constexpr size_t countof(T (&)[N]) { return N; } @@ -95,6 +104,201 @@ void RunSession(OrtAllocator* allocator, Ort::Session& session_object, } } +#ifdef USE_DML +struct DmlObjects { + ComPtr d3d12_device; + ComPtr command_queue; + ComPtr command_allocator; + ComPtr command_list; + ComPtr upload_buffer; + ComPtr dml_device; +}; + +static DmlObjects CreateDmlObjects() { + D3D12_COMMAND_QUEUE_DESC command_queue_description = + { + D3D12_COMMAND_LIST_TYPE_DIRECT, + 0, + D3D12_COMMAND_QUEUE_FLAG_NONE, + 0, + }; + + DmlObjects dml_objects; + + THROW_IF_FAILED(D3D12CreateDevice(nullptr, D3D_FEATURE_LEVEL_11_0, IID_PPV_ARGS(&dml_objects.d3d12_device))); + THROW_IF_FAILED(DMLCreateDevice(dml_objects.d3d12_device.Get(), DML_CREATE_DEVICE_FLAG_NONE, IID_PPV_ARGS(&dml_objects.dml_device))); + THROW_IF_FAILED(dml_objects.d3d12_device->CreateCommandQueue(&command_queue_description, IID_PPV_ARGS(&dml_objects.command_queue))); + THROW_IF_FAILED(dml_objects.d3d12_device->CreateCommandAllocator(D3D12_COMMAND_LIST_TYPE_DIRECT, IID_PPV_ARGS(&dml_objects.command_allocator))); + THROW_IF_FAILED(dml_objects.d3d12_device->CreateCommandList(0, D3D12_COMMAND_LIST_TYPE_DIRECT, dml_objects.command_allocator.Get(), nullptr, IID_PPV_ARGS(&dml_objects.command_list))); + + return dml_objects; +} + +static ComPtr CreateD3D12ResourceOfByteSize( + ID3D12Device* d3d_device, + size_t resource_byte_size, + D3D12_HEAP_TYPE heap_type = D3D12_HEAP_TYPE_DEFAULT, + D3D12_RESOURCE_STATES resource_state = D3D12_RESOURCE_STATE_COMMON, + D3D12_RESOURCE_FLAGS resource_flags = D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS) { + resource_byte_size = std::max(resource_byte_size, size_t(DML_MINIMUM_BUFFER_TENSOR_ALIGNMENT)); + + // DML needs the resources' sizes to be a multiple of 4 bytes + (resource_byte_size += 3) &= ~3; + + D3D12_HEAP_PROPERTIES heap_properties = {}; + heap_properties.Type = heap_type; + heap_properties.CPUPageProperty = D3D12_CPU_PAGE_PROPERTY_UNKNOWN; + heap_properties.MemoryPoolPreference = D3D12_MEMORY_POOL_UNKNOWN; + heap_properties.CreationNodeMask = 1; + heap_properties.VisibleNodeMask = 1; + + D3D12_RESOURCE_DESC resource_desc = {}; + resource_desc.Dimension = D3D12_RESOURCE_DIMENSION_BUFFER; + resource_desc.Alignment = 0; + resource_desc.Width = static_cast(resource_byte_size); + resource_desc.Height = 1; + resource_desc.DepthOrArraySize = 1; + resource_desc.MipLevels = 1; + resource_desc.Format = DXGI_FORMAT_UNKNOWN; + resource_desc.SampleDesc = {1, 0}; + resource_desc.Layout = D3D12_TEXTURE_LAYOUT_ROW_MAJOR; + resource_desc.Flags = resource_flags; + + ComPtr gpu_resource; + THROW_IF_FAILED(d3d_device->CreateCommittedResource( + &heap_properties, + D3D12_HEAP_FLAG_NONE, + &resource_desc, + resource_state, + nullptr, + IID_PPV_ARGS(&gpu_resource))); + + return gpu_resource; +} + +static void WaitForQueueToComplete(ID3D12CommandQueue* queue) { + ComPtr device; + THROW_IF_FAILED(queue->GetDevice(IID_PPV_ARGS(device.GetAddressOf()))); + ComPtr fence; + THROW_IF_FAILED(device->CreateFence(0, D3D12_FENCE_FLAG_NONE, IID_PPV_ARGS(fence.GetAddressOf()))); + THROW_IF_FAILED(queue->Signal(fence.Get(), 1)); + THROW_IF_FAILED(fence->SetEventOnCompletion(1, nullptr)); +} + +static void UploadDataToDml( + DmlObjects& dml_objects, + ID3D12Resource* destination_resource, + gsl::span source_data) { + // Get the size of the resource. + D3D12_RESOURCE_DESC resource_desc = destination_resource->GetDesc(); + assert(resource_desc.Dimension == D3D12_RESOURCE_DIMENSION_BUFFER); + const size_t data_size_in_bytes = static_cast(resource_desc.Width); + + // Create intermediate upload resource visible to both CPU and GPU. + dml_objects.upload_buffer = CreateD3D12ResourceOfByteSize(dml_objects.d3d12_device.Get(), data_size_in_bytes, D3D12_HEAP_TYPE_UPLOAD, D3D12_RESOURCE_STATE_GENERIC_READ, D3D12_RESOURCE_FLAG_NONE); + + // Copy CPU-side data to shared memory that is both CPU and GPU visible. + size_t clamped_data_byte_size = std::min(data_size_in_bytes, source_data.size()); + std::byte* upload_buffer_data = nullptr; + THROW_IF_FAILED(dml_objects.upload_buffer->Map(0, nullptr, reinterpret_cast(&upload_buffer_data))); + memcpy(upload_buffer_data, source_data.data(), clamped_data_byte_size); + dml_objects.upload_buffer->Unmap(0, nullptr); + + D3D12_RESOURCE_BARRIER resource_barrier{}; + resource_barrier.Type = D3D12_RESOURCE_BARRIER_TYPE_TRANSITION; + resource_barrier.Flags = D3D12_RESOURCE_BARRIER_FLAG_NONE; + resource_barrier.Transition = {}; + resource_barrier.Transition.pResource = destination_resource; + resource_barrier.Transition.Subresource = D3D12_RESOURCE_BARRIER_ALL_SUBRESOURCES; + resource_barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_COPY_DEST; + resource_barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_UNORDERED_ACCESS; + + // Issue deferred command to copy from the intermediate shared resource to the final GPU resource, + // and then execute the commands. + dml_objects.command_list->CopyResource(destination_resource, dml_objects.upload_buffer.Get()); + dml_objects.command_list->ResourceBarrier(1, &resource_barrier); + THROW_IF_FAILED(dml_objects.command_list->Close()); + ID3D12CommandList* command_lists[] = {dml_objects.command_list.Get()}; + dml_objects.command_queue->ExecuteCommandLists(ARRAYSIZE(command_lists), command_lists); + + THROW_IF_FAILED(dml_objects.command_allocator->Reset()); + THROW_IF_FAILED(dml_objects.command_list->Reset(dml_objects.command_allocator.Get(), nullptr)); +} + +static void DownloadDataFromDml( + DmlObjects& dml_objects, + ID3D12Resource* source_resource, + gsl::span destination_data) { + // Get the size of the resource. + D3D12_RESOURCE_DESC resource_desc = source_resource->GetDesc(); + assert(resource_desc.Dimension == D3D12_RESOURCE_DIMENSION_BUFFER); + const size_t data_size_in_bytes = static_cast(resource_desc.Width); + + // Create intermediate upload resource visible to both CPU and GPU. + ComPtr download_buffer = CreateD3D12ResourceOfByteSize(dml_objects.d3d12_device.Get(), data_size_in_bytes, D3D12_HEAP_TYPE_READBACK, D3D12_RESOURCE_STATE_COPY_DEST, D3D12_RESOURCE_FLAG_NONE); + + D3D12_RESOURCE_BARRIER resource_barrier = {}; + resource_barrier.Type = D3D12_RESOURCE_BARRIER_TYPE_TRANSITION; + resource_barrier.Flags = D3D12_RESOURCE_BARRIER_FLAG_NONE; + resource_barrier.Transition = {}; + resource_barrier.Transition.pResource = source_resource; + resource_barrier.Transition.Subresource = D3D12_RESOURCE_BARRIER_ALL_SUBRESOURCES; + resource_barrier.Transition.StateBefore = D3D12_RESOURCE_STATE_UNORDERED_ACCESS; + resource_barrier.Transition.StateAfter = D3D12_RESOURCE_STATE_COPY_SOURCE; + + // Copy GPU data into the download buffer. + dml_objects.command_list->ResourceBarrier(1, &resource_barrier); + dml_objects.command_list->CopyResource(download_buffer.Get(), source_resource); + THROW_IF_FAILED(dml_objects.command_list->Close()); + ID3D12CommandList* command_lists[] = {dml_objects.command_list.Get()}; + dml_objects.command_queue->ExecuteCommandLists(static_cast(std::size(command_lists)), command_lists); + WaitForQueueToComplete(dml_objects.command_queue.Get()); + THROW_IF_FAILED(dml_objects.command_allocator->Reset()); + THROW_IF_FAILED(dml_objects.command_list->Reset(dml_objects.command_allocator.Get(), nullptr)); + + // Copy from shared GPU/CPU memory to ordinary system RAM. + size_t clamped_data_byte_size = std::min(data_size_in_bytes, destination_data.size()); + std::byte* sourceData = nullptr; + D3D12_RANGE range = {0, clamped_data_byte_size}; + THROW_IF_FAILED(download_buffer->Map(0, &range, reinterpret_cast(&sourceData))); + memcpy(destination_data.data(), sourceData, clamped_data_byte_size); + download_buffer->Unmap(0, nullptr); +} + +Ort::Value CreateTensorValueFromExistingD3DResource( + const OrtDmlApi& ort_dml_api, + Ort::MemoryInfo const& memory_information, + ID3D12Resource* d3d_resource, + gsl::span tensor_dimensions, + ONNXTensorElementDataType element_data_type, + /*out*/ void** dml_ep_resource_wrapper // Must stay alive with Ort::Value. +) { + *dml_ep_resource_wrapper = nullptr; + + void* dml_allocator_resource; + Ort::ThrowOnError(ort_dml_api.CreateGPUAllocationFromD3DResource(d3d_resource, &dml_allocator_resource)); + auto deleter = [&](void*) { ort_dml_api.FreeGPUAllocation(dml_allocator_resource); }; + std::unique_ptr> dml_allocator_resource_cleanup(dml_allocator_resource, deleter); + + size_t tensor_byte_size = static_cast(d3d_resource->GetDesc().Width); + Ort::Value new_value( + Ort::Value::CreateTensor( + memory_information, + dml_allocator_resource, + tensor_byte_size, + tensor_dimensions.data(), + tensor_dimensions.size(), + element_data_type)); + + // Return values and the wrapped resource. + *dml_ep_resource_wrapper = dml_allocator_resource; + dml_allocator_resource_cleanup.release(); + + return new_value; +} + +#endif + template static void TestInference(Ort::Env& env, const std::basic_string& model_uri, const std::vector& inputs, @@ -180,7 +384,7 @@ static void TestInference(Ort::Env& env, const std::basic_string& mod } static constexpr PATH_TYPE MODEL_URI = TSTR("testdata/mul_1.onnx"); -#if defined(USE_CUDA) +#if defined(USE_CUDA) || defined(USE_DML) static constexpr PATH_TYPE CUDA_GRAPH_ANNOTATION_MODEL_URI = TSTR("testdata/mul_1_dynamic.onnx"); #endif static constexpr PATH_TYPE MATMUL_MODEL_URI = TSTR("testdata/matmul_1.onnx"); @@ -1458,9 +1662,12 @@ TEST(CApiTest, test_custom_op_library) { #elif USE_ROCM TestInference(*ort_env, CUSTOM_OP_LIBRARY_TEST_MODEL_URI, inputs, "output", expected_dims_y, expected_values_y, 3, nullptr, lib_name.c_str()); -#else +#elif USE_DML TestInference(*ort_env, CUSTOM_OP_LIBRARY_TEST_MODEL_URI, inputs, "output", expected_dims_y, - expected_values_y, 0, nullptr, lib_name.c_str()); + expected_values_y, 4, nullptr, lib_name.c_str()); +#else +TestInference(*ort_env, CUSTOM_OP_LIBRARY_TEST_MODEL_URI, inputs, "output", expected_dims_y, + expected_values_y, 0, nullptr, lib_name.c_str()); #endif } @@ -1966,7 +2173,7 @@ TEST(CApiTest, io_binding_cuda) { } #endif -#if defined(USE_CUDA) || defined(USE_TENSORRT) || defined(USE_ROCM) +#if defined(USE_CUDA) || defined(USE_TENSORRT) || defined(USE_ROCM) || defined(USE_DML) TEST(CApiTest, basic_cuda_graph) { const auto& api = Ort::GetApi(); Ort::SessionOptions session_options; @@ -2010,6 +2217,14 @@ TEST(CApiTest, basic_cuda_graph) { ASSERT_TRUE(api.SessionOptionsAppendExecutionProvider_ROCM( static_cast(session_options), rel_rocm_options.get()) == nullptr); +#elif defined(USE_DML) + // Enable dynamic DML graph in DML provider option. + session_options.AddConfigEntry("ep.dml.enable_graph_capture", "1"); + const OrtDmlApi* ort_dml_api; + Ort::ThrowOnError(api.GetExecutionProviderApi("DML", ORT_API_VERSION, reinterpret_cast(&ort_dml_api))); + + auto dml_objects = CreateDmlObjects(); + ort_dml_api->SessionOptionsAppendExecutionProvider_DML1(session_options, dml_objects.dml_device.Get(), dml_objects.command_queue.Get()); #endif Ort::Session session(*ort_env, MODEL_URI, session_options); @@ -2019,6 +2234,8 @@ TEST(CApiTest, basic_cuda_graph) { #define cudaMemcpyHostToDevice hipMemcpyHostToDevice #define cudaMemcpyDeviceToHost hipMemcpyDeviceToHost Ort::MemoryInfo info_mem("Hip", OrtAllocatorType::OrtArenaAllocator, 0, OrtMemTypeDefault); +#elif defined(USE_DML) + Ort::MemoryInfo info_mem("DML", OrtAllocatorType::OrtDeviceAllocator, 0, OrtMemTypeDefault); #else Ort::MemoryInfo info_mem("Cuda", OrtAllocatorType::OrtArenaAllocator, 0, OrtMemTypeDefault); #endif @@ -2032,7 +2249,14 @@ TEST(CApiTest, basic_cuda_graph) { auto input_data = allocator.GetAllocation(x_values.size() * sizeof(float)); ASSERT_NE(input_data.get(), nullptr); + +#ifdef USE_DML + ComPtr input_resource; + Ort::ThrowOnError(ort_dml_api->GetD3D12ResourceFromAllocation(allocator, input_data.get(), &input_resource)); + UploadDataToDml(dml_objects, input_resource.Get(), gsl::make_span(reinterpret_cast(x_values.data()), sizeof(float) * x_values.size())); +#else (void)cudaMemcpy(input_data.get(), x_values.data(), sizeof(float) * x_values.size(), cudaMemcpyHostToDevice); +#endif // Create an OrtValue tensor backed by data on CUDA memory Ort::Value bound_x = Ort::Value::CreateTensor(info_mem, reinterpret_cast(input_data.get()), x_values.size(), @@ -2058,21 +2282,48 @@ TEST(CApiTest, basic_cuda_graph) { // Check the values against the bound raw memory (needs copying from device to host first) std::array y_values; + +#ifdef USE_DML + ComPtr output_resource; + Ort::ThrowOnError(ort_dml_api->GetD3D12ResourceFromAllocation(allocator, output_data.get(), &output_resource)); + auto output_cpu_bytes = reinterpret_cast(y_values.data()); + DownloadDataFromDml(dml_objects, output_resource.Get(), gsl::make_span(output_cpu_bytes, sizeof(float) * y_values.size())); +#else (void)cudaMemcpy(y_values.data(), output_data.get(), sizeof(float) * y_values.size(), cudaMemcpyDeviceToHost); +#endif + ASSERT_THAT(y_values, ::testing::ContainerEq(expected_y)); // Replay the captured CUDA graph session.Run(Ort::RunOptions(), binding); + +#ifdef USE_DML + DownloadDataFromDml(dml_objects, output_resource.Get(), gsl::make_span(output_cpu_bytes, sizeof(float) * y_values.size())); +#else (void)cudaMemcpy(y_values.data(), output_data.get(), sizeof(float) * y_values.size(), cudaMemcpyDeviceToHost); +#endif + ASSERT_THAT(y_values, ::testing::ContainerEq(expected_y)); // Change the input and replay the CUDA graph again. x_values = {10.0f, 20.0f, 30.0f, 40.0f, 50.0f, 60.0f}; + +#ifdef USE_DML + UploadDataToDml(dml_objects, input_resource.Get(), gsl::make_span(reinterpret_cast(x_values.data()), sizeof(float) * x_values.size())); +#else (void)cudaMemcpy(input_data.get(), x_values.data(), sizeof(float) * x_values.size(), cudaMemcpyHostToDevice); +#endif + binding.SynchronizeInputs(); session.Run(Ort::RunOptions(), binding); + +#ifdef USE_DML + DownloadDataFromDml(dml_objects, output_resource.Get(), gsl::make_span(output_cpu_bytes, sizeof(float) * y_values.size())); +#else (void)cudaMemcpy(y_values.data(), output_data.get(), sizeof(float) * y_values.size(), cudaMemcpyDeviceToHost); +#endif + expected_y = {10.0f, 40.0f, 90.0f, 160.0f, 250.0f, 360.0f}; ASSERT_THAT(y_values, ::testing::ContainerEq(expected_y)); @@ -2086,7 +2337,7 @@ TEST(CApiTest, basic_cuda_graph) { #endif } -#if defined(USE_CUDA) +#if defined(USE_CUDA) || defined(USE_DML) struct CudaGraphInputOutputData_0 { const std::array x_shape = {3, 2}; std::array x_values = {1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f}; @@ -2124,13 +2375,28 @@ template static void RunWithCudaGraphAnnotation(T& cg_data, Ort::Session& session, Ort::MemoryInfo& info_mem, +#ifdef USE_DML + DmlObjects& dml_objects, +#endif Ort::MemoryAllocation& input_data, Ort::MemoryAllocation& output_data, const char* cuda_graph_annotation) { +#ifdef USE_DML + Ort::SessionOptions session_options; + Ort::Allocator allocator(session, info_mem); + const auto& api = Ort::GetApi(); + const OrtDmlApi* ort_dml_api; + Ort::ThrowOnError(api.GetExecutionProviderApi("DML", ORT_API_VERSION, reinterpret_cast(&ort_dml_api))); + + ComPtr input_resource; + Ort::ThrowOnError(ort_dml_api->GetD3D12ResourceFromAllocation(allocator, input_data.get(), &input_resource)); + UploadDataToDml(dml_objects, input_resource.Get(), gsl::make_span(reinterpret_cast(cg_data.x_values.data()), sizeof(float) * cg_data.x_values.size())); +#else (void)cudaMemcpy(input_data.get(), cg_data.x_values.data(), sizeof(float) * cg_data.x_values.size(), cudaMemcpyHostToDevice); +#endif // Create an OrtValue tensor backed by data on CUDA memory Ort::Value bound_x = Ort::Value::CreateTensor(info_mem, @@ -2160,32 +2426,59 @@ static void RunWithCudaGraphAnnotation(T& cg_data, session.Run(run_option, binding); // Check the values against the bound raw memory (needs copying from device to host first) +#ifdef USE_DML + ComPtr output_resource; + Ort::ThrowOnError(ort_dml_api->GetD3D12ResourceFromAllocation(allocator, output_data.get(), &output_resource)); + auto output_cpu_bytes = reinterpret_cast(cg_data.y_values.data()); + DownloadDataFromDml(dml_objects, output_resource.Get(), gsl::make_span(output_cpu_bytes, sizeof(float) * cg_data.y_values.size())); +#else (void)cudaMemcpy(cg_data.y_values.data(), output_data.get(), sizeof(float) * cg_data.y_values.size(), cudaMemcpyDeviceToHost); +#endif + ASSERT_THAT(cg_data.y_values, ::testing::ContainerEq(cg_data.expected_y)); // Replay the captured CUDA graph session.Run(run_option, binding); + +#ifdef USE_DML + DownloadDataFromDml(dml_objects, output_resource.Get(), gsl::make_span(output_cpu_bytes, sizeof(float) * cg_data.y_values.size())); +#else (void)cudaMemcpy(cg_data.y_values.data(), output_data.get(), sizeof(float) * cg_data.y_values.size(), cudaMemcpyDeviceToHost); +#endif ASSERT_THAT(cg_data.y_values, ::testing::ContainerEq(cg_data.expected_y)); // Change the input and replay the CUDA graph again. +#ifdef USE_DML + UploadDataToDml(dml_objects, + input_resource.Get(), + gsl::make_span(reinterpret_cast(cg_data.new_x_values.data()), + sizeof(float) * cg_data.new_x_values.size())); +#else (void)cudaMemcpy(input_data.get(), cg_data.new_x_values.data(), sizeof(float) * cg_data.new_x_values.size(), cudaMemcpyHostToDevice); +#endif + binding.SynchronizeInputs(); session.Run(run_option, binding); + +#ifdef USE_DML + DownloadDataFromDml(dml_objects, output_resource.Get(), gsl::make_span(output_cpu_bytes, sizeof(float) * cg_data.y_values.size())); +#else (void)cudaMemcpy(cg_data.y_values.data(), output_data.get(), sizeof(float) * cg_data.y_values.size(), cudaMemcpyDeviceToHost); +#endif + ASSERT_THAT(cg_data.y_values, ::testing::ContainerEq(cg_data.new_expected_y)); // Clean up @@ -2197,6 +2490,15 @@ TEST(CApiTest, basic_cuda_graph_with_annotation) { const auto& api = Ort::GetApi(); Ort::SessionOptions session_options; +#ifdef USE_DML + session_options.AddConfigEntry("ep.dml.enable_graph_capture", "1"); + const OrtDmlApi* ort_dml_api; + Ort::ThrowOnError(api.GetExecutionProviderApi("DML", ORT_API_VERSION, reinterpret_cast(&ort_dml_api))); + auto dml_objects = CreateDmlObjects(); + ort_dml_api->SessionOptionsAppendExecutionProvider_DML1(session_options, dml_objects.dml_device.Get(), dml_objects.command_queue.Get()); + + Ort::MemoryInfo info_mem("DML", OrtAllocatorType::OrtDeviceAllocator, 0, OrtMemTypeDefault); +#else // Enable cuda graph in cuda provider option. OrtCUDAProviderOptionsV2* cuda_options = nullptr; ASSERT_TRUE(api.CreateCUDAProviderOptions(&cuda_options) == nullptr); @@ -2209,9 +2511,10 @@ TEST(CApiTest, basic_cuda_graph_with_annotation) { ASSERT_TRUE(api.SessionOptionsAppendExecutionProvider_CUDA_V2( static_cast(session_options), rel_cuda_options.get()) == nullptr); + Ort::MemoryInfo info_mem("Cuda", OrtAllocatorType::OrtArenaAllocator, 0, OrtMemTypeDefault); +#endif Ort::Session session(*ort_env, CUDA_GRAPH_ANNOTATION_MODEL_URI, session_options); - Ort::MemoryInfo info_mem("Cuda", OrtAllocatorType::OrtArenaAllocator, 0, OrtMemTypeDefault); Ort::Allocator allocator(session, info_mem); auto allocator_info = allocator.GetInfo(); @@ -2226,9 +2529,15 @@ TEST(CApiTest, basic_cuda_graph_with_annotation) { ASSERT_NE(input_data.get(), nullptr); ASSERT_NE(output_data.get(), nullptr); +#ifdef USE_DML + RunWithCudaGraphAnnotation(cg_data_0, session, info_mem, dml_objects, input_data, output_data, nullptr); + RunWithCudaGraphAnnotation(cg_data_1, session, info_mem, dml_objects, input_data, output_data, "1"); + RunWithCudaGraphAnnotation(cg_data_2, session, info_mem, dml_objects, input_data, output_data, "2"); +#else RunWithCudaGraphAnnotation(cg_data_0, session, info_mem, input_data, output_data, nullptr); RunWithCudaGraphAnnotation(cg_data_1, session, info_mem, input_data, output_data, "1"); RunWithCudaGraphAnnotation(cg_data_2, session, info_mem, input_data, output_data, "2"); +#endif } #endif @@ -2280,6 +2589,22 @@ TEST(CApiTest, hip_graph_with_shape_nodes) { } #endif // defined(USE_ROCM) +#if defined(USE_DML) +TEST(CApiTest, dml_graph_with_shape_nodes) { + const auto& api = Ort::GetApi(); + + // Enable hip graph in rocm provider option. + const OrtDmlApi* ort_dml_api; + Ort::SessionOptions session_options; + session_options.AddConfigEntry("ep.dml.enable_graph_capture", "1"); + Ort::ThrowOnError(api.GetExecutionProviderApi("DML", ORT_API_VERSION, reinterpret_cast(&ort_dml_api))); + ort_dml_api->SessionOptionsAppendExecutionProvider_DML(session_options, 0); + + // Successful loading of the ONNX model with shape nodes with dml graph feature enabled + Ort::Session session(*ort_env, TSTR("testdata/cuda_graph_with_shape_nodes.onnx"), session_options); +} +#endif // defined(USE_DML) + #endif // REDUCED_OPS_BUILD #endif // defined(USE_CUDA) || defined(USE_TENSORRT) || defined(USE_ROCM) diff --git a/onnxruntime/test/util/default_providers.cc b/onnxruntime/test/util/default_providers.cc index 46230b61db..e353cc73b2 100644 --- a/onnxruntime/test/util/default_providers.cc +++ b/onnxruntime/test/util/default_providers.cc @@ -299,7 +299,7 @@ std::unique_ptr DefaultCannExecutionProvider() { std::unique_ptr DefaultDmlExecutionProvider() { #ifdef USE_DML ConfigOptions config_options{}; - if (auto factory = DMLProviderFactoryCreator::CreateFromDeviceOptions(config_options, nullptr, false, false, false)) { + if (auto factory = DMLProviderFactoryCreator::CreateFromDeviceOptions(config_options, nullptr, false, false)) { return factory->CreateProvider(); } #endif diff --git a/tools/ci_build/build.py b/tools/ci_build/build.py index 6d99ff07dd..d2119986dc 100644 --- a/tools/ci_build/build.py +++ b/tools/ci_build/build.py @@ -2089,6 +2089,10 @@ def run_onnxruntime_tests(args, source_dir, ctest_path, build_dir, configs): log.info("Testing CUDA Graph feature") run_subprocess([sys.executable, "onnxruntime_test_python_cudagraph.py"], cwd=cwd, dll_path=dll_path) + if args.use_dml: + log.info("Testing DML Graph feature") + run_subprocess([sys.executable, "onnxruntime_test_python_dmlgraph.py"], cwd=cwd, dll_path=dll_path) + if not args.disable_ml_ops and not args.use_tensorrt: run_subprocess([sys.executable, "onnxruntime_test_python_mlops.py"], cwd=cwd, dll_path=dll_path)