From 9435369550532989f6bcc009b1e928cca5fdf587 Mon Sep 17 00:00:00 2001 From: Justin Stoecker Date: Tue, 12 Apr 2022 11:59:00 -0700 Subject: [PATCH 01/19] Option to build with DML as an external project (#11180) --- cmake/external/dml.cmake | 48 ++++++++++++++++++++++++++++++- cmake/onnxruntime_providers.cmake | 11 +++++-- tools/ci_build/build.py | 14 +++++++-- 3 files changed, 66 insertions(+), 7 deletions(-) diff --git a/cmake/external/dml.cmake b/cmake/external/dml.cmake index f7da89f544..0f7730c872 100644 --- a/cmake/external/dml.cmake +++ b/cmake/external/dml.cmake @@ -1,6 +1,26 @@ # Copyright (c) Microsoft Corporation. All rights reserved. # Licensed under the MIT License. +# There are effectively three ways to consume DirectML in this repo: +# +# 1) Public = the build points at a pre-built copy of DirectML distributed as a NuGet package. +# 2) Custom = the build points at a local copy of DirectML (bin/, include/, lib/). The dml_INCLUDE_DIR and +# dml_LIB_DIR variables are also expected to be set to the custom build location. +# 3) Internal = the build points at the DirectML source repo and builds it as part of the main project. +# +# Build Type | onnxruntime_USE_CUSTOM_DIRECTML | dml_EXTERNAL_PROJECT +# -----------|---------------------------------|--------------------- +# Public | OFF | OFF +# Custom | ON | OFF +# Internal | ON | ON +# +# The "Public" build type is the default, and any mainline branches (e.g. master, rel-*) subject to CI +# should use the public build configuration. Topic branches can use the internal build type for testing, +# but they must be buildable with a public NuGet package before merging with a mainline branch. + +set(onnxruntime_USE_CUSTOM_DIRECTML OFF CACHE BOOL "Depend on a custom/internal build of DirectML.") +set(dml_EXTERNAL_PROJECT OFF CACHE BOOL "Build DirectML as a source dependency.") + if (NOT onnxruntime_USE_CUSTOM_DIRECTML) if (NOT(MSVC) OR NOT(WIN32)) message(FATAL_ERROR "NuGet packages are only supported for MSVC on Windows.") @@ -34,5 +54,31 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML) add_custom_target(RESTORE_PACKAGES ALL DEPENDS ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib) add_dependencies(RESTORE_PACKAGES nuget) else() - include_directories(${dml_INCLUDE_DIR}) + if (dml_EXTERNAL_PROJECT) + set(dml_preset_config $,debug,release>) + set(dml_preset_name ${onnxruntime_target_platform}-win-redist-${dml_preset_config}) + + include(ExternalProject) + ExternalProject_Add( + directml_repo + GIT_REPOSITORY https://dev.azure.com/microsoft/WindowsAI/_git/DirectML + GIT_TAG 2290bd6495fdf8c35822816213516d13f3742cc9 + GIT_SHALLOW OFF # not allowed when GIT_TAG is a commit SHA, which is preferred (it's stable, unlike branches) + GIT_PROGRESS ON + BUILD_IN_SOURCE ON + CONFIGURE_COMMAND ${CMAKE_COMMAND} --preset ${dml_preset_name} -DDML_BUILD_TESTS=OFF + BUILD_COMMAND ${CMAKE_COMMAND} --build --preset ${dml_preset_name} + INSTALL_COMMAND ${CMAKE_COMMAND} --install build/${dml_preset_name} + STEP_TARGETS install + ) + + # Target that consumers can use to link with the internal build of DirectML. + set(directml_install_path ${CMAKE_BINARY_DIR}/directml_repo-prefix/src/directml_repo/build/${dml_preset_name}/install) + add_library(DirectML INTERFACE) + target_link_libraries(DirectML INTERFACE ${directml_install_path}/lib/DirectML.lib) + add_dependencies(DirectML directml_repo-install) + include_directories(BEFORE ${directml_install_path}/include) + else() + include_directories(${dml_INCLUDE_DIR}) + endif() endif() diff --git a/cmake/onnxruntime_providers.cmake b/cmake/onnxruntime_providers.cmake index af18cd605a..fe22776205 100644 --- a/cmake/onnxruntime_providers.cmake +++ b/cmake/onnxruntime_providers.cmake @@ -1070,10 +1070,15 @@ if (onnxruntime_USE_DML) function(target_add_dml target) if (onnxruntime_USE_CUSTOM_DIRECTML) - if (dml_LIB_DIR) - target_link_libraries(${target} PRIVATE ${dml_LIB_DIR}/DirectML.lib) - else() + if (dml_EXTERNAL_PROJECT) + # Internal build of DirectML: link against the "DirectML" target. target_link_libraries(${target} PRIVATE DirectML) + else() + if (dml_LIB_DIR) + target_link_libraries(${target} PRIVATE ${dml_LIB_DIR}/DirectML.lib) + else() + target_link_libraries(${target} PRIVATE DirectML) + endif() endif() else() add_dependencies(${target} RESTORE_PACKAGES) diff --git a/tools/ci_build/build.py b/tools/ci_build/build.py index 432f2d0917..4b53a401c6 100644 --- a/tools/ci_build/build.py +++ b/tools/ci_build/build.py @@ -491,6 +491,8 @@ def parse_arguments(): parser.add_argument( "--dml_path", type=str, default="", help="Path to a custom DirectML installation (must have bin/, lib/, and include/ subdirectories).") + parser.add_argument( + "--dml_external_project", action='store_true', help="Build with DirectML as an external project.") parser.add_argument( "--use_winml", action='store_true', help="Build with WinML.") parser.add_argument( @@ -992,6 +994,12 @@ def generate_build_tree(cmake_path, source_dir, build_dir, cuda_home, cudnn_home "-Ddml_LIB_DIR=" + os.path.join(args.dml_path, "lib"), ] + if args.dml_external_project: + cmake_args += [ + "-Donnxruntime_USE_CUSTOM_DIRECTML=ON", + "-Ddml_EXTERNAL_PROJECT=ON", + ] + if args.use_gdk: cmake_args += [ "-DCMAKE_TOOLCHAIN_FILE=" + os.path.join(source_dir, 'cmake', 'gdk_toolchain.cmake'), @@ -999,8 +1007,8 @@ def generate_build_tree(cmake_path, source_dir, build_dir, cuda_home, cudnn_home "-DGDK_PLATFORM=" + args.gdk_platform, "-Donnxruntime_BUILD_UNIT_TESTS=OFF" # gtest doesn't build for GDK ] - if args.use_dml and not args.dml_path: - raise BuildError("You must set dml_path when building with the GDK.") + if args.use_dml and not (args.dml_path or args.dml_external_project): + raise BuildError("You must set dml_path or dml_external_project when building with the GDK.") if is_macOS() and not args.android: cmake_args += ["-DCMAKE_OSX_ARCHITECTURES=" + args.osx_arch] @@ -1333,7 +1341,7 @@ def setup_dml_build(args, cmake_path, build_dir, configs): raise BuildError("dml_path is invalid.", "dml_path='{}' expected_file='{}'." .format(args.dml_path, file_path)) - else: + elif not args.dml_external_project: for config in configs: # Run the RESTORE_PACKAGES target to perform the initial # NuGet setup. From 205b61c5d84989e4f8a538e7db125aee69c23960 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Tue, 10 May 2022 17:17:55 -0700 Subject: [PATCH 02/19] Fix bad merge in build.py --- tools/ci_build/build.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/tools/ci_build/build.py b/tools/ci_build/build.py index bcbcfdb328..fb1cecb8d8 100644 --- a/tools/ci_build/build.py +++ b/tools/ci_build/build.py @@ -1362,10 +1362,6 @@ def setup_dml_build(args, cmake_path, build_dir, configs): "dml_path='{}' expected_file='{}'." .format(args.dml_path, file_path)) elif not args.dml_external_project: - raise BuildError( - "dml_path is invalid.", "dml_path='{}' expected_file='{}'.".format(args.dml_path, file_path) - ) - else: for config in configs: # Run the RESTORE_PACKAGES target to perform the initial # NuGet setup. From 2660eb836427807d843efd563366450af54bda72 Mon Sep 17 00:00:00 2001 From: sumitsays Date: Wed, 11 May 2022 16:24:11 -0700 Subject: [PATCH 03/19] DML EP: Gelu (#11483) Co-authored-by: Sumit Agarwal --- .../src/External/DirectMLHelpers/ApiTraits.h | 18 ++++++++++++++++-- .../External/DirectMLHelpers/DirectMLSchema.h | 13 +++++++++++++ .../DirectMLHelpers/GeneratedSchemaHelpers.h | 13 ++++++++++++- .../src/Operators/DmlOperatorActivation.cpp | 2 ++ .../src/Operators/OperatorRegistration.cpp | 2 ++ .../dml/OperatorAuthorHelper/OperatorHelper.h | 1 + .../OperatorAuthorHelper/OperatorVersions.h | 1 + 7 files changed, 47 insertions(+), 3 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h index dd00eb4a80..1774a8a2b0 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h @@ -24,8 +24,8 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 153; - static constexpr size_t ActivationFunctionCount = 20; + static constexpr auto ValueCount = 154; + static constexpr size_t ActivationFunctionCount = 21; }; template <> @@ -1113,6 +1113,12 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_SHRINK; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_GELU; +}; + template struct OperatorTypeTraits @@ -2055,6 +2061,12 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_SHRINK> using DescType = DML_ACTIVATION_SHRINK_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_GELU> +{ + using DescType = DML_ACTIVATION_GELU_OPERATOR_DESC; +}; + // Calls a visitor functor, supplying an empty operator desc corresponding to the given DML_OPERATOR_TYPE as // the first argument. // @@ -2382,6 +2394,8 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_ACTIVATION_THRESHOLDED_RELU_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_SHRINK: return std::invoke(std::forward(visitor), DML_ACTIVATION_SHRINK_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ACTIVATION_GELU: + return std::invoke(std::forward(visitor), DML_ACTIVATION_GELU_OPERATOR_DESC{}, std::forward(args)...); default: ORT_THROW_HR(E_INVALIDARG); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h index 32d1eda07e..137ea2b030 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h @@ -2525,6 +2525,19 @@ constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA { DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_ACTIVATION_GELU_OPERATOR_SCHEMA_FIELDS[2] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_GELU_OPERATOR_SCHEMA { + "DML_OPERATOR_ACTIVATION_GELU", + DML_OPERATOR_ACTIVATION_GELU, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_IN_PLACE_EXECUTION, + 2, + DML_ACTIVATION_GELU_OPERATOR_SCHEMA_FIELDS, +}; + constexpr DML_SCHEMA_FIELD DML_RNN_ZERO_OPERATOR_SCHEMA_FIELDS[3] { DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "SequenceLengthsTensor", false }, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h index 227c6aa46c..e0a705330f 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h @@ -1539,7 +1539,13 @@ inline std::vector GetFields(const DML_ACTIVATION_SHRINK_OPERATOR OperatorField(&DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Threshold))), }; } - +inline std::vector GetFields(const DML_ACTIVATION_GELU_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ACTIVATION_GELU_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ACTIVATION_GELU_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + }; +} inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) { switch (operatorType) @@ -1700,6 +1706,7 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_ACTIVATION_TANH: return DML_ACTIVATION_TANH_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_THRESHOLDED_RELU: return DML_ACTIVATION_THRESHOLDED_RELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_SHRINK: return DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA; + case DML_OPERATOR_ACTIVATION_GELU: return DML_ACTIVATION_GELU_OPERATOR_SCHEMA; default: ORT_THROW_HR(E_INVALIDARG); @@ -2337,6 +2344,10 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ACTIVATION_GELU: + return AbstractOperatorDesc( + &DML_ACTIVATION_GELU_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); default: ORT_THROW_HR(E_INVALIDARG); return AbstractOperatorDesc( diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp index c0d08f83d0..c7c9e50c90 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp @@ -91,6 +91,7 @@ public: case DML_OPERATOR_ACTIVATION_SIGMOID: case DML_OPERATOR_ACTIVATION_TANH: case DML_OPERATOR_ACTIVATION_SOFTSIGN: + case DML_OPERATOR_ACTIVATION_GELU: // No additional parameters to set. break; @@ -169,5 +170,6 @@ DML_OP_DEFINE_CREATION_FUNCTION(Softmax, DmlOperatorActivationTempla DML_OP_DEFINE_CREATION_FUNCTION(LogSoftmax, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(Hardmax, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(Shrink, DmlOperatorActivationTemplate); +DML_OP_DEFINE_CREATION_FUNCTION(Gelu, DmlOperatorActivationTemplate); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index edf3d3467d..b072d1509b 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -222,6 +222,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(Atanh); DML_OP_EXTERN_CREATION_FUNCTION(Erf); DML_OP_EXTERN_CREATION_FUNCTION(Where); DML_OP_EXTERN_CREATION_FUNCTION(Shrink); +DML_OP_EXTERN_CREATION_FUNCTION(Gelu); DML_OP_EXTERN_CREATION_FUNCTION(OneHot); DML_OP_EXTERN_CREATION_FUNCTION(EyeLike); DML_OP_EXTERN_CREATION_FUNCTION(MaxUnpool); @@ -627,6 +628,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, ParametricSoftplus, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Dropout, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 9, Shrink, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported)}, + {REG_INFO_MS( 1, Gelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // Uncategorized {REG_INFO( 7, MatMul, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index 32c34a19fa..7a3dd34760 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -1465,6 +1465,7 @@ using ShapeInferenceHelper_Softplus = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_ParametricSoftplus = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Dropout = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Shrink = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_Gelu = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Identity7 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Identity13 = GetOutputShapeAsInputShapeHelper; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h index 303f3ffe20..4ba4f944bd 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h @@ -353,6 +353,7 @@ namespace OperatorHelper static const int sc_sinceVer_DequantizeLinear = 1; static const int sc_sinceVer_ConvTransposeWithDynamicPads = 1; static const int sc_sinceVer_QLinearAdd = 1; + static const int sc_sinceVer_Gelu = 1; } // namespace MsftOperatorSet1 } // namespace OperatorHelper From 3a867d83d5c8af4d308b7e4212f99d96613a7ae0 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 18 May 2022 17:34:15 -0700 Subject: [PATCH 04/19] DirectML opset14 type updates for Add/Sub/Mul/Div and Relu/PRelu (#11560) * Add more type support for Add/Sub/Mul/Div and Relu/PRelu * Remove stale Remap64bitDmlDataTypeTo32bit --- .../DmlExecutionProvider/src/DmlCommon.cpp | 10 ---- .../dml/DmlExecutionProvider/src/DmlCommon.h | 1 - .../src/GraphPartitioner.cpp | 15 +---- .../src/GraphPartitioner.h | 4 +- .../src/GraphTransformer.cpp | 6 +- .../src/Operators/DmlOperatorCumSum.cpp | 3 +- .../src/Operators/OperatorRegistration.cpp | 60 +++++++++++++++---- .../src/Operators/OperatorUtility.cpp | 2 + .../DmlExecutionProvider/src/TensorDesc.cpp | 56 ----------------- .../dml/DmlExecutionProvider/src/TensorDesc.h | 1 - .../dml/OperatorAuthorHelper/OperatorHelper.h | 4 +- .../OperatorAuthorHelper/OperatorVersions.h | 11 ++++ 12 files changed, 71 insertions(+), 102 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.cpp index c4d3d60717..43204be90f 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.cpp @@ -30,16 +30,6 @@ DML_TENSOR_DATA_TYPE GetDmlDataTypeFromMlDataTypeNoThrow(MLOperatorTensorDataTyp }; } -DML_TENSOR_DATA_TYPE Remap64bitDmlDataTypeTo32bit(DML_TENSOR_DATA_TYPE dmlElementType) noexcept -{ - switch (dmlElementType) - { - case DML_TENSOR_DATA_TYPE_UINT64: return DML_TENSOR_DATA_TYPE_UINT32; break; - case DML_TENSOR_DATA_TYPE_INT64: return DML_TENSOR_DATA_TYPE_INT32; break; - default: return dmlElementType; - } -} - bool IsSigned(DML_TENSOR_DATA_TYPE dataType) { switch (dataType) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h index e43b982364..0c3abcd255 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h @@ -14,7 +14,6 @@ namespace Dml DML_TENSOR_DATA_TYPE GetDmlDataTypeFromMlDataType(MLOperatorTensorDataType tensorDataType); DML_TENSOR_DATA_TYPE GetDmlDataTypeFromMlDataTypeNoThrow(MLOperatorTensorDataType tensorDataType) noexcept; - DML_TENSOR_DATA_TYPE Remap64bitDmlDataTypeTo32bit(DML_TENSOR_DATA_TYPE dmlElementType) noexcept; MLOperatorTensorDataType GetMlDataTypeFromDmlDataType(DML_TENSOR_DATA_TYPE tensorDataType); size_t ComputeByteSizeFromDimensions(gsl::span dimensions, MLOperatorTensorDataType tensorDataType); size_t ComputeByteSizeFromTensor(IMLOperatorTensor& tensor); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp index 9a9c4b81bb..ad52d4056c 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp @@ -164,14 +164,10 @@ namespace Dml bool DoesNodeContainSupportedDataTypes( const onnxruntime::Node& node, - bool allow64BitInputThroughStrides, - _In_opt_ const std::unordered_map* nodeNameToPartitionMap, // Only used when allow64BitInputThroughStrides is true _In_opt_ const InternalRegistrationInfo* regInfo, uint32_t supportedDeviceDataTypeMask // Each bit corresponds to each DML_TENSOR_DATA_TYPE. ) { - ORT_THROW_HR_IF(E_INVALIDARG, allow64BitInputThroughStrides && !nodeNameToPartitionMap); - std::vector constantCpuInputs; if (regInfo != nullptr) @@ -253,13 +249,9 @@ namespace Dml const onnxruntime::Node& node, const onnxruntime::KernelRegistry& registry, uint32_t supportedDeviceDataTypeMask, // Each bit corresponds to each DML_TENSOR_DATA_TYPE. - const InternalRegistrationInfoMap& internalRegInfoMap, - bool allow64BitInputThroughStrides, - _In_opt_ const std::unordered_map* nodeNameToPartitionMap + const InternalRegistrationInfoMap& internalRegInfoMap ) { - ORT_THROW_HR_IF(E_INVALIDARG, allow64BitInputThroughStrides && !nodeNameToPartitionMap); - const onnxruntime::KernelCreateInfo* createInfo; Status st = registry.TryFindKernel(node, onnxruntime::kDmlExecutionProvider, &createInfo); if (!st.IsOK()) @@ -279,7 +271,7 @@ namespace Dml } // Check whether the node uses any data types which are unsupported by the device. - if (!DoesNodeContainSupportedDataTypes(node, allow64BitInputThroughStrides, nodeNameToPartitionMap, internalRegInfo.get(), supportedDeviceDataTypeMask)) + if (!DoesNodeContainSupportedDataTypes(node, internalRegInfo.get(), supportedDeviceDataTypeMask)) { return false; } @@ -309,8 +301,7 @@ namespace Dml // registration. Determine if that registration supports usage as a graph node. for (auto registry : dmlRegistries) { - bool allow64BitInputThroughStrides = true; - if (IsNodeSupportedByDml(node, *registry, supportedDeviceDataTypeMask, internalRegInfoMap, allow64BitInputThroughStrides, nodeNameToPartitionMap)) + if (IsNodeSupportedByDml(node, *registry, supportedDeviceDataTypeMask, internalRegInfoMap)) { *isDmlNode = true; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.h index 1b2744ecb4..e0fd8af31d 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.h @@ -64,9 +64,7 @@ namespace Dml const onnxruntime::Node& node, const onnxruntime::KernelRegistry& registry, uint32_t supportedDeviceDataTypeMask, // Each bit corresponds to each DML_TENSOR_DATA_TYPE. - const Windows::AI::MachineLearning::Adapter::InternalRegistrationInfoMap& internalRegInfoMap, - bool allow64BitInputThroughStrides, - _In_opt_ const std::unordered_map* nodeNameToPartitionMap // Only used when allow64BitInputThroughStrides is true + const Windows::AI::MachineLearning::Adapter::InternalRegistrationInfoMap& internalRegInfoMap ); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphTransformer.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphTransformer.cpp index 54223e450f..dc8399059f 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphTransformer.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphTransformer.cpp @@ -89,14 +89,12 @@ namespace Dml // We need to predict whether the nodes will be assigned to the DML transformer by Lotus, // which occurs in IExecutionProvider::GetCapability. - bool allow64BitInputThroughStrides = false; if (!IsNodeSupportedByDml( node, *registry, m_providerImpl->GetSupportedDeviceDataTypeMask(), - *m_providerImpl->GetInternalRegistrationInfoMap().get(), - allow64BitInputThroughStrides, - nullptr)) + *m_providerImpl->GetInternalRegistrationInfoMap().get() + )) { // Can't fuse nodes that don't belong to this execution provider continue; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCumSum.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCumSum.cpp index 2a80555583..8d42a9d38d 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCumSum.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCumSum.cpp @@ -51,6 +51,7 @@ public: } }; -DML_OP_DEFINE_CREATION_FUNCTION(CumSum, DmlOperatorCumSum); +DML_OP_DEFINE_CREATION_FUNCTION(CumSum11, DmlOperatorCumSum); +DML_OP_DEFINE_CREATION_FUNCTION(CumSum14, DmlOperatorCumSum); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index b072d1509b..700b6398c5 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -53,8 +53,8 @@ DEFINE_ENUM_FLAG_OPERATORS(Dml::SupportedTensorDataTypes); enum class DmlGraphSupport : uint32_t { - Supported = 0, - NotSupported = 1, + Supported = 0, + NotSupported = 1, }; DEFINE_ENUM_FLAG_OPERATORS(DmlGraphSupport); @@ -236,7 +236,8 @@ DML_OP_EXTERN_CREATION_FUNCTION(ConstantOfShape); DML_OP_EXTERN_CREATION_FUNCTION(IsInf); DML_OP_EXTERN_CREATION_FUNCTION(Mod); DML_OP_EXTERN_CREATION_FUNCTION(BitShift); -DML_OP_EXTERN_CREATION_FUNCTION(CumSum); +DML_OP_EXTERN_CREATION_FUNCTION(CumSum11); +DML_OP_EXTERN_CREATION_FUNCTION(CumSum14); DML_OP_EXTERN_CREATION_FUNCTION(GatherElements); DML_OP_EXTERN_CREATION_FUNCTION(GatherND); DML_OP_EXTERN_CREATION_FUNCTION(Range); @@ -278,6 +279,7 @@ constexpr static std::array supportedTypeListFloat1 constexpr static std::array supportedTypeListFloat16to32Ints32 = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::UInt32}; constexpr static std::array supportedTypeListFloat16to32Ints8to32 = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Ints8Bit | SupportedTensorDataTypes::Ints16Bit | SupportedTensorDataTypes::Ints32Bit}; constexpr static std::array supportedTypeListFloat16to32Ints8to64 = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Ints8Bit | SupportedTensorDataTypes::Ints16Bit | SupportedTensorDataTypes::Ints32Bit | SupportedTensorDataTypes::Ints64Bit}; +constexpr static std::array supportedTypeListFloat16to32SignedInts8to32 = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Int8 | SupportedTensorDataTypes::Int16 | SupportedTensorDataTypes::Int32}; constexpr static std::array supportedTypeListFloat16to32Ints32to64 = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Ints32Bit | SupportedTensorDataTypes::Ints64Bit}; constexpr static std::array supportedTypeListUInt8to64 = {SupportedTensorDataTypes::UInt8to64}; constexpr static std::array supportedTypeListNumericDefault = { SupportedTensorDataTypes::NumericDefault }; @@ -388,6 +390,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, InstanceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 9, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. + // TODO: Add additional type constraints in BatchNormalization-15, with scale and bias (T1) being different from input X (T). + // Add training-mode support to BatchNormalization-15 https://github.com/onnx/onnx/pull/3333 + // {REG_INFO( 14, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. + // {REG_INFO( 15, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. {REG_INFO( 7, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 13, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, MeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -395,8 +401,14 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 13, MeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, LpNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, RNN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, + // TODO: Allow recurrent operations to be batchwise https://github.com/onnx/onnx/pull/3217 + // {REG_INFO( 14, RNN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, {REG_INFO( 7, GRU, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, + // TODO: Allow recurrent operations to be batchwise https://github.com/onnx/onnx/pull/3217 + // {REG_INFO( 14, GRU, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, {REG_INFO( 7, LSTM, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, + // TODO: Allow recurrent operations to be batchwise https://github.com/onnx/onnx/pull/3217 + // {REG_INFO( 14, LSTM, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, {REG_INFO_MS( 1, ConvTransposeWithDynamicPads, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(2))}, // Data Reorganization Layers @@ -441,10 +453,13 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 11, ScatterND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmlGraphSupport::Supported)}, {REG_INFO( 13, ScatterND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmlGraphSupport::Supported)}, {REG_INFO( 9, EyeLike, typeNameListEyeLike, supportedTypeListScalars8to32, DmlGraphSupport::Supported)}, + // TODO: Add Trilu-14 to fill diagonal matrix https://github.com/onnx/onnx/pull/3291 + // {REG_INFO( 14, Trilu, typeNameListTrilu, supportedTypeListScalars8to32, DmlGraphSupport::Supported)}, // Data reorganization that merely changes the dimensions while keeping the data identical. {REG_INFO_ID( 7, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, {REG_INFO_ID( 13, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_ID( 14, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, {REG_INFO_ID( 7, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, {REG_INFO_ID( 9, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, {REG_INFO_ID( 11, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, @@ -457,6 +472,8 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO_ID( 13, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, {REG_INFO_ID( 7, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, {REG_INFO_ID( 13, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + // TODO: Add allowzero attribute. + // {REG_INFO_ID(14, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, // Elementwise {REG_INFO( 7, Sqrt, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -480,14 +497,18 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO_VER( 11, Clip, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1,2))}, {REG_INFO_VER( 12, Clip, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported, requiredConstantCpuInputs(1,2))}, {REG_INFO_VER( 13, Clip, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported, requiredConstantCpuInputs(1,2))}, - {REG_INFO( 7, Add, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported)}, - {REG_INFO( 13, Add, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported)}, - {REG_INFO( 7, Sub, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported)}, - {REG_INFO( 13, Sub, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported)}, - {REG_INFO( 7, Mul, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported)}, - {REG_INFO( 13, Mul, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported)}, - {REG_INFO( 7, Div, typeNameListDefault, supportedTypeListFloat16to32Ints32, DmlGraphSupport::Supported)}, - {REG_INFO( 13, Div, typeNameListDefault, supportedTypeListFloat16to32Ints32, DmlGraphSupport::Supported)}, + {REG_INFO( 7, Add, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 13, Add, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 14, Add, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 7, Sub, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 13, Sub, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 14, Sub, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 7, Mul, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 13, Mul, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 14, Mul, typeNameListDefault, supportedTypeListFloat16to32Ints8to64, DmlGraphSupport::Supported)}, + {REG_INFO( 7, Div, typeNameListDefault, supportedTypeListFloat16to32Ints8to32, DmlGraphSupport::Supported)}, + {REG_INFO( 13, Div, typeNameListDefault, supportedTypeListFloat16to32Ints8to32, DmlGraphSupport::Supported)}, + {REG_INFO( 14, Div, typeNameListDefault, supportedTypeListFloat16to32Ints8to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Sum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, {REG_INFO( 8, Sum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, {REG_INFO( 13, Sum, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), 2)}, @@ -598,6 +619,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO_VER( 10, Upsample, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1) /*scales*/)}, {REG_INFO_VER( 10, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1) /*scales*/)}, {REG_INFO_VER( 11, Resize, typeNameListTwo, supportedTypeListResize11, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3) /*roi, scales, sizes*/, std::nullopt, QueryResize)}, + // TODO: Resize-13 support nearest rounding mode attribute https://github.com/onnx/onnx/pull/3026 {REG_INFO_VER( 13, Resize, typeNameListTwo, supportedTypeListResize13, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3) /*roi, scales, sizes*/, std::nullopt, QueryResize)}, // Activation Functions @@ -609,9 +631,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, ScaledTanh, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Relu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 13, Relu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO( 14, Relu, typeNameListDefault, supportedTypeListFloat16to32SignedInts8to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, LeakyRelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, PRelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 9, PRelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO( 9, PRelu, typeNameListDefault, supportedTypeListFloat16to32SignedInts8to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, ThresholdedRelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 10, ThresholdedRelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Elu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -619,10 +642,16 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, Selu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + // TODO: Update Softmax-13/LogSoftmax-13/Hardmax-13 family ops behavior to align with other frameworks https://github.com/onnx/onnx/pull/2879 + // {REG_INFO( 13, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + // TODO: Update Softmax-13/LogSoftmax-13/Hardmax-13 family ops behavior to align with other frameworks https://github.com/onnx/onnx/pull/2879 + // {REG_INFO( 13, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + // TODO: Update Softmax-13/LogSoftmax-13/Hardmax-13 family ops behavior to align with other frameworks https://github.com/onnx/onnx/pull/2879 + // {REG_INFO( 13, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softsign, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softplus, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, ParametricSoftplus, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -637,6 +666,8 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, Cast, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, {REG_INFO( 9, Cast, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, {REG_INFO( 13, Cast, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, + // TODO: Add CastLike-15 https://github.com/onnx/onnx/pull/3558 + // {REG_INFO( 15, CastLike, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, {REG_INFO( 7, MemcpyFromHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO( 7, MemcpyToHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO_VER( 7, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported)}, @@ -644,6 +675,8 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO_VER( 11, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, {REG_INFO( 9, OneHot, typeNameListThree, supportedTypeListOneHot, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, {REG_INFO( 11, OneHot, typeNameListThree, supportedTypeListOneHot, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + // Shape-1, Shape-13, Shape-15 rely on CPU. + // Size-1 relies on CPU. // Fused operators {REG_INFO_MSDML(1, FusedConv, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -662,7 +695,8 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 11, BitShift, typeNameListDefault, supportedTypeListUInt8to64, DmlGraphSupport::Supported)}, {REG_INFO( 11, Round, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 10, ReverseSequence, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO( 11, CumSum, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO_VER( 11, CumSum, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO_VER( 14, CumSum, typeNameListDefault, supportedTypeListFloat16to32Ints32to64, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, {REG_INFO( 11, Range, typeNameListDefault, supportedTypeListRange, DmlGraphSupport::Supported, requiredConstantCpuInputs(0,1,2))}, {REG_INFO( 9, MaxUnpool, typeNameListTwo, supportedTypeListMaxUnpool, DmlGraphSupport::Supported, requiredConstantCpuInputs(2))}, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp index b9302ed01e..2412a41d16 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp @@ -163,6 +163,7 @@ namespace Dml // The filter for activation functions maps to what DML's fused op internally fuses at the shader level. OperatorInfo{ "Add", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_Add, {"Relu", "LeakyRelu"} }, OperatorInfo{ "Add", onnxruntime::kOnnxDomain, OnnxOperatorSet13::sc_sinceVer_Add, {"Relu", "LeakyRelu"} }, + OperatorInfo{ "Add", onnxruntime::kOnnxDomain, OnnxOperatorSet14::sc_sinceVer_Add, {"Relu", "LeakyRelu"} }, OperatorInfo{ "Sum", onnxruntime::kOnnxDomain, OnnxOperatorSet8::sc_sinceVer_Sum, {"Relu", "LeakyRelu"}, 2 }, OperatorInfo{ "Sum", onnxruntime::kOnnxDomain, OnnxOperatorSet13::sc_sinceVer_Sum, {"Relu", "LeakyRelu"}, 2 }, }; @@ -179,6 +180,7 @@ namespace Dml OperatorInfo{ "ScaledTanh", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_ScaledTanh }, OperatorInfo{ "Relu", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_Relu }, OperatorInfo{ "Relu", onnxruntime::kOnnxDomain, OnnxOperatorSet13::sc_sinceVer_Relu }, + OperatorInfo{ "Relu", onnxruntime::kOnnxDomain, OnnxOperatorSet14::sc_sinceVer_Relu }, OperatorInfo{ "LeakyRelu", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_LeakyRelu }, OperatorInfo{ "PRelu", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_PRelu }, OperatorInfo{ "PRelu", onnxruntime::kOnnxDomain, OnnxOperatorSet9::sc_sinceVer_PRelu }, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp index b8382d59b2..513bc125b9 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.cpp @@ -201,62 +201,6 @@ TensorDesc::TensorDesc( assert(m_bufferTensorDesc.TotalTensorSizeInBytes >= ComputeByteSizeFromDimensions(nonBroadcastDimensions, dataType)); } -void TensorDesc::Remap64bitDmlDataTypeTo32bit() -{ - if (m_bufferTensorDesc.DataType != DML_TENSOR_DATA_TYPE_UINT64 && - m_bufferTensorDesc.DataType != DML_TENSOR_DATA_TYPE_INT64) - { - return; // Nothing to do. - } - - uint64_t endPaddingInBytes = 0; - - // A workaround for older devices is to use strides to fake 64-bit memory access - // while only the lower 32 bits contains the data. This trick obviously doesn't - // work if the data element is genuine 64-bit. It also doesn't work if the data - // element is negative as the signed bit will be incorrectly interpreted. - m_bufferTensorDesc.DataType = Dml::Remap64bitDmlDataTypeTo32bit(m_bufferTensorDesc.DataType); - - // If the strides haven't been calculated yet, initialize them as packed. - if (m_bufferTensorDesc.Strides == nullptr) - { - uint32_t stride = 1; - for (int i = m_bufferTensorDesc.DimensionCount - 1; i >= 0; i--) - { - m_strides[i] = stride; - stride *= m_sizes[i]; - } - } - - // Double the stride values to emulate 64-bit integer support. - for (uint32_t i = 0; i < m_bufferTensorDesc.DimensionCount; ++i) - { - m_strides[i] *= 2; - } - - // The physical size of the tensor will have an extra 4 bytes at the end. - // DMLCalcBufferTensorSize calculates the minimum implied size, which is based on the last - // addressable element of the tensor plus the space for the last element. However, the size - // of the last element is now halved from 8 bytes to 4 bytes. - // - // Example: - // Original Tensor: size={2,3}, strides={3,1}, type=int64, size = (1+{1,2}*{3,1})*sizeof(int64) = 6 * 8 = 48 - // Emulated Tensor: size={2,3}, strides={6,2}, type=int32, size = (1+{1,2}*{6,2})*sizeof(int32) = 11 * 4 = 44 - // - // DirectML itself won't read/write the last 4 bytes, but we want the total size to be accurate - // so that the entire region can be zeroed. - endPaddingInBytes = sizeof(uint32_t); - - m_bufferTensorDesc.Strides = m_strides; - - m_bufferTensorDesc.TotalTensorSizeInBytes = DMLCalcBufferTensorSize( - m_bufferTensorDesc.DataType, - m_bufferTensorDesc.DimensionCount, - m_sizes, - m_strides - ) + endPaddingInBytes; -} - gsl::span TensorDesc::GetStrides() const { if (m_bufferTensorDesc.Strides == nullptr) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h index b48725eb11..867b0c29c0 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/TensorDesc.h @@ -37,7 +37,6 @@ namespace Dml inline DML_TENSOR_DATA_TYPE GetDmlDataType() const { return m_bufferTensorDesc.DataType; } inline MLOperatorTensorDataType GetMlOperatorDataType() const { return m_mlOperatorTensorDataType; } void ForceUnsignedDataType(); - void Remap64bitDmlDataTypeTo32bit(); inline bool IsValid() const { return m_tensorType != DML_TENSOR_TYPE_INVALID; } inline uint32_t GetDimensionCount() const { return m_bufferTensorDesc.DimensionCount; } diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index 7a3dd34760..609a0c33ec 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -1469,6 +1469,7 @@ using ShapeInferenceHelper_Gelu = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Identity7 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Identity13 = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_Identity14 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_MatMul = MatMulHelper; using ShapeInferenceHelper_MatMulInteger = MatMulHelper; using ShapeInferenceHelper_QLinearMatMul = QLinearMatMulHelper; @@ -1489,7 +1490,8 @@ using ShapeInferenceHelper_RandomNormalLike = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Multinomial = MultinomialHelper; using ShapeInferenceHelper_ReverseSequence = GetOutputShapeAsInputShapeHelper; -using ShapeInferenceHelper_CumSum = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_CumSum11 = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_CumSum14 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Range = RangeHelper; using ShapeInferenceHelper_FusedConv = ConvHelper; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h index 4ba4f944bd..e01fbb1ea4 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h @@ -338,6 +338,17 @@ namespace OperatorHelper static const int sc_sinceVer_ReduseSum = 13; } // namespace OnnxOperatorSet13 + namespace OnnxOperatorSet14 + { + static const int sc_sinceVer_Add = 14; + static const int sc_sinceVer_CumSum = 14; + static const int sc_sinceVer_Div = 14; + static const int sc_sinceVer_Identity = 14; + static const int sc_sinceVer_Mul = 14; + static const int sc_sinceVer_Relu = 14; + static const int sc_sinceVer_Sub = 14; + } // namespace OnnxOperatorSet14 + namespace MsftOperatorSet1 { static const int sc_sinceVer_FusedConv = 1; From d2519ec0c26c31e632e4dd72ad9550f1f9a18df7 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Thu, 19 May 2022 14:34:11 -0700 Subject: [PATCH 05/19] DirectML EP add CastLike-15 and Reshape-14 (#11568) * Add CastLike15 * Add Reshape14 * Fix allowzero comment * Rename REG_INFO_ID to REG_INFO_COPY to be clearer to readers --- .../src/Operators/DmlOperatorCast.cpp | 10 +++- .../src/Operators/OperatorRegistration.cpp | 39 +++++++-------- .../dml/OperatorAuthorHelper/Attributes.h | 1 + .../OperatorAuthorHelper/OperatorHelper.cpp | 50 ++++++++++++------- .../dml/OperatorAuthorHelper/OperatorHelper.h | 3 ++ .../OperatorAuthorHelper/OperatorVersions.h | 6 +++ 6 files changed, 68 insertions(+), 41 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCast.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCast.cpp index 311da999cc..76b9b308fe 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCast.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorCast.cpp @@ -15,7 +15,11 @@ public: const MLOperatorKernelCreationContext& kernelInfo ) : DmlOperator(kernelInfo) { - Initialize(kernelInfo); + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetInputCount() >= 1); + ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() == 1); + std::vector> inputIndices = { 0 }; // For CastLike, the second tensor ('target_type') is not bound. + std::vector> outputIndices = { 0 }; + DmlOperator::Initialize(kernelInfo, inputIndices, outputIndices); std::vector inputDescs = GetDmlInputDescs(); std::vector outputDescs = GetDmlOutputDescs(); @@ -38,10 +42,12 @@ public: m_compiledOperator.Get(), m_persistentResourceBinding ? &*m_persistentResourceBinding : nullptr, gsl::make_span(inputTensors), - gsl::make_span(outputTensors))); + gsl::make_span(outputTensors) + )); } }; DML_OP_DEFINE_CREATION_FUNCTION(Cast, DmlOperatorCast); +DML_OP_DEFINE_CREATION_FUNCTION(CastLike15, DmlOperatorCast); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index 700b6398c5..d39f177d37 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -188,6 +188,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(Affine); DML_OP_EXTERN_CREATION_FUNCTION(Dropout); DML_OP_EXTERN_CREATION_FUNCTION(MatMul); DML_OP_EXTERN_CREATION_FUNCTION(Cast); +DML_OP_EXTERN_CREATION_FUNCTION(CastLike15); DML_OP_EXTERN_CREATION_FUNCTION(MemcpyFromHost); DML_OP_EXTERN_CREATION_FUNCTION(MemcpyToHost); DML_OP_EXTERN_CREATION_FUNCTION(TopK7); @@ -349,7 +350,7 @@ constexpr auto requiredConstantCpuInputs(Args... args) // Identity operators use Copy, alias their first input, and use elementwise identity operators // when needed for striding support, but issue actual copies outside the graph. -#define REG_INFO_ID(version, operatorName, ...) \ +#define REG_INFO_COPY(version, operatorName, ...) \ #operatorName, OnnxOperatorSet##version::sc_sinceVer_##operatorName, onnxruntime::kOnnxDomain, CreateCopy, ShapeInferenceFunction, true, ##__VA_ARGS__, // MS-domain operators @@ -457,23 +458,22 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation // {REG_INFO( 14, Trilu, typeNameListTrilu, supportedTypeListScalars8to32, DmlGraphSupport::Supported)}, // Data reorganization that merely changes the dimensions while keeping the data identical. - {REG_INFO_ID( 7, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 13, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 14, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 7, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 9, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 11, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 13, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 7, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 11, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 13, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, - {REG_INFO_ID( 7, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 11, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_ID( 13, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, - {REG_INFO_ID( 7, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, - {REG_INFO_ID( 13, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, - // TODO: Add allowzero attribute. - // {REG_INFO_ID(14, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO_COPY( 7, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(13, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(14, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY( 7, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY( 9, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(11, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(13, Flatten, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY( 7, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(11, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(13, Squeeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO_COPY( 7, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(11, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, + {REG_INFO_COPY(13, Unsqueeze, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO_COPY( 7, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO_COPY(13, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, + {REG_INFO_COPY(14, Reshape, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, // Elementwise {REG_INFO( 7, Sqrt, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -666,8 +666,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, Cast, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, {REG_INFO( 9, Cast, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, {REG_INFO( 13, Cast, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, - // TODO: Add CastLike-15 https://github.com/onnx/onnx/pull/3558 - // {REG_INFO( 15, CastLike, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, + {REG_INFO_VER( 15, CastLike, typeNameListTwo, supportedTypeListCast, DmlGraphSupport::Supported)}, {REG_INFO( 7, MemcpyFromHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO( 7, MemcpyToHost, typeNameListDefault, supportedTypeListAll)}, {REG_INFO_VER( 7, TopK, typeNameListTopK, supportedTypeListTopK, DmlGraphSupport::Supported)}, diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h index 71392bf155..99e7d1d390 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h @@ -9,6 +9,7 @@ namespace AttrName static constexpr const char* ActivationAlpha = "activation_alpha"; static constexpr const char* ActivationBeta = "activation_beta"; static constexpr const char* Activations = "activations"; + static constexpr const char* AllowZero = "allowzero"; static constexpr const char* Alpha = "alpha"; static constexpr const char* AutoPad = "auto_pad"; static constexpr const char* Axes = "axes"; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp index efe5d3f4a1..268e12706d 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp @@ -1971,32 +1971,44 @@ namespace OperatorHelper int inferDim = -1; DimensionType inElementCount = ComputeElementCountFromDimensions(inputDimensions); + bool allowZero = shapeInfo.template GetOptionalAttribute(AttrName::AllowZero, 0); - for (int i = 0, ci = gsl::narrow_cast(m_shapeDims.size()); i < ci; ++i) + if (allowZero) { - switch (m_shapeDims[i]) + // Just take the shape directly (no special handling for 0). + for (int i = 0, ci = gsl::narrow_cast(m_shapeDims.size()); i < ci; ++i) { - case -1: - ML_CHECK_VALID_ARGUMENT(inferDim == -1, "Only one dimension can be inferred."); - inferDim = i; - break; - - case 0: - outputDimensions[i] = inputDimensions[i]; - outElementCount *= outputDimensions[i]; - break; - - default: outputDimensions[i] = m_shapeDims[i]; - outElementCount *= outputDimensions[i]; - break; } } - - if (inferDim != -1) + else { - outputDimensions[inferDim] = inElementCount / outElementCount; - outElementCount *= outputDimensions[inferDim]; + // Special handling where 0 size means to copy the corresponding input tensor dimension. + for (int i = 0, ci = gsl::narrow_cast(m_shapeDims.size()); i < ci; ++i) + { + switch (m_shapeDims[i]) + { + case -1: + ML_CHECK_VALID_ARGUMENT(inferDim == -1, "Only one dimension can be inferred."); + inferDim = i; + break; + + case 0: + outputDimensions[i] = inputDimensions[i]; + outElementCount *= outputDimensions[i]; + break; + + default: + outputDimensions[i] = m_shapeDims[i]; + outElementCount *= outputDimensions[i]; + break; + } + } + + if (inferDim != -1) + { + outputDimensions[inferDim] = inElementCount / outElementCount; + } } return { EdgeShapes(outputDimensions) }; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index 609a0c33ec..188eb3dff0 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -1365,6 +1365,7 @@ using ShapeInferenceHelper_EyeLike = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Expand = ExpandHelper; using ShapeInferenceHelper_Reshape7 = ReshapeHelper; using ShapeInferenceHelper_Reshape13 = ReshapeHelper; +using ShapeInferenceHelper_Reshape14 = ReshapeHelper; using ShapeInferenceHelper_ConstantOfShape = ConstantOfShapeHelper; using ShapeInferenceHelper_Tile = TileHelper; using ShapeInferenceHelper_Resize10 = VersionedOpsetHelper; @@ -1494,6 +1495,8 @@ using ShapeInferenceHelper_CumSum11 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_CumSum14 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Range = RangeHelper; +using ShapeInferenceHelper_CastLike15 = GetOutputShapeAsInputShapeHelper; + using ShapeInferenceHelper_FusedConv = ConvHelper; using ShapeInferenceHelper_FusedConvTranspose = ConvTransposeHelper; using ShapeInferenceHelper_FusedInstanceNormalization = GetOutputShapeAsInputShapeHelper; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h index e01fbb1ea4..ec6e4dca67 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h @@ -346,9 +346,15 @@ namespace OperatorHelper static const int sc_sinceVer_Identity = 14; static const int sc_sinceVer_Mul = 14; static const int sc_sinceVer_Relu = 14; + static const int sc_sinceVer_Reshape = 14; static const int sc_sinceVer_Sub = 14; } // namespace OnnxOperatorSet14 + namespace OnnxOperatorSet15 + { + static const int sc_sinceVer_CastLike = 15; + } // namespace OnnxOperatorSet14 + namespace MsftOperatorSet1 { static const int sc_sinceVer_FusedConv = 1; From aa3a825816555d1f859385a256804dfceefebd2f Mon Sep 17 00:00:00 2001 From: sumitsays Date: Tue, 7 Jun 2022 14:31:55 -0700 Subject: [PATCH 06/19] Added Softmax/Hardmax/LogSoftmax-13 (#11772) * Added Softmax/Hardmax/LogSoftmax-13 * Removed redundant method specifier Co-authored-by: Sumit Agarwal --- .../src/External/DirectMLHelpers/ApiHelpers.h | 6 +++ .../src/External/DirectMLHelpers/ApiTraits.h | 46 ++++++++++++++++- .../External/DirectMLHelpers/DirectMLSchema.h | 45 +++++++++++++++++ .../DirectMLHelpers/GeneratedSchemaHelpers.h | 42 ++++++++++++++++ .../src/Operators/DmlOperatorActivation.cpp | 49 +++++++++++++++---- .../src/Operators/OperatorRegistration.cpp | 12 ++--- .../dml/OperatorAuthorHelper/OperatorHelper.h | 3 ++ .../OperatorAuthorHelper/OperatorVersions.h | 3 ++ 8 files changed, 189 insertions(+), 17 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h index 28ab7e167f..8c85e4ec1d 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h @@ -9,10 +9,12 @@ union ActivationOperatorDescUnion DML_ACTIVATION_ELU_OPERATOR_DESC elu; DML_ACTIVATION_CELU_OPERATOR_DESC celu; DML_ACTIVATION_HARDMAX_OPERATOR_DESC hardmax; + DML_ACTIVATION_HARDMAX1_OPERATOR_DESC hardmax1; DML_ACTIVATION_HARD_SIGMOID_OPERATOR_DESC hardSigmoid; DML_ACTIVATION_LEAKY_RELU_OPERATOR_DESC leakyRelu; DML_ACTIVATION_LINEAR_OPERATOR_DESC linear; DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_DESC logSoftmax; + DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_DESC logSoftmax1; DML_ACTIVATION_PARAMETERIZED_RELU_OPERATOR_DESC parameterizedRelu; DML_ACTIVATION_PARAMETRIC_SOFTPLUS_OPERATOR_DESC parametricSoftplus; DML_ACTIVATION_RELU_OPERATOR_DESC relu; @@ -20,6 +22,7 @@ union ActivationOperatorDescUnion DML_ACTIVATION_SCALED_ELU_OPERATOR_DESC scaledElu; DML_ACTIVATION_SIGMOID_OPERATOR_DESC sigmoid; DML_ACTIVATION_SOFTMAX_OPERATOR_DESC softmax; + DML_ACTIVATION_SOFTMAX1_OPERATOR_DESC softmax1; DML_ACTIVATION_SOFTPLUS_OPERATOR_DESC softplus; DML_ACTIVATION_SOFTSIGN_OPERATOR_DESC softsign; DML_ACTIVATION_TANH_OPERATOR_DESC tanh; @@ -41,11 +44,13 @@ struct ActivationOperatorDesc case DML_OPERATOR_ACTIVATION_ELU: return { activationType, ¶ms.elu }; case DML_OPERATOR_ACTIVATION_CELU: return { activationType, ¶ms.celu }; case DML_OPERATOR_ACTIVATION_HARDMAX: return { activationType, ¶ms.hardmax }; + case DML_OPERATOR_ACTIVATION_HARDMAX1: return { activationType, ¶ms.hardmax1 }; case DML_OPERATOR_ACTIVATION_HARD_SIGMOID: return { activationType, ¶ms.sigmoid }; case DML_OPERATOR_ACTIVATION_IDENTITY: return { activationType, ¶ms.identity }; case DML_OPERATOR_ACTIVATION_LEAKY_RELU: return { activationType, ¶ms.leakyRelu }; case DML_OPERATOR_ACTIVATION_LINEAR: return { activationType, ¶ms.linear }; case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX: return { activationType, ¶ms.logSoftmax }; + case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1: return { activationType, ¶ms.logSoftmax1 }; case DML_OPERATOR_ACTIVATION_PARAMETERIZED_RELU: return { activationType, ¶ms.parameterizedRelu }; case DML_OPERATOR_ACTIVATION_PARAMETRIC_SOFTPLUS: return { activationType, ¶ms.parametricSoftplus }; case DML_OPERATOR_ACTIVATION_RELU: return { activationType, ¶ms.relu }; @@ -53,6 +58,7 @@ struct ActivationOperatorDesc case DML_OPERATOR_ACTIVATION_SCALED_TANH: return { activationType, ¶ms.scaledTanh }; case DML_OPERATOR_ACTIVATION_SIGMOID: return { activationType, ¶ms.sigmoid }; case DML_OPERATOR_ACTIVATION_SOFTMAX: return { activationType, ¶ms.softmax }; + case DML_OPERATOR_ACTIVATION_SOFTMAX1: return { activationType, ¶ms.softmax1 }; case DML_OPERATOR_ACTIVATION_SOFTPLUS: return { activationType, ¶ms.softplus }; case DML_OPERATOR_ACTIVATION_SOFTSIGN: return { activationType, ¶ms.softsign }; case DML_OPERATOR_ACTIVATION_TANH: return { activationType, ¶ms.tanh }; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h index 1774a8a2b0..7d4c75e6ca 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h @@ -24,8 +24,8 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 154; - static constexpr size_t ActivationFunctionCount = 21; + static constexpr auto ValueCount = 157; + static constexpr size_t ActivationFunctionCount = 24; }; template <> @@ -1011,6 +1011,12 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_HARDMAX; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_HARDMAX1; +}; + template <> struct OperatorDescTraits { @@ -1041,6 +1047,12 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_LOG_SOFTMAX; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1; +}; + template <> struct OperatorDescTraits { @@ -1083,6 +1095,12 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_SOFTMAX; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_SOFTMAX1; +}; + template <> struct OperatorDescTraits { @@ -1959,6 +1977,12 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_HARDMAX> using DescType = DML_ACTIVATION_HARDMAX_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_HARDMAX1> +{ + using DescType = DML_ACTIVATION_HARDMAX1_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_HARD_SIGMOID> { @@ -1989,6 +2013,12 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_LOG_SOFTMAX using DescType = DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1> +{ + using DescType = DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_PARAMETERIZED_RELU> { @@ -2031,6 +2061,12 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_SOFTMAX> using DescType = DML_ACTIVATION_SOFTMAX_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_SOFTMAX1> +{ + using DescType = DML_ACTIVATION_SOFTMAX1_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_SOFTPLUS> { @@ -2360,6 +2396,8 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_ACTIVATION_CELU_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_HARDMAX: return std::invoke(std::forward(visitor), DML_ACTIVATION_HARDMAX_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ACTIVATION_HARDMAX1: + return std::invoke(std::forward(visitor), DML_ACTIVATION_HARDMAX1_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_HARD_SIGMOID: return std::invoke(std::forward(visitor), DML_ACTIVATION_HARD_SIGMOID_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_IDENTITY: @@ -2370,6 +2408,8 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_ACTIVATION_LINEAR_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX: return std::invoke(std::forward(visitor), DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1: + return std::invoke(std::forward(visitor), DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_PARAMETERIZED_RELU: return std::invoke(std::forward(visitor), DML_ACTIVATION_PARAMETERIZED_RELU_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_PARAMETRIC_SOFTPLUS: @@ -2384,6 +2424,8 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_ACTIVATION_SIGMOID_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_SOFTMAX: return std::invoke(std::forward(visitor), DML_ACTIVATION_SOFTMAX_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_ACTIVATION_SOFTMAX1: + return std::invoke(std::forward(visitor), DML_ACTIVATION_SOFTMAX1_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_SOFTPLUS: return std::invoke(std::forward(visitor), DML_ACTIVATION_SOFTPLUS_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_SOFTSIGN: diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h index 137ea2b030..a993300291 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h @@ -2288,6 +2288,21 @@ constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA { DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA_FIELDS[4] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, "Axes", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA { + "DML_OPERATOR_ACTIVATION_HARDMAX1", + DML_OPERATOR_ACTIVATION_HARDMAX1, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 4, + DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA_FIELDS, +}; + constexpr DML_SCHEMA_FIELD DML_ACTIVATION_HARD_SIGMOID_OPERATOR_SCHEMA_FIELDS[4] { DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, @@ -2358,6 +2373,21 @@ constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_SCHEMA { DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA_FIELDS[4] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, "Axes", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA { + "DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1", + DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 4, + DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA_FIELDS, +}; + constexpr DML_SCHEMA_FIELD DML_ACTIVATION_PARAMETERIZED_RELU_OPERATOR_SCHEMA_FIELDS[3] { DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "SlopeTensor", false }, @@ -2456,6 +2486,21 @@ constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_SOFTMAX_OPERATOR_SCHEMA { DML_ACTIVATION_SOFTMAX_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA_FIELDS[4] { + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "AxisCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT_ARRAY, "Axes", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA { + "DML_OPERATOR_ACTIVATION_SOFTMAX1", + DML_OPERATOR_ACTIVATION_SOFTMAX1, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 4, + DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA_FIELDS, +}; + constexpr DML_SCHEMA_FIELD DML_ACTIVATION_SOFTPLUS_OPERATOR_SCHEMA_FIELDS[3] { DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h index e0a705330f..b37389ce17 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h @@ -1404,6 +1404,15 @@ inline std::vector GetFields(const DML_ACTIVATION_HARDMAX_OPERATO OperatorField(&DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), }; } +inline std::vector GetFields(const DML_ACTIVATION_HARDMAX1_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.AxisCount))), + OperatorField(&DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Axes), desc.AxisCount)), + }; +} inline std::vector GetFields(const DML_ACTIVATION_HARD_SIGMOID_OPERATOR_DESC& desc) { return { @@ -1444,6 +1453,15 @@ inline std::vector GetFields(const DML_ACTIVATION_LOG_SOFTMAX_OPE OperatorField(&DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), }; } +inline std::vector GetFields(const DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.AxisCount))), + OperatorField(&DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Axes), desc.AxisCount)), + }; +} inline std::vector GetFields(const DML_ACTIVATION_PARAMETERIZED_RELU_OPERATOR_DESC& desc) { return { @@ -1500,6 +1518,15 @@ inline std::vector GetFields(const DML_ACTIVATION_SOFTMAX_OPERATO OperatorField(&DML_ACTIVATION_SOFTMAX_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), }; } +inline std::vector GetFields(const DML_ACTIVATION_SOFTMAX1_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.AxisCount))), + OperatorField(&DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Axes), desc.AxisCount)), + }; +} inline std::vector GetFields(const DML_ACTIVATION_SOFTPLUS_OPERATOR_DESC& desc) { return { @@ -1689,11 +1716,13 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_ACTIVATION_ELU: return DML_ACTIVATION_ELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_CELU: return DML_ACTIVATION_CELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_HARDMAX: return DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA; + case DML_OPERATOR_ACTIVATION_HARDMAX1: return DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_HARD_SIGMOID: return DML_ACTIVATION_HARD_SIGMOID_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_IDENTITY: return DML_ACTIVATION_IDENTITY_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_LEAKY_RELU: return DML_ACTIVATION_LEAKY_RELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_LINEAR: return DML_ACTIVATION_LINEAR_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX: return DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_SCHEMA; + case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1: return DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_PARAMETERIZED_RELU: return DML_ACTIVATION_PARAMETERIZED_RELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_PARAMETRIC_SOFTPLUS: return DML_ACTIVATION_PARAMETRIC_SOFTPLUS_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_RELU: return DML_ACTIVATION_RELU_OPERATOR_SCHEMA; @@ -1701,6 +1730,7 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_ACTIVATION_SCALED_TANH: return DML_ACTIVATION_SCALED_TANH_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_SIGMOID: return DML_ACTIVATION_SIGMOID_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_SOFTMAX: return DML_ACTIVATION_SOFTMAX_OPERATOR_SCHEMA; + case DML_OPERATOR_ACTIVATION_SOFTMAX1: return DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_SOFTPLUS: return DML_ACTIVATION_SOFTPLUS_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_SOFTSIGN: return DML_ACTIVATION_SOFTSIGN_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_TANH: return DML_ACTIVATION_TANH_OPERATOR_SCHEMA; @@ -2276,6 +2306,10 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ACTIVATION_HARDMAX1: + return AbstractOperatorDesc( + &DML_ACTIVATION_HARDMAX1_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); case DML_OPERATOR_ACTIVATION_HARD_SIGMOID: return AbstractOperatorDesc( &DML_ACTIVATION_HARD_SIGMOID_OPERATOR_SCHEMA, @@ -2296,6 +2330,10 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_ACTIVATION_LOG_SOFTMAX_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1: + return AbstractOperatorDesc( + &DML_ACTIVATION_LOG_SOFTMAX1_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); case DML_OPERATOR_ACTIVATION_PARAMETERIZED_RELU: return AbstractOperatorDesc( &DML_ACTIVATION_PARAMETERIZED_RELU_OPERATOR_SCHEMA, @@ -2324,6 +2362,10 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_ACTIVATION_SOFTMAX_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_ACTIVATION_SOFTMAX1: + return AbstractOperatorDesc( + &DML_ACTIVATION_SOFTMAX1_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); case DML_OPERATOR_ACTIVATION_SOFTPLUS: return AbstractOperatorDesc( &DML_ACTIVATION_SOFTPLUS_OPERATOR_SCHEMA, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp index c7c9e50c90..61baa10cfb 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp @@ -25,7 +25,7 @@ public: ActivationOperatorDescUnion operatorDesc = {}; - int coerceAxis = TensorAxis::DoNotCoerce; + std::vector dmlAxes; switch (operatorType) { @@ -39,7 +39,29 @@ public: case DML_OPERATOR_ACTIVATION_HARDMAX: { const uint32_t onnxDimCount = gsl::narrow_cast(kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0).size()); - coerceAxis = HandleNegativeAxis(kernelCreationContext.GetOptionalAttribute(AttrName::Axis, 1), onnxDimCount); + int axis = HandleNegativeAxis(kernelCreationContext.GetOptionalAttribute(AttrName::Axis, 1), onnxDimCount); + std::vector onnxAxes(onnxDimCount - axis); + std::iota(onnxAxes.begin(), onnxAxes.end(), static_cast(axis)); + + dmlAxes.resize(onnxDimCount - axis); + GetDmlAdjustedAxes(onnxAxes, onnxDimCount, m_inputTensorDescs.front().GetDimensionCount(), /*out*/ dmlAxes); + + operatorDesc.hardmax1.Axes = dmlAxes.data(); + operatorDesc.hardmax1.AxisCount = gsl::narrow_cast(dmlAxes.size()); + } + break; + + case DML_OPERATOR_ACTIVATION_SOFTMAX1: + case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1: + case DML_OPERATOR_ACTIVATION_HARDMAX1: + { + const uint32_t onnxDimCount = gsl::narrow_cast(kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0).size()); + int onnxAxis = HandleNegativeAxis(kernelCreationContext.GetOptionalAttribute(AttrName::Axis, -1), onnxDimCount); + + dmlAxes.push_back(GetDmlAdjustedAxis(onnxAxis, onnxDimCount, m_inputTensorDescs.front().GetDimensionCount())); + + operatorDesc.hardmax1.Axes = dmlAxes.data(); + operatorDesc.hardmax1.AxisCount = gsl::narrow_cast(dmlAxes.size()); } break; @@ -100,12 +122,6 @@ public: break; } - if (coerceAxis != TensorAxis::DoNotCoerce) - { - m_inputTensorDescs[0] = CreateTensorDescFromInput(kernelCreationContext, 0, coerceAxis); - m_outputTensorDescs[0] = CreateTensorDescFromOutput(kernelCreationContext, 0, coerceAxis); - } - gsl::span outputSizes = m_outputTensorDescs[0].GetSizes(); std::vector inputDescs; std::vector outputDescs; @@ -135,9 +151,24 @@ public: operatorDesc.elu.OutputTensor = outputDescs.data(); } - DML_OPERATOR_DESC opDesc = { operatorType, &operatorDesc }; + DML_OPERATOR_DESC opDesc = { remappedOperatorType(operatorType), &operatorDesc }; SetDmlOperatorDesc(opDesc, kernelCreationContext); } + +private: + DML_OPERATOR_TYPE remappedOperatorType(const DML_OPERATOR_TYPE operatorType) const { + switch (operatorType) + { + case DML_OPERATOR_ACTIVATION_HARDMAX: + return DML_OPERATOR_ACTIVATION_HARDMAX1; + case DML_OPERATOR_ACTIVATION_SOFTMAX: + return DML_OPERATOR_ACTIVATION_SOFTMAX1; + case DML_OPERATOR_ACTIVATION_LOG_SOFTMAX: + return DML_OPERATOR_ACTIVATION_LOG_SOFTMAX1; + default: + return operatorType; + } + } }; // A specific type of operation for registration. diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index d39f177d37..bf1e4c762c 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -179,8 +179,11 @@ DML_OP_EXTERN_CREATION_FUNCTION(Elu); DML_OP_EXTERN_CREATION_FUNCTION(Celu); DML_OP_EXTERN_CREATION_FUNCTION(Selu); DML_OP_EXTERN_CREATION_FUNCTION(Softmax); +DML_OP_EXTERN_CREATION_FUNCTION(Softmax13); DML_OP_EXTERN_CREATION_FUNCTION(LogSoftmax); +DML_OP_EXTERN_CREATION_FUNCTION(LogSoftmax13); DML_OP_EXTERN_CREATION_FUNCTION(Hardmax); +DML_OP_EXTERN_CREATION_FUNCTION(Hardmax13); DML_OP_EXTERN_CREATION_FUNCTION(Softsign); DML_OP_EXTERN_CREATION_FUNCTION(Softplus); DML_OP_EXTERN_CREATION_FUNCTION(ParametricSoftplus); @@ -642,16 +645,13 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, Selu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - // TODO: Update Softmax-13/LogSoftmax-13/Hardmax-13 family ops behavior to align with other frameworks https://github.com/onnx/onnx/pull/2879 - // {REG_INFO( 13, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO( 13, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - // TODO: Update Softmax-13/LogSoftmax-13/Hardmax-13 family ops behavior to align with other frameworks https://github.com/onnx/onnx/pull/2879 - // {REG_INFO( 13, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO( 13, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - // TODO: Update Softmax-13/LogSoftmax-13/Hardmax-13 family ops behavior to align with other frameworks https://github.com/onnx/onnx/pull/2879 - // {REG_INFO( 13, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO( 13, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softsign, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softplus, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, ParametricSoftplus, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index 188eb3dff0..daceecab5c 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -1459,8 +1459,11 @@ using ShapeInferenceHelper_Elu = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Celu = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Selu = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Softmax = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_Softmax13 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_LogSoftmax = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_LogSoftmax13 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Hardmax = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_Hardmax13 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Softsign = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Softplus = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_ParametricSoftplus = GetOutputShapeAsInputShapeHelper; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h index ec6e4dca67..598d282bf2 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h @@ -336,6 +336,9 @@ namespace OperatorHelper static const int sc_sinceVer_Transpose = 13; static const int sc_sinceVer_Unsqueeze = 13; static const int sc_sinceVer_ReduseSum = 13; + static const int sc_sinceVer_Softmax = 13; + static const int sc_sinceVer_LogSoftmax = 13; + static const int sc_sinceVer_Hardmax = 13; } // namespace OnnxOperatorSet13 namespace OnnxOperatorSet14 From f5fe4f253c336a77e575a090a11bd180ff087d63 Mon Sep 17 00:00:00 2001 From: sumitsays Date: Wed, 8 Jun 2022 11:38:14 -0700 Subject: [PATCH 07/19] Registered Softmax/Hardmax/LogSoftmax-13 as Versioned Operator (#11787) * Added Softmax/Hardmax/LogSoftmax-13 * Removed redundant method specifier * Registered softmax/hardmax/logsoftmax as verisioned operator Co-authored-by: Sumit Agarwal --- .../src/Operators/DmlOperatorActivation.cpp | 3 +++ .../src/Operators/OperatorRegistration.cpp | 6 +++--- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp index 61baa10cfb..10c41ad10d 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorActivation.cpp @@ -198,8 +198,11 @@ DML_OP_DEFINE_CREATION_FUNCTION(Softplus, DmlOperatorActivationTempla DML_OP_DEFINE_CREATION_FUNCTION(ParametricSoftplus, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(Dropout, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(Softmax, DmlOperatorActivationTemplate); +DML_OP_DEFINE_CREATION_FUNCTION(Softmax13, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(LogSoftmax, DmlOperatorActivationTemplate); +DML_OP_DEFINE_CREATION_FUNCTION(LogSoftmax13, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(Hardmax, DmlOperatorActivationTemplate); +DML_OP_DEFINE_CREATION_FUNCTION(Hardmax13, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(Shrink, DmlOperatorActivationTemplate); DML_OP_DEFINE_CREATION_FUNCTION(Gelu, DmlOperatorActivationTemplate); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index bf1e4c762c..7c40dae020 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -645,13 +645,13 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, Selu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 13, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO_VER( 13, Softmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 13, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO_VER( 13, LogSoftmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 11, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 13, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO_VER( 13, Hardmax, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softsign, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Softplus, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, ParametricSoftplus, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, From 0f0b640b4b1a142a17574fa60159d33c34c2e614 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 8 Jun 2022 18:05:11 -0700 Subject: [PATCH 08/19] Reformat build.py for WindowsAI branch (#11794) --- tools/ci_build/build.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/tools/ci_build/build.py b/tools/ci_build/build.py index fb1cecb8d8..b4439353ba 100644 --- a/tools/ci_build/build.py +++ b/tools/ci_build/build.py @@ -519,7 +519,7 @@ def parse_arguments(): "--winml_root_namespace_override", type=str, help="Specify the namespace that WinML builds into." ) parser.add_argument( - "--dml_external_project", action='store_true', help="Build with DirectML as an external project." + "--dml_external_project", action="store_true", help="Build with DirectML as an external project." ) parser.add_argument( "--use_telemetry", action="store_true", help="Only official builds can set this flag to enable telemetry." @@ -1358,9 +1358,9 @@ def setup_dml_build(args, cmake_path, build_dir, configs): for expected_file in ["bin/DirectML.dll", "lib/DirectML.lib", "include/DirectML.h"]: file_path = os.path.join(args.dml_path, expected_file) if not os.path.exists(file_path): - raise BuildError("dml_path is invalid.", - "dml_path='{}' expected_file='{}'." - .format(args.dml_path, file_path)) + raise BuildError( + "dml_path is invalid.", "dml_path='{}' expected_file='{}'.".format(args.dml_path, file_path) + ) elif not args.dml_external_project: for config in configs: # Run the RESTORE_PACKAGES target to perform the initial From 5e54611427d6c98cea014c200b24edee57c92295 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 8 Jun 2022 19:08:00 -0700 Subject: [PATCH 09/19] DML EP add Trilu-14 and Resize-13 nearest mode and others (#11782) * Add Trilu-14 kernel * Support Resize with rounding direction for round_prefer_ceil/round_prefer_floor * Add batch normalization query and RNN query * Appease CPPLINT.cfg per https://raw.githubusercontent.com/google/styleguide/gh-pages/cpplint/cpplint.py to reduce the noise --- onnxruntime/core/providers/dml/CPPLINT.cfg | 1 + .../src/External/DirectMLHelpers/ApiTraits.h | 47 +++++++++++++++- .../External/DirectMLHelpers/DirectMLSchema.h | 55 ++++++++++++++++++ .../DirectMLHelpers/GeneratedSchemaHelpers.h | 54 +++++++++++++++++- .../src/MLOperatorAuthorImpl.cpp | 3 +- .../DmlOperatorBatchNormalization.cpp | 9 +++ .../src/Operators/DmlOperatorEyeLike.cpp | 2 +- .../DmlOperatorRecurrentNeuralNetwork.cpp | 26 ++++++++- .../src/Operators/DmlOperatorResize.cpp | 54 +++++++++--------- .../src/Operators/DmlOperatorTrilu.cpp | 56 +++++++++++++++++++ .../src/Operators/OperatorRegistration.cpp | 45 ++++++++------- .../src/Operators/OperatorUtility.cpp | 5 +- .../dml/OperatorAuthorHelper/Attributes.h | 3 + .../dml/OperatorAuthorHelper/OperatorHelper.h | 1 + .../OperatorAuthorHelper/OperatorVersions.h | 7 +++ 15 files changed, 311 insertions(+), 57 deletions(-) create mode 100644 onnxruntime/core/providers/dml/CPPLINT.cfg create mode 100644 onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorTrilu.cpp diff --git a/onnxruntime/core/providers/dml/CPPLINT.cfg b/onnxruntime/core/providers/dml/CPPLINT.cfg new file mode 100644 index 0000000000..e7dbd3164b --- /dev/null +++ b/onnxruntime/core/providers/dml/CPPLINT.cfg @@ -0,0 +1 @@ +filter=-whitespace/braces,-whitespace/parens,-whitespace/line_length,-whitespace/indent diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h index 7d4c75e6ca..8b5cf36937 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiTraits.h @@ -24,7 +24,7 @@ struct EnumTraits template <> struct EnumTraits { - static constexpr auto ValueCount = 157; + static constexpr auto ValueCount = 160; static constexpr size_t ActivationFunctionCount = 24; }; @@ -993,6 +993,24 @@ struct OperatorDescTraits static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_BATCH_NORMALIZATION_TRAINING; }; +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_RESAMPLE2; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_RESAMPLE_GRAD1; +}; + +template <> +struct OperatorDescTraits +{ + static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_DIAGONAL_MATRIX1; +}; + template <> struct OperatorDescTraits { @@ -1959,6 +1977,24 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_BATCH_NORMALIZATION_TR using DescType = DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_DESC; }; +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_RESAMPLE2> +{ + using DescType = DML_RESAMPLE2_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_RESAMPLE_GRAD1> +{ + using DescType = DML_RESAMPLE_GRAD1_OPERATOR_DESC; +}; + +template <> +struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_DIAGONAL_MATRIX1> +{ + using DescType = DML_DIAGONAL_MATRIX1_OPERATOR_DESC; +}; + template <> struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_ELU> { @@ -2390,6 +2426,12 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args return std::invoke(std::forward(visitor), DML_ROI_ALIGN_GRAD_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_BATCH_NORMALIZATION_TRAINING: return std::invoke(std::forward(visitor), DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_RESAMPLE2: + return std::invoke(std::forward(visitor), DML_RESAMPLE2_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_RESAMPLE_GRAD1: + return std::invoke(std::forward(visitor), DML_RESAMPLE_GRAD1_OPERATOR_DESC{}, std::forward(args)...); + case DML_OPERATOR_DIAGONAL_MATRIX1: + return std::invoke(std::forward(visitor), DML_DIAGONAL_MATRIX1_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_ELU: return std::invoke(std::forward(visitor), DML_ACTIVATION_ELU_OPERATOR_DESC{}, std::forward(args)...); case DML_OPERATOR_ACTIVATION_CELU: @@ -2588,6 +2630,9 @@ inline gsl::czstring ToString(DML_OPERATOR_TYPE value) case DML_OPERATOR_ELEMENT_WISE_QUANTIZED_LINEAR_ADD: return "DML_OPERATOR_ELEMENT_WISE_QUANTIZED_LINEAR_ADD"; case DML_OPERATOR_ROI_ALIGN_GRAD: return "DML_OPERATOR_ROI_ALIGN_GRAD"; case DML_OPERATOR_BATCH_NORMALIZATION_TRAINING: return "DML_OPERATOR_BATCH_NORMALIZATION_TRAINING"; + case DML_OPERATOR_RESAMPLE2: return "DML_OPERATOR_RESAMPLE2"; + case DML_OPERATOR_RESAMPLE_GRAD1: return "DML_OPERATOR_RESAMPLE_GRAD1"; + case DML_OPERATOR_DIAGONAL_MATRIX1: return "DML_OPERATOR_DIAGONAL_MATRIX1"; default: assert(false); return ""; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h index a993300291..42de619a87 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLSchema.h @@ -2247,6 +2247,61 @@ constexpr DML_OPERATOR_SCHEMA DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_SCHEMA { DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_SCHEMA_FIELDS, }; +constexpr DML_SCHEMA_FIELD DML_RESAMPLE2_OPERATOR_SCHEMA_FIELDS[8]{ + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "InterpolationMode", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "RoundingDirection", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "DimensionCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY, "Scales", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY, "InputPixelOffsets", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY, "OutputPixelOffsets", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_RESAMPLE2_OPERATOR_SCHEMA{ + "DML_OPERATOR_RESAMPLE2", + DML_OPERATOR_RESAMPLE2, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 8, + DML_RESAMPLE2_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA_FIELDS[8]{ + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputGradientTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputGradientTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "InterpolationMode", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "RoundingDirection", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "DimensionCount", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY, "Scales", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY, "InputPixelOffsets", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT_ARRAY, "OutputPixelOffsets", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA{ + "DML_OPERATOR_RESAMPLE_GRAD1", + DML_OPERATOR_RESAMPLE_GRAD1, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 8, + DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA_FIELDS, +}; + +constexpr DML_SCHEMA_FIELD DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA_FIELDS[6]{ + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", true }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "ValueDataType", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_SCALAR_UNION, "Value", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_INT, "DiagonalFillBegin", false }, + DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_INT, "DiagonalFillEnd", false }, +}; + +constexpr DML_OPERATOR_SCHEMA DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA{ + "DML_OPERATOR_DIAGONAL_MATRIX1", + DML_OPERATOR_DIAGONAL_MATRIX1, + DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE, + 6, + DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA_FIELDS, +}; + constexpr DML_SCHEMA_FIELD DML_ACTIVATION_ELU_OPERATOR_SCHEMA_FIELDS[3] { DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false }, DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false }, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h index b37389ce17..aaf02ca146 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/GeneratedSchemaHelpers.h @@ -1381,6 +1381,43 @@ inline std::vector GetFields(const DML_BATCH_NORMALIZATION_TRAINI OperatorField(&DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_SCHEMA.Fields[8], ToOperatorFieldType(static_cast(desc.FusedActivation))), }; } +inline std::vector GetFields(const DML_RESAMPLE2_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.InterpolationMode))), + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.RoundingDirection))), + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.DimensionCount))), + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[5], ToOperatorFieldType(static_cast(desc.Scales), desc.DimensionCount)), + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[6], ToOperatorFieldType(static_cast(desc.InputPixelOffsets), desc.DimensionCount)), + OperatorField(&DML_RESAMPLE2_OPERATOR_SCHEMA.Fields[7], ToOperatorFieldType(static_cast(desc.OutputPixelOffsets), desc.DimensionCount)), + }; +} +inline std::vector GetFields(const DML_RESAMPLE_GRAD1_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputGradientTensor))), + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputGradientTensor))), + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.InterpolationMode))), + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.RoundingDirection))), + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.DimensionCount))), + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[5], ToOperatorFieldType(static_cast(desc.Scales), desc.DimensionCount)), + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[6], ToOperatorFieldType(static_cast(desc.InputPixelOffsets), desc.DimensionCount)), + OperatorField(&DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA.Fields[7], ToOperatorFieldType(static_cast(desc.OutputPixelOffsets), desc.DimensionCount)), + }; +} +inline std::vector GetFields(const DML_DIAGONAL_MATRIX1_OPERATOR_DESC& desc) +{ + return { + OperatorField(&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast(desc.InputTensor))), + OperatorField(&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast(desc.OutputTensor))), + OperatorField(&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast(desc.ValueDataType))), + OperatorField(&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast(desc.Value))), + OperatorField(&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast(desc.DiagonalFillBegin))), + OperatorField(&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA.Fields[5], ToOperatorFieldType(static_cast(desc.DiagonalFillEnd))), + }; +} inline std::vector GetFields(const DML_ACTIVATION_ELU_OPERATOR_DESC& desc) { return { @@ -1713,6 +1750,9 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType) case DML_OPERATOR_ELEMENT_WISE_QUANTIZED_LINEAR_ADD: return DML_ELEMENT_WISE_QUANTIZED_LINEAR_ADD_OPERATOR_SCHEMA; case DML_OPERATOR_ROI_ALIGN_GRAD: return DML_ROI_ALIGN_GRAD_OPERATOR_SCHEMA; case DML_OPERATOR_BATCH_NORMALIZATION_TRAINING: return DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_SCHEMA; + case DML_OPERATOR_RESAMPLE2: return DML_RESAMPLE2_OPERATOR_SCHEMA; + case DML_OPERATOR_RESAMPLE_GRAD1: return DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA; + case DML_OPERATOR_DIAGONAL_MATRIX1: return DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_ELU: return DML_ACTIVATION_ELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_CELU: return DML_ACTIVATION_CELU_OPERATOR_SCHEMA; case DML_OPERATOR_ACTIVATION_HARDMAX: return DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA; @@ -2286,7 +2326,7 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_ELEMENT_WISE_QUANTIZED_LINEAR_ADD_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); - case DML_OPERATOR_ROI_ALIGN_GRAD: + case DML_OPERATOR_ROI_ALIGN_GRAD: return AbstractOperatorDesc( &DML_ROI_ALIGN_GRAD_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); @@ -2294,6 +2334,18 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc) return AbstractOperatorDesc( &DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_SCHEMA, GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_RESAMPLE2: + return AbstractOperatorDesc( + &DML_RESAMPLE2_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_RESAMPLE_GRAD1: + return AbstractOperatorDesc( + &DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); + case DML_OPERATOR_DIAGONAL_MATRIX1: + return AbstractOperatorDesc( + &DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA, + GetFields(*static_cast(opDesc.Desc))); case DML_OPERATOR_ACTIVATION_ELU: return AbstractOperatorDesc( &DML_ACTIVATION_ELU_OPERATOR_SCHEMA, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp index 7a61f56eb6..6c1b670502 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp @@ -1060,7 +1060,8 @@ HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::SetDmlOperator( SetDmlProperties(dmlProperties); m_graphNodeCreateInfo->op = op; - m_graphNodeCreateInfo->desc = std::make_unique(SchemaHelpers::ConvertOperatorDesc(*desc)); + AbstractOperatorDesc abstractDesc = SchemaHelpers::ConvertOperatorDesc(*desc); + m_graphNodeCreateInfo->desc = std::make_unique(std::move(abstractDesc)); return S_OK; } diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp index 5f1a116802..d97dbc3d4a 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp @@ -63,6 +63,15 @@ public: } }; +void CALLBACK QueryBatchNormalization(IMLOperatorSupportQueryContextPrivate* context, /*out*/ bool* isSupported) +{ + // training_mode=1 is unsupported as it isn't needed for inference (https://github.com/onnx/onnx/pull/3333). + + MLOperatorAttributes attributes(context); + int32_t trainingMode = attributes.GetOptionalAttribute(AttrName::TrainingMode, 0); + *isSupported = (trainingMode == 0); +} + DML_OP_DEFINE_CREATION_FUNCTION(BatchNormalization, DmlOperatorBatchNormalization); DML_OP_DEFINE_CREATION_FUNCTION(FusedBatchNormalization, DmlOperatorBatchNormalization); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorEyeLike.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorEyeLike.cpp index f5cf913571..ef67ddccda 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorEyeLike.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorEyeLike.cpp @@ -24,7 +24,7 @@ public: assert(inputDescs.size() <= 1); assert(outputDescs.size() == 1); - auto outputTensorShapeDescription = kernelCreationContext.GetTensorShapeDescription();; + auto outputTensorShapeDescription = kernelCreationContext.GetTensorShapeDescription(); std::vector outputDimensions = outputTensorShapeDescription.GetOutputTensorShape(0); ML_CHECK_VALID_ARGUMENT(outputDimensions.size() <= OperatorHelper::NchwDimensionCount); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRecurrentNeuralNetwork.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRecurrentNeuralNetwork.cpp index bca338e892..88b827f61f 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRecurrentNeuralNetwork.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorRecurrentNeuralNetwork.cpp @@ -1,7 +1,7 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include "precomp.h" +#include "./precomp.h" namespace Dml { @@ -13,7 +13,7 @@ class DmlOperatorRecurrentBase: public DmlOperator, public RecurrentHelper public: using Self = DmlOperatorRecurrentBase; - DmlOperatorRecurrentBase(const MLOperatorKernelCreationContext& kernelInfo): + explicit DmlOperatorRecurrentBase(const MLOperatorKernelCreationContext& kernelInfo): DmlOperator(kernelInfo), RecurrentHelper(kernelInfo, kernelInfo.GetTensorShapeDescription()) { @@ -404,6 +404,28 @@ private: }; }; +void CALLBACK QueryRecurrentNeuralNetwork(IMLOperatorSupportQueryContextPrivate* context, /*out*/ bool* isSupported) +{ + // layout=1 for batchwise operation is unsupported, added in opset 14 for RNN, GRU, and LSTM + // (https://github.com/onnx/onnx/pull/3217, https://github.com/onnx/onnx/pull/2284). + // Currently (2022-05-27) the ORT CPU execution provider (lstm_base.h) does not support it either, + // with no models warranting it. When needed, it can be achieved with no new DML API's by just + // swapping the size and strides in the TensorDesc before filling in the *_OPERATOR_DESC, where: + // + // layout=0: (default, consistent with opset 7) + // X.shape = [seq_length, batch_size, input_size] + // Y.shape = [seq_length, num_directions, batch_size, hidden_size] + // initial_h.shape = Y_h.shape = initial_c.shape = Y_c.shape = [num_directions, batch_size, hidden_size] + // layout=1: + // X.shape = [batch_size, seq_length, input_size] + // Y.shape = [batch_size, seq_length, num_directions, hidden_size] + // initial_h.shape = Y_h.shape = initial_c.shape = Y_c.shape = [batch_size, num_directions, hidden_size] + + MLOperatorAttributes attributes(context); + int32_t layout = attributes.GetOptionalAttribute(AttrName::Layout, 0); + *isSupported = (layout == 0); +} + DML_OP_DEFINE_CREATION_FUNCTION(RNN, DmlOperatorRecurrentNeuralNetwork); DML_OP_DEFINE_CREATION_FUNCTION(GRU, DmlOperatorGatedRecurrentUnit); DML_OP_DEFINE_CREATION_FUNCTION(LSTM, DmlOperatorLongShortTermUnit); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorResize.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorResize.cpp index 1eb7532742..68931a9317 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorResize.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorResize.cpp @@ -22,7 +22,7 @@ constexpr NameAndIndex nearestNeighborRoundingModes[] = {"round_prefer_floor", 0}, // round halves down {"round_prefer_ceil", 1}, // round halves up {"floor", 2}, // round always down - // {"ceil", 3}, // round always up (requires a DirectML API addition) + {"ceil", 3}, // round always up }; void ComputePixelOffsetsAndScales( @@ -212,14 +212,16 @@ public: ); // Find any useless dimensions of size 1 that occur in both input and output. + // This enables higher dimension cases (where models prepend unnecessary + // dimensions) beyond DML's supported dimension count of 4. for (size_t i = 0, rank = m_outputDimensions.size(); i < rank; ++i) { - if (m_inputDimensions[i] == 1 && m_outputDimensions[i] == 1) { squeezableDimensionIndices.push_back(gsl::narrow_cast(i)); } } + RemoveValuesByIndex(squeezableDimensionIndices, /*keepOneValue*/ true, /*inout*/ squeezedInputShape); RemoveValuesByIndex(squeezableDimensionIndices, /*keepOneValue*/ true, /*inout*/ paddedScales); RemoveValuesByIndex(squeezableDimensionIndices, /*keepOneValue*/ true, /*inout*/ inputPixelOffsets); @@ -248,48 +250,56 @@ public: std::string mode = kernelCreationContext.GetOptionalAttribute(AttrName::Mode, "NEAREST"); DML_INTERPOLATION_MODE interpolationMode = Dml::MapStringToInteropolationMode(mode); - // DML's nearest neighbor mode uses round-halves-up (or round_prefer_ceil) via floor(input.x + 0.5). - // So to support floor, adjust the input by half a pixel. - // round_prefer_floor is not supported without an API extension, - // but existing code already default to treating it as round_prefer_ceil. - // So continue that. + // Map ONNX to DML's mode using offsets and rounding direction. + // These offsets are in addition to the coordinate transform offsets. + DML_AXIS_DIRECTION roundingDirection = DML_AXIS_DIRECTION_DECREASING; if (interpolationMode == DML_INTERPOLATION_MODE_NEAREST_NEIGHBOR) { std::string nearestMode = kernelCreationContext.GetOptionalAttribute(AttrName::NearestMode, "round_prefer_floor"); + float offsetAdjustment = 0.5f; auto optionalNearestModeValue = TryMapStringToIndex(nearestMode, nearestNeighborRoundingModes); if (optionalNearestModeValue) { + // The round_prefer_floor mode rounds values to the nearest integer, with half ties rounded toward + // negative infinity. The increasing rounding direction is correct, albeit unintuitive, because + // floor(x + 0.5) would return the wrong result, whereas the correct implementation is ceil(x - 0.5). + // The input offset is positive because positive input offsets translate the output rightward and + // downward, which (from the perspective of the output) is equivalent to panning the input + // toward further negative coordinates. switch (*optionalNearestModeValue) { - case 0: // round_prefer_floor - case 1: // round_prefer_ceil - break; - case 2: // floor - for (auto& offset : inputPixelOffsets) - { - offset += 0.5; - } - break; + case 0: /*round_prefer_floor*/ roundingDirection = DML_AXIS_DIRECTION_INCREASING; offsetAdjustment = 0.5; break; + case 1: /*round_prefer_ceil */ roundingDirection = DML_AXIS_DIRECTION_DECREASING; offsetAdjustment = -0.5; break; + case 2: /*floor */ roundingDirection = DML_AXIS_DIRECTION_DECREASING; offsetAdjustment = 0.0; break; + case 3: /*ceil */ roundingDirection = DML_AXIS_DIRECTION_INCREASING; offsetAdjustment = 0.0; break; default: assert(false); } } + if (offsetAdjustment != 0.0f) + { + for (auto& offset : inputPixelOffsets) + { + offset += offsetAdjustment; + } + } } // Create the operator description. std::vector inputDescs = GetDmlInputDescs(); std::vector outputDescs = GetDmlOutputDescs(); - DML_RESAMPLE1_OPERATOR_DESC operatorDesc = {}; + DML_RESAMPLE2_OPERATOR_DESC operatorDesc = {}; operatorDesc.InputTensor = inputDescs.data(); operatorDesc.OutputTensor = outputDescs.data(); operatorDesc.InterpolationMode = interpolationMode; + operatorDesc.RoundingDirection = roundingDirection; operatorDesc.Scales = paddedScales.data(); operatorDesc.DimensionCount = gsl::narrow_cast(paddedScales.size()); operatorDesc.InputPixelOffsets = inputPixelOffsets.data(); operatorDesc.OutputPixelOffsets = outputPixelOffsets.data(); - DML_OPERATOR_DESC opDesc = { DML_OPERATOR_RESAMPLE1, &operatorDesc }; + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_RESAMPLE2, &operatorDesc }; SetDmlOperatorDesc(opDesc, kernelCreationContext); } }; @@ -323,14 +333,6 @@ void CALLBACK QueryResize(IMLOperatorSupportQueryContextPrivate* context, bool* return; } - // DML's nearest neighbor mode uses half pixels rounded down. - std::string nearestMode = attributes.GetOptionalAttribute(AttrName::NearestMode, "round_prefer_floor"); - auto optionalNearestModeValue = TryMapStringToIndex(nearestMode, nearestNeighborRoundingModes); - if (!optionalNearestModeValue) - { - return; - } - // Ignore parameter "cubic_coeff_a" since Cubic interpolation unsupported in DML. // Ignore parameter "extrapolation_value" as DML clamps to the input rather than reading black pixels. diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorTrilu.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorTrilu.cpp new file mode 100644 index 0000000000..0d1350d926 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorTrilu.cpp @@ -0,0 +1,56 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "./precomp.h" + +namespace Dml +{ + +class DmlOperatorTrilu : public DmlOperator +{ +public: + explicit DmlOperatorTrilu(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext) + { + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() >= 1, "Trilu expects 1-2 inputs."); + ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() == 1, "Trilu expects 1 output."); + + std::vector> inputIndices = {0}; // Use only the first tensor. The second tensor is CPU-based (k). + std::vector> outputIndices = {0}; + DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + assert(inputDescs.size() == 1); + assert(outputDescs.size() == 1); + + // Read the diagonal offset from the 2nd tensor (defaults to 0 if absent). + int32_t k = 0; + if (kernelCreationContext.IsInputValid(1)) + { + MLOperatorTensor kTensor = kernelCreationContext.GetConstantInputTensor(1); + k = gsl::narrow_cast(ReadScalarTensorCastToInt64(kTensor)); + } + + auto outputTensorShapeDescription = kernelCreationContext.GetTensorShapeDescription(); + std::vector outputDimensions = outputTensorShapeDescription.GetOutputTensorShape(0); + ML_CHECK_VALID_ARGUMENT(outputDimensions.size() <= OperatorHelper::NchwDimensionCount); + + const bool keepUpperDiagonal = kernelCreationContext.GetOptionalAttribute(AttrName::Upper, 0); + + DML_DIAGONAL_MATRIX1_OPERATOR_DESC operatorDesc = {}; + operatorDesc.InputTensor = inputDescs.data(); + operatorDesc.OutputTensor = outputDescs.data(); + operatorDesc.DiagonalFillBegin = keepUpperDiagonal ? INT32_MIN : k + 1; + operatorDesc.DiagonalFillEnd = keepUpperDiagonal ? k : INT32_MAX; + operatorDesc.ValueDataType = m_inputTensorDescs[0].GetDmlDataType(); + CastToClampedScalarUnion(operatorDesc.ValueDataType, 0.0f, /*out*/&operatorDesc.Value); + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_DIAGONAL_MATRIX1, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +DML_OP_DEFINE_CREATION_FUNCTION(Trilu, DmlOperatorTrilu); + +} // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index 7c40dae020..d0198cadb6 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -255,11 +255,14 @@ DML_OP_EXTERN_CREATION_FUNCTION(QLinearMatMul); DML_OP_EXTERN_CREATION_FUNCTION(DynamicQuantizeLinear); DML_OP_EXTERN_CREATION_FUNCTION(MatMulInteger); DML_OP_EXTERN_CREATION_FUNCTION(ConvInteger); +DML_OP_EXTERN_CREATION_FUNCTION(Trilu); DML_OP_EXTERN_QUERY_FUNCTION(MaxPool); DML_OP_EXTERN_QUERY_FUNCTION(Slice); DML_OP_EXTERN_QUERY_FUNCTION(Resize); DML_OP_EXTERN_QUERY_FUNCTION(EinSum); +DML_OP_EXTERN_QUERY_FUNCTION(RecurrentNeuralNetwork); +DML_OP_EXTERN_QUERY_FUNCTION(BatchNormalization); constexpr static std::array typeNameListDefault = {"T"}; constexpr static std::array typeNameListTwo = { "T1", "T2" }; @@ -274,7 +277,7 @@ constexpr static std::array typeNameListScatterGather = { "T", " constexpr static std::array typeNameListScatterGatherND = { "T" }; // Tind is curiously missing, only allowing 64-bit. constexpr static std::array typeNameListSlice10 = { "T", "Tind" }; constexpr static std::array typeNameListWhere = { "B", "T" }; -constexpr static std::array typeNameListEyeLike = { "T2" }; +constexpr static std::array typeNameListEyeLike = { "T1", "T2" }; constexpr static std::array supportedTypeListAll = {SupportedTensorDataTypes::All}; constexpr static std::array supportedTypeListFloat32 = {SupportedTensorDataTypes::Float32}; @@ -287,7 +290,8 @@ constexpr static std::array supportedTypeListFloat1 constexpr static std::array supportedTypeListFloat16to32Ints32to64 = {SupportedTensorDataTypes::Float16to32 | SupportedTensorDataTypes::Ints32Bit | SupportedTensorDataTypes::Ints64Bit}; constexpr static std::array supportedTypeListUInt8to64 = {SupportedTensorDataTypes::UInt8to64}; constexpr static std::array supportedTypeListNumericDefault = { SupportedTensorDataTypes::NumericDefault }; -constexpr static std::array supportedTypeListAllScalars = { SupportedTensorDataTypes::AllScalars }; +constexpr static std::array supportedTypeListAllScalars = {SupportedTensorDataTypes::AllScalars}; +constexpr static std::array supportedTypeListEyeLike = { SupportedTensorDataTypes::AllScalars, SupportedTensorDataTypes::AllScalars}; constexpr static std::array supportedTypeListBool = {SupportedTensorDataTypes::Bool}; constexpr static std::array supportedTypeListPow12 = {SupportedTensorDataTypes::Int32 | SupportedTensorDataTypes::Float16to32, SupportedTensorDataTypes::NumericDefault}; constexpr static std::array supportedTypeListTopK = {SupportedTensorDataTypes::NumericDefault | SupportedTensorDataTypes::Ints64Bit, SupportedTensorDataTypes::Int64}; @@ -393,11 +397,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO_VER( 10, RoiAlign, typeNameListTwo, supportedTypeListRoiAlign, DmlGraphSupport::Supported)}, {REG_INFO( 7, InstanceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 9, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. + {REG_INFO( 9, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. + {REG_INFO( 14, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v14 adds training_mode attribute // TODO: Add additional type constraints in BatchNormalization-15, with scale and bias (T1) being different from input X (T). - // Add training-mode support to BatchNormalization-15 https://github.com/onnx/onnx/pull/3333 - // {REG_INFO( 14, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. - // {REG_INFO( 15, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. + // {REG_INFO( 15, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v15 adds differing types for scale and bias vs input. {REG_INFO( 7, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 13, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, MeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -405,27 +408,24 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 13, MeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, LpNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, RNN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, - // TODO: Allow recurrent operations to be batchwise https://github.com/onnx/onnx/pull/3217 - // {REG_INFO( 14, RNN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, + {REG_INFO( 14, RNN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(), std::nullopt, QueryRecurrentNeuralNetwork)}, {REG_INFO( 7, GRU, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, - // TODO: Allow recurrent operations to be batchwise https://github.com/onnx/onnx/pull/3217 - // {REG_INFO( 14, GRU, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, + {REG_INFO( 14, GRU, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(), std::nullopt, QueryRecurrentNeuralNetwork)}, {REG_INFO( 7, LSTM, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, - // TODO: Allow recurrent operations to be batchwise https://github.com/onnx/onnx/pull/3217 - // {REG_INFO( 14, LSTM, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, + {REG_INFO( 14, LSTM, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(), std::nullopt, QueryRecurrentNeuralNetwork)}, {REG_INFO_MS( 1, ConvTransposeWithDynamicPads, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(2))}, // Data Reorganization Layers {REG_INFO_VER( 7, Split, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_VER( 11, Split, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, // Adds negative axis. - {REG_INFO_VER( 13, Split, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, // Moves splits from constant parameter to dynamic input. + {REG_INFO_VER( 11, Split, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, // Adds negative axis. + {REG_INFO_VER( 13, Split, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, // Moves splits from constant parameter to dynamic input. {REG_INFO( 7, Transpose, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, {REG_INFO( 13, Transpose, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, {REG_INFO( 7, Concat, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO( 11, Concat, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, // Adds negative axis. - {REG_INFO( 13, Concat, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, // Adds negative axis. + {REG_INFO( 11, Concat, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, // Adds negative axis. + {REG_INFO( 13, Concat, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, // Adds negative axis. {REG_INFO_VER( 7, Slice, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, - {REG_INFO_VER( 10, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3, 4), std::nullopt, QuerySlice)}, // Adds negative axes. + {REG_INFO_VER( 10, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3, 4), std::nullopt, QuerySlice)}, // Adds negative axes. {REG_INFO_VER( 11, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3, 4), std::nullopt, QuerySlice)}, {REG_INFO_VER( 13, Slice, typeNameListSlice10, supportedTypeListSlice10, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3, 4), std::nullopt, QuerySlice)}, {REG_INFO_VER( 7, Pad, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, @@ -456,9 +456,8 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 13, ScatterElements, typeNameListScatterGather, supportedTypeListScatterGather, DmlGraphSupport::Supported)}, {REG_INFO( 11, ScatterND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmlGraphSupport::Supported)}, {REG_INFO( 13, ScatterND, typeNameListScatterGatherND, supportedTypeListScatterGatherND, DmlGraphSupport::Supported)}, - {REG_INFO( 9, EyeLike, typeNameListEyeLike, supportedTypeListScalars8to32, DmlGraphSupport::Supported)}, - // TODO: Add Trilu-14 to fill diagonal matrix https://github.com/onnx/onnx/pull/3291 - // {REG_INFO( 14, Trilu, typeNameListTrilu, supportedTypeListScalars8to32, DmlGraphSupport::Supported)}, + {REG_INFO( 9, EyeLike, typeNameListEyeLike, supportedTypeListEyeLike, DmlGraphSupport::Supported)}, + {REG_INFO( 14, Trilu, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported, requiredConstantCpuInputs(1))}, // Data reorganization that merely changes the dimensions while keeping the data identical. {REG_INFO_COPY( 7, Identity, typeNameListDefault, supportedTypeListAllScalars, DmlGraphSupport::Supported)}, @@ -485,7 +484,8 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 13, Reciprocal, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Pow, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 12, Pow, typeNameListPow12, supportedTypeListPow12, DmlGraphSupport::Supported)}, - {REG_INFO( 13, Pow, typeNameListPow12, supportedTypeListPow12, DmlGraphSupport::Supported)}, + {REG_INFO( 13, Pow, typeNameListPow12, supportedTypeListPow12, DmlGraphSupport::Supported)}, // 13 added bfloat16 to T. + {REG_INFO( 15, Pow, typeNameListPow12, supportedTypeListPow12, DmlGraphSupport::Supported)}, // 15 added bfloat16 to T1. {REG_INFO( 7, Exp, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 13, Exp, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, Log, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, @@ -622,7 +622,6 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO_VER( 10, Upsample, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1) /*scales*/)}, {REG_INFO_VER( 10, Resize, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(1) /*scales*/)}, {REG_INFO_VER( 11, Resize, typeNameListTwo, supportedTypeListResize11, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3) /*roi, scales, sizes*/, std::nullopt, QueryResize)}, - // TODO: Resize-13 support nearest rounding mode attribute https://github.com/onnx/onnx/pull/3026 {REG_INFO_VER( 13, Resize, typeNameListTwo, supportedTypeListResize13, DmlGraphSupport::Supported, requiredConstantCpuInputs(1, 2, 3) /*roi, scales, sizes*/, std::nullopt, QueryResize)}, // Activation Functions @@ -699,7 +698,7 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 11, Range, typeNameListDefault, supportedTypeListRange, DmlGraphSupport::Supported, requiredConstantCpuInputs(0,1,2))}, {REG_INFO( 9, MaxUnpool, typeNameListTwo, supportedTypeListMaxUnpool, DmlGraphSupport::Supported, requiredConstantCpuInputs(2))}, - {REG_INFO( 11, MaxUnpool, typeNameListTwo, supportedTypeListMaxUnpool, DmlGraphSupport::Supported, requiredConstantCpuInputs(2))}, // 11 is identical to 9. + {REG_INFO( 11, MaxUnpool, typeNameListTwo, supportedTypeListMaxUnpool, DmlGraphSupport::Supported, requiredConstantCpuInputs(2))}, // 11 is identical to 9. {REG_INFO_MS( 1, QLinearAdd, typeNameListDefault, supportedTypeListInteger8, DmlGraphSupport::Supported)}, {REG_INFO( 10, QLinearConv, typeNameListFour, supportedTypeListQLinearConv, DmlGraphSupport::Supported)}, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp index 2412a41d16..88070d5543 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp @@ -148,6 +148,7 @@ namespace Dml OperatorInfo{ "ConvTranspose", onnxruntime::kOnnxDomain, OnnxOperatorSet11::sc_sinceVer_ConvTranspose }, OperatorInfo{ "BatchNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_BatchNormalization }, OperatorInfo{ "BatchNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet9::sc_sinceVer_BatchNormalization }, + OperatorInfo{ "BatchNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet14::sc_sinceVer_BatchNormalization }, OperatorInfo{ "InstanceNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_InstanceNormalization }, OperatorInfo{ "MeanVarianceNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_MeanVarianceNormalization }, OperatorInfo{ "MeanVarianceNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet9::sc_sinceVer_MeanVarianceNormalization }, @@ -448,7 +449,7 @@ namespace Dml return *index; } ML_INVALID_ARGUMENT("Unknown interpolation mode"); - return (DML_INTERPOLATION_MODE)0; + return static_cast(0); } #pragma warning(pop) @@ -466,7 +467,7 @@ namespace Dml return *index; } ML_INVALID_ARGUMENT("Unknown depth/space order"); - return (DML_DEPTH_SPACE_ORDER)0; + return static_cast(0); } #pragma warning(pop) diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h index 99e7d1d390..86540ac30b 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/Attributes.h @@ -52,6 +52,7 @@ namespace AttrName static constexpr const char* LinearBeforeReset = "linear_before_reset"; static constexpr const char* Lambda = "lambd"; // Deliberate typo to match ONNX spec. static constexpr const char* Largest = "largest"; + static constexpr const char* Layout = "layout"; static constexpr const char* Low = "low"; static constexpr const char* Max = "max"; static constexpr const char* Mean = "mean"; @@ -87,8 +88,10 @@ namespace AttrName static constexpr const char* Tiles = "tiles"; static constexpr const char* TimeAxis = "time_axis"; static constexpr const char* To = "to"; + static constexpr const char* TrainingMode = "training_mode"; static constexpr const char* TransA = "transA"; static constexpr const char* TransB = "transB"; + static constexpr const char* Upper = "upper"; static constexpr const char* Value = "value"; static constexpr const char* WidthScale = "width_scale"; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index daceecab5c..f9c4f5eae9 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -1361,6 +1361,7 @@ using ShapeInferenceHelper_Unsqueeze7 = VersionedOpsetHelper using ShapeInferenceHelper_Unsqueeze11 = VersionedOpsetHelper; using ShapeInferenceHelper_Unsqueeze13 = VersionedOpsetHelper; using ShapeInferenceHelper_EyeLike = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_Trilu = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_Expand = ExpandHelper; using ShapeInferenceHelper_Reshape7 = ReshapeHelper; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h index 598d282bf2..24d3c464fd 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorVersions.h @@ -344,18 +344,25 @@ namespace OperatorHelper namespace OnnxOperatorSet14 { static const int sc_sinceVer_Add = 14; + static const int sc_sinceVer_BatchNormalization = 14; static const int sc_sinceVer_CumSum = 14; static const int sc_sinceVer_Div = 14; static const int sc_sinceVer_Identity = 14; + static const int sc_sinceVer_GRU = 14; + static const int sc_sinceVer_LSTM = 14; static const int sc_sinceVer_Mul = 14; static const int sc_sinceVer_Relu = 14; static const int sc_sinceVer_Reshape = 14; + static const int sc_sinceVer_RNN = 14; static const int sc_sinceVer_Sub = 14; + static const int sc_sinceVer_Trilu = 14; } // namespace OnnxOperatorSet14 namespace OnnxOperatorSet15 { static const int sc_sinceVer_CastLike = 15; + static const int sc_sinceVer_BatchNormalization = 15; + static const int sc_sinceVer_Pow = 15; } // namespace OnnxOperatorSet14 namespace MsftOperatorSet1 From c1b5f343621bdc63440084b7054992ead0510a5a Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Fri, 10 Jun 2022 15:04:48 -0700 Subject: [PATCH 10/19] DML EP BatchNormalization-15 (#11814) * Add external helper DirectMLX.h * Add BatchNormalization-15 using DMLX to achieve casting if types are different * Shape helper and some reformatting * Additional linting issues --- onnxruntime/core/providers/dml/CPPLINT.cfg | 2 +- .../inc/IWinmlExecutionProvider.h | 2 +- .../src/AbiCustomRegistry.cpp | 15 +- .../dml/DmlExecutionProvider/src/DmlCommon.h | 74 - .../DmlExecutionProvider/src/ErrorHandling.h | 2 +- .../src/External/CPPLINT.cfg | 1 + .../src/External/DirectMLHelpers/DirectMLX.h | 4050 +++++++++++++++++ .../src/GraphDescBuilder.cpp | 3 +- .../src/GraphPartitioner.cpp | 4 +- .../src/MLOperatorAuthorImpl.cpp | 3959 ++++++++-------- .../src/MLOperatorAuthorImpl.h | 145 +- .../src/Operators/DmlOperator.cpp | 4 +- .../DmlOperatorBatchNormalization.cpp | 69 +- .../src/Operators/DmlOperatorSlice.cpp | 7 +- .../src/Operators/OperatorRegistration.cpp | 10 +- .../src/Operators/OperatorUtility.cpp | 1 + .../dml/DmlExecutionProvider/src/precomp.h | 1 + .../OperatorAuthorHelper/OperatorHelper.cpp | 23 + .../dml/OperatorAuthorHelper/OperatorHelper.h | 20 +- 19 files changed, 6353 insertions(+), 2039 deletions(-) create mode 100644 onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/CPPLINT.cfg create mode 100644 onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLX.h diff --git a/onnxruntime/core/providers/dml/CPPLINT.cfg b/onnxruntime/core/providers/dml/CPPLINT.cfg index e7dbd3164b..02d14c65cc 100644 --- a/onnxruntime/core/providers/dml/CPPLINT.cfg +++ b/onnxruntime/core/providers/dml/CPPLINT.cfg @@ -1 +1 @@ -filter=-whitespace/braces,-whitespace/parens,-whitespace/line_length,-whitespace/indent +filter=-whitespace/braces,-whitespace/parens,-whitespace/line_length,-whitespace/indent,-whitespace/newline diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/IWinmlExecutionProvider.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/IWinmlExecutionProvider.h index 781930a186..073c0ebf62 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/IWinmlExecutionProvider.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/inc/IWinmlExecutionProvider.h @@ -92,7 +92,7 @@ namespace Windows::AI::MachineLearning::Adapter const onnxruntime::Node& node, MLOperatorTensorGetter& constantInputGetter, const void* executionHandle, - DmlGraphNodeCreateInfo* graphNodeCreateInfo + /*out*/ DmlGraphNodeCreateInfo* graphNodeCreateInfo )>; struct GraphNodeFactoryRegistration diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/AbiCustomRegistry.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/AbiCustomRegistry.cpp index 33df3d6df4..255659521d 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/AbiCustomRegistry.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/AbiCustomRegistry.cpp @@ -457,7 +457,7 @@ HRESULT STDMETHODCALLTYPE AbiCustomRegistry::RegisterOperatorKernel( constantCpuInputCapture, shapeInferrerCapture.Get(), &defaultAttributesCapture); - return Status::OK(); + return Status::OK(); }; onnxruntime::KernelCreateInfo create_info(builder.Build(), lotusKernelCreateFn); @@ -472,11 +472,18 @@ HRESULT STDMETHODCALLTYPE AbiCustomRegistry::RegisterOperatorKernel( if (supportsGraph) { GraphNodeFactoryRegistration graphReg; - graphReg.factory = - [kernelFactoryCapture, + graphReg.factory = [ + kernelFactoryCapture, shapeInferrerCapture, defaultAttributesCapture, - constantCpuInputCapture](const onnxruntime::Node& node, MLOperatorTensorGetter& constantInputGetter, const void* executionHandle, DmlGraphNodeCreateInfo* graphNodeCreateInfo) + constantCpuInputCapture + ] + ( + const onnxruntime::Node& node, + MLOperatorTensorGetter& constantInputGetter, + const void* executionHandle, + /*out*/ DmlGraphNodeCreateInfo* graphNodeCreateInfo + ) { onnxruntime::ProtoHelperNodeContext nodeContext(node); onnxruntime::OpNodeProtoHelper protoHelper(&nodeContext); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h index 0c3abcd255..9bf8c58f7a 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/DmlCommon.h @@ -22,80 +22,6 @@ namespace Dml bool IsSigned(DML_TENSOR_DATA_TYPE dataType); - /** Calculates the minimum number of bytes required to store a buffer tensor with the specified type, sizes, and - strides. The formula can be expressed as the following: - - IndexOfLastElement = dot(Sizes - 1, Strides); - MinimumImpliedSizeInBytes = roundup((IndexOfLastElement + 1) * ElementSizeInBytes, 4) - - In other words, the minimum size of a tensor is the index of the one-past-the-end element, multiplied by the - element size (e.g. 2 bytes for a FLOAT16 tensor). Additionally DirectML requires that all buffers bound must have - a total size which is DWORD-aligned, and hence the minimum implied size in bytes must be rounded up to the nearest - 4-byte boundary. - */ - inline UINT64 DMLCalcBufferTensorSize( - DML_TENSOR_DATA_TYPE dataType, - UINT dimensionCount, - _In_reads_(dimensionCount) const UINT* sizes, - _In_reads_opt_(dimensionCount) const UINT* strides) - { - UINT elementSizeInBytes = 0; - switch (dataType) - { - case DML_TENSOR_DATA_TYPE_FLOAT64: - case DML_TENSOR_DATA_TYPE_UINT64: - case DML_TENSOR_DATA_TYPE_INT64: - elementSizeInBytes = 8; - break; - - case DML_TENSOR_DATA_TYPE_FLOAT32: - case DML_TENSOR_DATA_TYPE_UINT32: - case DML_TENSOR_DATA_TYPE_INT32: - elementSizeInBytes = 4; - break; - - case DML_TENSOR_DATA_TYPE_FLOAT16: - case DML_TENSOR_DATA_TYPE_UINT16: - case DML_TENSOR_DATA_TYPE_INT16: - elementSizeInBytes = 2; - break; - - case DML_TENSOR_DATA_TYPE_UINT8: - case DML_TENSOR_DATA_TYPE_INT8: - elementSizeInBytes = 1; - break; - - default: - return 0; // Invalid data type - } - - UINT64 minimumImpliedSizeInBytes = 0; - if (!strides) - { - minimumImpliedSizeInBytes = sizes[0]; - for (UINT i = 1; i < dimensionCount; ++i) - { - minimumImpliedSizeInBytes *= sizes[i]; - } - minimumImpliedSizeInBytes *= elementSizeInBytes; - } - else - { - UINT indexOfLastElement = 0; - for (UINT i = 0; i < dimensionCount; ++i) - { - indexOfLastElement += (sizes[i] - 1) * strides[i]; - } - - minimumImpliedSizeInBytes = (indexOfLastElement + 1) * elementSizeInBytes; - } - - // Round up to the nearest 4 bytes. - minimumImpliedSizeInBytes = (minimumImpliedSizeInBytes + 3) & ~3ui64; - - return minimumImpliedSizeInBytes; - } - template void CastToClampedScalarUnion(DML_TENSOR_DATA_TYPE dataType, T value, DML_SCALAR_UNION* outputValue) { diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ErrorHandling.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ErrorHandling.h index c59ad56e43..c39343fc68 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ErrorHandling.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/ErrorHandling.h @@ -3,7 +3,7 @@ #pragma once #ifdef ORT_NO_EXCEPTIONS -#define ORT_CATCH_RETURN +#define ORT_CATCH_RETURN #else #define ORT_CATCH_RETURN CATCH_RETURN() #endif diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/CPPLINT.cfg b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/CPPLINT.cfg new file mode 100644 index 0000000000..bf14c49304 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/CPPLINT.cfg @@ -0,0 +1 @@ +filter=-whitespace/comments,-readability/todo,-whitespace/end_of_line,-runtime/indentation_namespace diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLX.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLX.h new file mode 100644 index 0000000000..3f56ad5a70 --- /dev/null +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/DirectMLX.h @@ -0,0 +1,4050 @@ +//********************************************************* +// +// Copyright (c) Microsoft. All rights reserved. +// This code is licensed under the MIT License (MIT). +// THIS CODE IS PROVIDED *AS IS* WITHOUT WARRANTY OF +// ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING ANY +// IMPLIED WARRANTIES OF FITNESS FOR A PARTICULAR +// PURPOSE, MERCHANTABILITY, OR NON-INFRINGEMENT. +// +//********************************************************* +// clang-format off + +#pragma once +#include "DirectML.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include // For Microsoft::WRL::ComPtr + +#if DMLX_USE_ABSEIL + #if __cpp_lib_span + #include + #endif +#elif __cplusplus >= 201703L && __has_include() + // stl optional is only available in cpp17 and above. + #include +#elif __has_include("dml_optional_extensions.h") + #include "dml_optional_extensions.h" + #define DMLX_OPTIONAL_EXTENDED +#endif + +/** Calculates the minimum number of bytes required to store a buffer tensor with the specified type, sizes, and + strides. The formula can be expressed as the following: + + IndexOfLastElement = dot(Sizes - 1, Strides); + MinimumImpliedSizeInBytes = roundup((IndexOfLastElement + 1) * ElementSizeInBytes, 4) + + In other words, the minimum size of a tensor is the index of the one-past-the-end element, multiplied by the + element size (e.g. 2 bytes for a FLOAT16 tensor). Additionally DirectML requires that all buffers bound must have + a total size which is DWORD-aligned, and hence the minimum implied size in bytes must be rounded up to the nearest + 4-byte boundary. + */ + +inline UINT64 DMLCalcBufferTensorSize( + DML_TENSOR_DATA_TYPE dataType, + UINT dimensionCount, + _In_reads_(dimensionCount) const UINT* sizes, + _In_reads_opt_(dimensionCount) const UINT* strides) +{ + UINT elementSizeInBytes = 0; + switch (dataType) + { + case DML_TENSOR_DATA_TYPE_FLOAT32: + case DML_TENSOR_DATA_TYPE_UINT32: + case DML_TENSOR_DATA_TYPE_INT32: + elementSizeInBytes = 4; + break; + + case DML_TENSOR_DATA_TYPE_FLOAT16: + case DML_TENSOR_DATA_TYPE_UINT16: + case DML_TENSOR_DATA_TYPE_INT16: + elementSizeInBytes = 2; + break; + + case DML_TENSOR_DATA_TYPE_UINT8: + case DML_TENSOR_DATA_TYPE_INT8: + elementSizeInBytes = 1; + break; + + case DML_TENSOR_DATA_TYPE_FLOAT64: + case DML_TENSOR_DATA_TYPE_UINT64: + case DML_TENSOR_DATA_TYPE_INT64: + elementSizeInBytes = 8; + break; + + default: + return 0; // Invalid data type + } + + UINT64 minimumImpliedSizeInBytes = 0; + if (!strides) + { + minimumImpliedSizeInBytes = 1; + for (UINT i = 0; i < dimensionCount; ++i) + { + minimumImpliedSizeInBytes *= sizes[i]; + } + minimumImpliedSizeInBytes *= elementSizeInBytes; + } + else + { + UINT indexOfLastElement = 0; + for (UINT i = 0; i < dimensionCount; ++i) + { + indexOfLastElement += (sizes[i] - 1) * strides[i]; + } + + minimumImpliedSizeInBytes = (indexOfLastElement + 1) * elementSizeInBytes; + } + + // Round up to the nearest 4 bytes. + minimumImpliedSizeInBytes = (minimumImpliedSizeInBytes + 3) & ~3ull; + + return minimumImpliedSizeInBytes; +} + +namespace dml +{ + namespace detail + { + // Provide non-member size() and data(). Defaults to standard library implementation (if available) +#if __cpp_lib_nonmember_container_access + template + constexpr auto size(const C& c) -> decltype(c.size()) + { + return std::size(c); + } + + template + constexpr std::size_t size(const T(&array)[N]) noexcept + { + return std::size(array); + } + + template + constexpr auto data(C& c) -> decltype(c.data()) + { + return std::data(c); + } + + template + constexpr T* data(T(&array)[N]) noexcept + { + return std::data(array); + } +#else + template + constexpr auto size(const C& c) -> decltype(c.size()) + { + return c.size(); + } + + template + constexpr std::size_t size(const T(&array)[N]) noexcept + { + return N; + } + + template + constexpr auto data(C& c) -> decltype(c.data()) + { + return c.data(); + } + + template + constexpr T* data(T(&array)[N]) noexcept + { + return array; + } +#endif + + template + class span + { + public: + span() = default; + + constexpr span(std::initializer_list i) : m_begin(i.begin()), m_end(i.end()) {} + constexpr span(T* begin, T* end) : m_begin(begin), m_end(end) {} + constexpr span(T* begin, size_t elementCount) : m_begin(begin), m_end(begin + elementCount) {} + + template + constexpr span(ContiguousContainer&& container) + : m_begin(dml::detail::data(container)), m_end(m_begin + dml::detail::size(container)) {} + + template + constexpr span(T(&a)[N]) noexcept : span(a, N) {} + + T* data() noexcept { return m_begin; } + T* begin() noexcept { return m_begin; } + T* end() noexcept { return m_end; } + T const* data() const noexcept { return m_begin; } + T const* begin() const noexcept { return m_begin; } + T const* end() const noexcept { return m_end; } + bool empty() const noexcept { return m_end == m_begin; } + size_t size() const noexcept { return m_end - m_begin; } + size_t size_bytes() const noexcept { return sizeof(T) * size(); } + T& operator[](size_t index) const noexcept { return m_begin[index]; } + span subspan(size_t index, size_t count) { return span(m_begin + index, m_begin + index + count); } + + protected: + T* m_begin = nullptr; + T* m_end = nullptr; + }; + } + +#if DMLX_USE_ABSEIL + template + using Optional = absl::optional; + + constexpr absl::nullopt_t NullOpt = absl::nullopt; + + template + using SmallVector = absl::InlinedVector; + + template + using Span = absl::Span; + + using absl::make_unique; +#else + #ifndef DMLX_OPTIONAL_EXTENDED + template + using Optional = std::optional; + constexpr std::nullopt_t NullOpt = std::nullopt; + #endif + + template + using SmallVector = std::vector; + + #if __cpp_lib_span + template + using Span = std::span; + #elif DMLX_USE_GSL + template + using Span = gsl::span; + #else + template + using Span = dml::detail::span; + #endif + + using std::make_unique; +#endif + +#if __cpp_exceptions + #if DMLX_USE_WIL + #define DMLX_THROW_IF_FAILED(_hr) THROW_IF_FAILED(_hr) + #define DMLX_THROW(_hr) THROW_HR(_hr) + #else + #define DMLX_THROW_IF_FAILED(_hr) if (FAILED(_hr)) { throw std::runtime_error(#_hr); } + #define DMLX_THROW(_hr) throw std::runtime_error(#_hr); + #endif +#else + #define DMLX_THROW_IF_FAILED(_hr) if (FAILED(_hr)) { std::abort(); } + #define DMLX_THROW(_hr) { std::abort(); } +#endif + + class Graph; + class Expression; + + using TensorDimensions = SmallVector; + using TensorStrides = SmallVector; + + // The custom properties returned by a TensorPolicy. + struct TensorProperties + { + Optional strides; + uint64_t totalTensorSizeInBytes; + uint32_t guaranteedBaseOffsetAlignment; + }; + + // Provides a way to customize the properties that DMLX automatically sets on tensors. Callers may provide their + // own TensorPolicy implementation to provide custom strides, total tensor sizes, and alignment. TensorPolicy + // objects can be set using Graph::SetTensorPolicy(). + class TensorPolicy + { + public: + // A function type that returns a TensorProperties object given a tensor data type, flags, and sizes. + using Func = std::function< + TensorProperties (DML_TENSOR_DATA_TYPE dataType, DML_TENSOR_FLAGS flags, Span sizes) + >; + + TensorPolicy() = default; + /*implicit*/ TensorPolicy(Func impl) + : m_impl(impl) + {} + + TensorProperties Get( + DML_TENSOR_DATA_TYPE dataType, + DML_TENSOR_FLAGS flags, + Span sizes) const + { + // Empty/uninitialized policy falls back to default. + if (!m_impl) + { + return ComputeDefault(dataType, flags, sizes); + } + + return m_impl(dataType, flags, sizes); + } + + // Returns the default tensor policy, which doesn't produce any changes to tensor layout, has no guaranteed + // alignment, and which uses DMLCalcBufferTensorSize to compute the total tensor size. + static TensorPolicy Default() + { + return TensorPolicy(); + } + + // A tensor policy that returns strides which produce tensors with a layout transposed to dimension order + // (0, 2, ..., n, 1). This is often referred to as "NHWC" or "interleaved channel" layout. This is useful, + // for example, when applied to 2D Convolution to produce outputs in an NHWC layout (as opposed to NCHW, which + // is the DirectML default for 2D Convolution). + // + // Examples of the transposes produced by this policy: + // NCW -> NWC + // NCHW -> NHWC + // NCDHW -> NDHWC + static TensorPolicy InterleavedChannel() + { + return TensorPolicy(&ComputeInterleavedChannel); + } + + private: + static TensorProperties ComputeDefault( + DML_TENSOR_DATA_TYPE dataType, + DML_TENSOR_FLAGS /*flags*/, + Span sizes) + { + uint32_t dimensionCount = static_cast(sizes.size()); + TensorProperties props; + props.strides = NullOpt; // no strides + props.totalTensorSizeInBytes = DMLCalcBufferTensorSize(dataType, dimensionCount, sizes.data(), nullptr); + props.guaranteedBaseOffsetAlignment = 0; + return props; + } + + static TensorProperties ComputeInterleavedChannel( + DML_TENSOR_DATA_TYPE dataType, + DML_TENSOR_FLAGS /*flags*/, + Span sizes) + { + uint32_t dimensionCount = static_cast(sizes.size()); + TensorStrides strides(dimensionCount); + + enum Axes { N, C, /* spatial dimensions ... */ }; + + // N dimension strides + if (dimensionCount >= 1) + { + strides[N] = 1; + for (uint32_t i = 1; i < dimensionCount; ++i) + { + strides[N] *= sizes[i]; + } + } + + // C dimension strides + if (dimensionCount >= 2) + { + strides[C] = 1; + } + + // Spatial dimension strides + if (dimensionCount >= 3) + { + uint32_t stride = sizes[C]; + for (uint32_t i = dimensionCount - 1; i >= 2; --i) + { + strides[i] = stride; + stride *= sizes[i]; + } + } + + TensorProperties props; + props.strides = std::move(strides); + props.totalTensorSizeInBytes = DMLCalcBufferTensorSize(dataType, dimensionCount, sizes.data(), props.strides->data()); + props.guaranteedBaseOffsetAlignment = 0; + return props; + } + + Func m_impl; + }; + + struct TensorDesc + { + public: + using Dimensions = TensorDimensions; + using Strides = TensorStrides; + + DML_TENSOR_DATA_TYPE dataType = DML_TENSOR_DATA_TYPE_UNKNOWN; + DML_TENSOR_FLAGS flags = DML_TENSOR_FLAG_NONE; + Dimensions sizes; + Optional strides; + uint64_t totalTensorSizeInBytes = 0; + uint32_t guaranteedBaseOffsetAlignment = 0; + + TensorDesc() = default; + + TensorDesc(DML_TENSOR_DATA_TYPE dataType, Dimensions sizes, const TensorPolicy& policy = {}) + : TensorDesc(dataType, DML_TENSOR_FLAG_NONE, sizes, policy) + {} + + TensorDesc(DML_TENSOR_DATA_TYPE dataType, DML_TENSOR_FLAGS flags, Dimensions sizes, const TensorPolicy& policy = {}) + { + TensorProperties props = policy.Get(dataType, flags, sizes); + Initialize( + dataType, + flags, + std::move(sizes), + std::move(props.strides), + props.totalTensorSizeInBytes, + props.guaranteedBaseOffsetAlignment); + } + + TensorDesc( + DML_TENSOR_DATA_TYPE dataType, + DML_TENSOR_FLAGS flags, + Dimensions sizes, + Optional strides, + uint64_t totalTensorSizeInBytes, + uint32_t guaranteedBaseOffsetAlignment) + { + Initialize(dataType, flags, std::move(sizes), std::move(strides), totalTensorSizeInBytes, guaranteedBaseOffsetAlignment); + } + + /* implicit */ TensorDesc(const DML_TENSOR_DESC& desc) + : TensorDesc(*static_cast(desc.Desc)) + { + assert(desc.Type == DML_TENSOR_TYPE_BUFFER); + assert(desc.Desc != nullptr); + } + + /* implicit */ TensorDesc(const DML_BUFFER_TENSOR_DESC& desc) + { + this->dataType = desc.DataType; + this->flags = desc.Flags; + this->sizes.assign(desc.Sizes, desc.Sizes + desc.DimensionCount); + if (desc.Strides) + { + this->strides.emplace(); + this->strides->assign(desc.Strides, desc.Strides + desc.DimensionCount); + } + this->totalTensorSizeInBytes = desc.TotalTensorSizeInBytes; + this->guaranteedBaseOffsetAlignment = desc.GuaranteedBaseOffsetAlignment; + } + + // Returns an equivalent DML_TENSOR_DESC or DML_BUFFER_TENSOR_DESC. The returned object contains pointers + // into the TensorDesc, so it is only valid as long as the TensorDesc itself is alive. + template + T* AsPtr() + { + // "sizeof(T) == -1" is always false; this is just to make the static_assert dependent on the template + // parameter and therefore not evaluated until template instantiation + static_assert(sizeof(T) == -1, "Invalid type"); + } + + template <> + DML_BUFFER_TENSOR_DESC* AsPtr() + { + assert(!strides || sizes.size() == strides->size()); + + m_bufferDesc.DataType = this->dataType; + m_bufferDesc.Flags = this->flags; + m_bufferDesc.DimensionCount = static_cast(sizes.size()); + m_bufferDesc.Sizes = this->sizes.data(); + m_bufferDesc.Strides = this->strides ? this->strides->data() : nullptr; + m_bufferDesc.TotalTensorSizeInBytes = this->totalTensorSizeInBytes; + m_bufferDesc.GuaranteedBaseOffsetAlignment = this->guaranteedBaseOffsetAlignment; + return &m_bufferDesc; + } + + template <> + DML_TENSOR_DESC* AsPtr() + { + m_tensorDesc = DML_TENSOR_DESC{ DML_TENSOR_TYPE_BUFFER, AsPtr() }; + return &m_tensorDesc; + } + + private: + DML_BUFFER_TENSOR_DESC m_bufferDesc; + DML_TENSOR_DESC m_tensorDesc; + + void Initialize( + DML_TENSOR_DATA_TYPE tensorDataType, + DML_TENSOR_FLAGS tensorFlags, + Dimensions tensorSizes, + Optional tensorStrides, + uint64_t totalTensorSizeInBytesVal, + uint32_t guaranteedBaseOffsetAlignmentVal) + { + assert(!tensorStrides || tensorStrides->size() == static_cast(tensorSizes.size())); + + this->dataType = tensorDataType; + this->flags = tensorFlags; + this->sizes = std::move(tensorSizes); + this->strides = std::move(tensorStrides); + this->totalTensorSizeInBytes = totalTensorSizeInBytesVal; + this->guaranteedBaseOffsetAlignment = guaranteedBaseOffsetAlignmentVal; + } + }; + + namespace detail + { + class GraphBuilder; + class NodeOutput; + + // A node in the graph which represents a graph input. + struct InputNode + { + uint32_t inputIndex; + }; + + // A node in the graph which represents a DML operator. + struct OperatorNode + { + Microsoft::WRL::ComPtr op; + + // The inputs to this node + std::vector inputs; + }; + + // Used for representing reshapes and type punning + struct ReinterpretNode + { + NodeOutput* input; + }; + + enum class NodeType + { + Invalid, + Input, + Operator, + Reinterpret, + }; + + // Identifies a node in the graph. + struct NodeID + { + NodeType type; + uint32_t index; // The index of this node in the GraphBuilder + }; + + // Represents one of the outputs of a node. + class NodeOutput + { + public: + NodeOutput(GraphBuilder* owner, NodeID node, uint32_t outputIndex, TensorDesc tensorDesc) + : m_owner(owner) + , m_node(node) + , m_outputIndex(outputIndex) + , m_tensorDesc(std::move(tensorDesc)) + {} + + // Retrieves the GraphBuilder that owns this object. + GraphBuilder* GetGraphBuilder() const { return m_owner; } + + NodeID GetNode() const { return m_node; } + uint32_t GetOutputIndex() const { return m_outputIndex; } + const TensorDesc& GetOutputDesc() const { return m_tensorDesc; } + + private: + GraphBuilder* m_owner; + NodeID m_node; + + // An operator can have multiple outputs; this index identifies which one of the operator's outputs this + // NodeOutput represents. + uint32_t m_outputIndex; + + TensorDesc m_tensorDesc; + }; + + struct GraphDesc + { + uint32_t inputCount; + uint32_t outputCount; + std::vector nodes; + std::vector inputEdges; + std::vector outputEdges; + std::vector intermediateEdges; + }; + + class GraphBuilder + { + public: + GraphBuilder(IDMLDevice* device, TensorPolicy tensorPolicy = {}) + : m_device(device) + , m_tensorPolicy(tensorPolicy) + {} + + IDMLDevice* GetDevice() const + { + return m_device.Get(); + } + + void SetTensorPolicy(TensorPolicy policy) { m_tensorPolicy = std::move(policy); } + const TensorPolicy& GetTensorPolicy() const { return m_tensorPolicy; } + TensorPolicy& GetTensorPolicy() { return m_tensorPolicy; } + + // Creates a DML operator node owned by this graph builder and returns a NodeInfo identifier. The + // inputs to this node must be supplied in the correct order matching the DML operator. + NodeID CreateOperatorNode(DML_OPERATOR_TYPE type, const void* desc, Span inputs); + NodeID CreateInputNode(uint32_t inputIndex); + NodeID CreateReinterpretNode(NodeOutput* input); + NodeOutput* CreateNodeOutput(NodeID node, uint32_t outputIndex, TensorDesc tensorDesc); + GraphDesc GetGraphDesc(Span outputs) const; + + private: + Microsoft::WRL::ComPtr m_device; + TensorPolicy m_tensorPolicy; + std::vector m_inputNodes; + std::vector m_operatorNodes; + std::vector m_reinterpretNodes; + std::deque m_nodeOutputs; // deque doesn't invalidate references to elements when it resizes + }; + + } // namespace detail + + class Expression + { + public: + /*implicit*/ Expression(detail::NodeOutput* nodeOutput = nullptr) + : m_nodeOutput(nodeOutput) + {} + + // Returns a struct containing the required properties of the tensor to hold the output of this expression, + // once evaluated. + const TensorDesc& GetOutputDesc() const { return Impl()->GetOutputDesc(); } + + // For internal use only + detail::NodeOutput* Impl() const { return m_nodeOutput; } + + explicit operator bool() const + { + return m_nodeOutput != nullptr; + } + + private: + detail::NodeOutput* m_nodeOutput; // weak; this is owned by the GraphBuilder + }; + + class Graph + { + public: + explicit Graph(IDMLDevice* device, TensorPolicy tensorPolicy = {}) + : m_graphBuilder(make_unique(device, tensorPolicy)) + {} + + // For internal use only + detail::GraphBuilder* Impl() { return m_graphBuilder.get(); } + + // Sets/gets the tensor policy. If not set, defaults to TensorPolicy::Default(). Tensor policies can be used + // to control properties (such as strides) on output tensors produced by this Graph. + void SetTensorPolicy(TensorPolicy policy) { m_graphBuilder->SetTensorPolicy(std::move(policy)); } + const TensorPolicy& GetTensorPolicy() const { return m_graphBuilder->GetTensorPolicy(); } + TensorPolicy& GetTensorPolicy() { return m_graphBuilder->GetTensorPolicy(); } + + Microsoft::WRL::ComPtr Compile( + DML_EXECUTION_FLAGS flags, + Span outputs, + uint32_t inputCount = 0) const + { + detail::GraphDesc graph = m_graphBuilder->GetGraphDesc(outputs); + + // If supplied, the requested number of inputs to the compiled operator can be larger than the actual + // number of input nodes on the graph (e.g. in the case of unused empty inputs), but never smaller. + assert(inputCount == 0 || inputCount >= graph.inputCount); + + std::vector graphNodes(graph.nodes.size()); + for (size_t i = 0; i < graphNodes.size(); ++i) + { + graphNodes[i] = { DML_GRAPH_NODE_TYPE_OPERATOR, &graph.nodes[i] }; + } + + std::vector inputEdges(graph.inputEdges.size()); + for (size_t i = 0; i < inputEdges.size(); ++i) + { + inputEdges[i] = { DML_GRAPH_EDGE_TYPE_INPUT, &graph.inputEdges[i] }; + } + + std::vector outputEdges(graph.outputEdges.size()); + for (size_t i = 0; i < outputEdges.size(); ++i) + { + outputEdges[i] = { DML_GRAPH_EDGE_TYPE_OUTPUT, &graph.outputEdges[i] }; + } + + std::vector intermediateEdges(graph.intermediateEdges.size()); + for (size_t i = 0; i < intermediateEdges.size(); ++i) + { + intermediateEdges[i] = { DML_GRAPH_EDGE_TYPE_INTERMEDIATE, &graph.intermediateEdges[i] }; + } + + DML_GRAPH_DESC graphDesc = {}; + graphDesc.InputCount = inputCount ? inputCount : graph.inputCount; + graphDesc.OutputCount = graph.outputCount; + graphDesc.NodeCount = static_cast(graphNodes.size()); + graphDesc.Nodes = graphNodes.data(); + graphDesc.InputEdgeCount = static_cast(inputEdges.size()); + graphDesc.InputEdges = inputEdges.data(); + graphDesc.OutputEdgeCount = static_cast(outputEdges.size()); + graphDesc.OutputEdges = outputEdges.data(); + graphDesc.IntermediateEdgeCount = static_cast(intermediateEdges.size()); + graphDesc.IntermediateEdges = intermediateEdges.data(); + + Microsoft::WRL::ComPtr device1; + DMLX_THROW_IF_FAILED(m_graphBuilder->GetDevice()->QueryInterface(IID_PPV_ARGS(&device1))); + + Microsoft::WRL::ComPtr compiledGraph; + DMLX_THROW_IF_FAILED(device1->CompileGraph(&graphDesc, flags, IID_PPV_ARGS(&compiledGraph))); + + return compiledGraph; + } + + private: + std::unique_ptr m_graphBuilder; + }; + + // Represents an activation to be fused with an existing operator. The meaning of param1 and param2 depend on the + // activation to be fused. + // + // For HARD_SIGMOID, LINEAR, PARAMETRIC_SOFTPLUS, and SCALED_TANH: param1 = Alpha and param2 = Beta + // For ELU, LEAKY_RELU, THRESHOLDED_RELU, and CELU: param1 = Alpha. param2 is unused. + // For SCALED_ELU, param1 = Alpha and param2 = Gamma. + // For SHRINK, param1 = Bias and param2 = Threshold + // For SOFTPLUS, param1 = Steepness. + // For all other activations, both param1 and param2 are unused. + struct FusedActivation + { + DML_OPERATOR_TYPE activation = DML_OPERATOR_INVALID; + float param1 = 0.0f; + float param2 = 0.0f; + + FusedActivation() = default; + + explicit FusedActivation(DML_OPERATOR_TYPE activation, float param1 = 0.0f, float param2 = 0.0f) + : activation(activation), param1(param1), param2(param2) + {} + + static FusedActivation None() + { + return FusedActivation(); + } + + static FusedActivation Elu(float alpha = 1.0f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_ELU, alpha); + } + + static FusedActivation HardSigmoid(float alpha = 0.2f, float beta = 0.5f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_HARD_SIGMOID, alpha, beta); + } + + static FusedActivation Identity() + { + return FusedActivation(DML_OPERATOR_ACTIVATION_IDENTITY); + } + + static FusedActivation LeakyRelu(float alpha = 0.01f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_LEAKY_RELU, alpha); + } + + static FusedActivation Linear(float alpha, float beta) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_LINEAR, alpha, beta); + } + + static FusedActivation ParametricSoftplus(float alpha, float beta) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_PARAMETRIC_SOFTPLUS, alpha, beta); + } + + static FusedActivation Relu() + { + return FusedActivation(DML_OPERATOR_ACTIVATION_RELU); + } + + static FusedActivation ScaledElu(float alpha = 1.67326319217681884765625f, float gamma = 1.05070102214813232421875f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_SCALED_ELU, alpha, gamma); + } + + static FusedActivation ScaledTanh(float alpha = 1.0f, float beta = 0.5f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_SCALED_TANH, alpha, beta); + } + + static FusedActivation Sigmoid() + { + return FusedActivation(DML_OPERATOR_ACTIVATION_SIGMOID); + } + + static FusedActivation Softplus(float steepness = 1.0f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_SOFTPLUS, steepness); + } + + static FusedActivation Softsign() + { + return FusedActivation(DML_OPERATOR_ACTIVATION_SOFTSIGN); + } + + static FusedActivation Tanh() + { + return FusedActivation(DML_OPERATOR_ACTIVATION_TANH); + } + + static FusedActivation ThresholdedRelu(float alpha = 1.0f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_THRESHOLDED_RELU, alpha); + } + + static FusedActivation Shrink(float bias = 0.0f, float threshold = 0.5f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_SHRINK, bias, threshold); + } + + static FusedActivation Celu(float alpha = 1.0f) + { + return FusedActivation(DML_OPERATOR_ACTIVATION_CELU, alpha); + } + }; + + // Implementation detail helper for determining if a list of expressions share the same GraphBuilder. + namespace detail + { + inline bool HasSameOwner(Span exprs) + { + if (exprs.size() == 0) + { + return true; + } + + detail::GraphBuilder* owner = exprs.begin()->Impl()->GetGraphBuilder(); + for (Expression expr : exprs) + { + if (expr.Impl()->GetGraphBuilder() != owner) + { + return false; + } + } + + return true; + } + + inline bool HasSameOwner(std::initializer_list exprs) + { + Span span(exprs.begin(), exprs.size()); + return HasSameOwner(span); + } + + inline bool HasSameDataType(Span exprs) + { + if (exprs.size() == 0) + { + return true; + } + + DML_TENSOR_DATA_TYPE dataType = exprs.begin()->Impl()->GetOutputDesc().dataType; + for (Expression expr : exprs) + { + if (expr.Impl()->GetOutputDesc().dataType != dataType) + { + return false; + } + } + + return true; + } + + inline bool HasSameDataType(std::initializer_list exprs) + { + Span span(exprs.begin(), exprs.size()); + return HasSameDataType(span); + } + } // namespace detail + + // Expression implementation helpers + namespace detail + { + template + Expression ElementWiseUnary(Expression input, const Optional& scaleBias) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); // Same as input + + TDesc desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.ScaleBias = scaleBias ? &scaleBias.value() : nullptr; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(OperatorType, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + template + Expression ElementWiseUnary(Expression input, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UNKNOWN) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + if (outputDataType == DML_TENSOR_DATA_TYPE_UNKNOWN) + { + outputDataType = inputTensor.dataType; + } + TensorDesc outputTensor(outputDataType, inputTensor.sizes, builder->GetTensorPolicy()); + + TDesc desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(OperatorType, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + template + Expression ElementWiseBinary(Expression a, Expression b) + { + assert(detail::HasSameOwner({ a, b })); + detail::GraphBuilder* builder = a.Impl()->GetGraphBuilder(); + + TensorDesc aTensor = a.Impl()->GetOutputDesc(); + TensorDesc bTensor = b.Impl()->GetOutputDesc(); + TensorDesc outputTensor(aTensor.dataType, aTensor.sizes, builder->GetTensorPolicy()); // Same as input + + TDesc desc = {}; + desc.ATensor = aTensor.AsPtr(); + desc.BTensor = bTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { a.Impl(), b.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(OperatorType, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + template + Expression ElementWiseComparison(Expression a, Expression b, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + assert(detail::HasSameOwner({ a, b })); + detail::GraphBuilder* builder = a.Impl()->GetGraphBuilder(); + + TensorDesc aTensor = a.Impl()->GetOutputDesc(); + TensorDesc bTensor = b.Impl()->GetOutputDesc(); + TensorDesc outputTensor(outputDataType, aTensor.sizes, builder->GetTensorPolicy()); + + TDesc desc = {}; + desc.ATensor = aTensor.AsPtr(); + desc.BTensor = bTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { a.Impl(), b.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(OperatorType, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + // Used to reserve some space on the stack for setting up fused activation operator descs. + struct FusedActivationStorage + { + DML_OPERATOR_DESC opDesc; + + // All fuseable activation descs have a common layout: two tensor desc pointers and up to 2 optional + // float parameters, so just use LINEAR as an archetype + DML_ACTIVATION_LINEAR_OPERATOR_DESC activationDesc; + }; + + // Returns the correct value for filling out fused activation fields in the DML API, e.g. + // DML_CONVOLUTION_OPERATOR_DESC::FusedActivation. The descs themselves are stored in the `storage` outptr. + inline const DML_OPERATOR_DESC* GetFusedActivationPtr( + FusedActivation fusedActivation, + _Out_ FusedActivationStorage* storage) + { + if (fusedActivation.activation == DML_OPERATOR_INVALID) + { + // No fused activation + return nullptr; + } + + storage->activationDesc.InputTensor = nullptr; + storage->activationDesc.OutputTensor = nullptr; + storage->activationDesc.Alpha = fusedActivation.param1; + storage->activationDesc.Beta = fusedActivation.param2; + + storage->opDesc.Type = fusedActivation.activation; + storage->opDesc.Desc = &storage->activationDesc; + + return &storage->opDesc; + } + + } // namespace detail + + inline Expression InputTensor(Graph& graph, uint32_t inputIndex, TensorDesc desc) + { + detail::GraphBuilder* builder = graph.Impl(); + + detail::NodeID node = builder->CreateInputNode(inputIndex); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(desc)); + return output; + } + + inline Expression Identity(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Abs(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression ACos(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Add(Expression a, Expression b) + { + assert(detail::HasSameOwner({ a, b })); + detail::GraphBuilder* builder = a.Impl()->GetGraphBuilder(); + + TensorDesc aTensor = a.Impl()->GetOutputDesc(); + TensorDesc bTensor = b.Impl()->GetOutputDesc(); + TensorDesc outputTensor(aTensor.dataType, aTensor.sizes, builder->GetTensorPolicy()); // Same as input + + DML_ELEMENT_WISE_ADD_OPERATOR_DESC desc = {}; + desc.ATensor = aTensor.AsPtr(); + desc.BTensor = bTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { a.Impl(), b.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_ADD, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Add(Expression a, Expression b, FusedActivation fusedActivation) + { + assert(detail::HasSameOwner({ a, b })); + detail::GraphBuilder* builder = a.Impl()->GetGraphBuilder(); + + TensorDesc aTensor = a.Impl()->GetOutputDesc(); + TensorDesc bTensor = b.Impl()->GetOutputDesc(); + TensorDesc outputTensor(aTensor.dataType, aTensor.sizes, builder->GetTensorPolicy()); // Same as input + detail::FusedActivationStorage storage; + + DML_ELEMENT_WISE_ADD1_OPERATOR_DESC desc = {}; + desc.ATensor = aTensor.AsPtr(); + desc.BTensor = bTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.FusedActivation = detail::GetFusedActivationPtr(fusedActivation, &storage); + + detail::NodeOutput* const inputs[] = { a.Impl(), b.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_ADD1, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression ASin(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression ATan(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + +#if DML_TARGET_VERSION >= 0x3100 + + inline Expression ATanYX(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + +#endif // DML_TARGET_VERSION >= 0x3100 + + inline Expression Ceil(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Clip(Expression input, float min, float max, const Optional& scaleBias = NullOpt) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); // Same as input + + DML_ELEMENT_WISE_CLIP_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.ScaleBias = scaleBias ? &scaleBias.value() : nullptr; + desc.Min = min; + desc.Max = max; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_CLIP, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + +#if DML_TARGET_VERSION >= 0x3100 + + inline Expression ClipGrad(Expression input, Expression inputGradient, float min, float max) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc inputGradientTensor = inputGradient.Impl()->GetOutputDesc(); + TensorDesc outputGradientTensor(inputGradientTensor.dataType, inputGradientTensor.sizes, builder->GetTensorPolicy()); + + DML_ELEMENT_WISE_CLIP_GRAD_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.InputGradientTensor = inputGradientTensor.AsPtr(); + desc.OutputGradientTensor = outputGradientTensor.AsPtr(); + desc.Min = min; + desc.Max = max; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_CLIP_GRAD, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputGradientTensor)); + + return output; + } + +#endif // DML_TARGET_VERSION >= 0x3100 + + inline Expression Cos(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Divide(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Exp(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Floor(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Log(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression LogicalAnd(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Equals(Expression a, Expression b, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + return detail::ElementWiseComparison< + DML_OPERATOR_ELEMENT_WISE_LOGICAL_EQUALS, + DML_ELEMENT_WISE_LOGICAL_EQUALS_OPERATOR_DESC>(a, b, outputDataType); + } + + inline Expression GreaterThan(Expression a, Expression b, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + return detail::ElementWiseComparison< + DML_OPERATOR_ELEMENT_WISE_LOGICAL_GREATER_THAN, + DML_ELEMENT_WISE_LOGICAL_GREATER_THAN_OPERATOR_DESC>(a, b, outputDataType); + } + + inline Expression GreaterThanOrEqual(Expression a, Expression b, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + return detail::ElementWiseComparison< + DML_OPERATOR_ELEMENT_WISE_LOGICAL_GREATER_THAN_OR_EQUAL, + DML_ELEMENT_WISE_LOGICAL_GREATER_THAN_OR_EQUAL_OPERATOR_DESC>(a, b, outputDataType); + } + + inline Expression LessThan(Expression a, Expression b, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + return detail::ElementWiseComparison< + DML_OPERATOR_ELEMENT_WISE_LOGICAL_LESS_THAN, + DML_ELEMENT_WISE_LOGICAL_LESS_THAN_OPERATOR_DESC>(a, b, outputDataType); + } + + inline Expression LessThanOrEqual(Expression a, Expression b, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + return detail::ElementWiseComparison< + DML_OPERATOR_ELEMENT_WISE_LOGICAL_LESS_THAN_OR_EQUAL, + DML_ELEMENT_WISE_LOGICAL_LESS_THAN_OR_EQUAL_OPERATOR_DESC>(a, b, outputDataType); + } + + inline Expression LogicalNot(Expression input) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); // Same as input + + DML_ELEMENT_WISE_LOGICAL_NOT_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_LOGICAL_NOT, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression LogicalOr(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression LogicalXor(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Max(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Mean(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Min(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Multiply(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Pow(Expression input, Expression exponent, const Optional& scaleBias = NullOpt) + { + assert(detail::HasSameOwner({ input, exponent })); + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc exponentTensor = exponent.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); // Same as input + + DML_ELEMENT_WISE_POW_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.ExponentTensor = exponentTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.ScaleBias = scaleBias ? &scaleBias.value() : nullptr; + + detail::NodeOutput* const inputs[] = { input.Impl(), exponent.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_POW, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Pow(Expression input, float exponent, const Optional& scaleBias = NullOpt) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); // Same as input + + DML_ELEMENT_WISE_CONSTANT_POW_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.ScaleBias = scaleBias ? &scaleBias.value() : nullptr; + desc.Exponent = exponent; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_CONSTANT_POW, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Recip(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Sin(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Sqrt(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + +#if DML_TARGET_VERSION >= 0x3100 + + inline Expression DifferenceSquare(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + +#endif // DML_TARGET_VERSION >= 0x3100 + + inline Expression Subtract(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression Tan(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Threshold(Expression input, float min, const Optional& scaleBias = NullOpt) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); // Same as input + + DML_ELEMENT_WISE_THRESHOLD_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.ScaleBias = scaleBias ? &scaleBias.value() : nullptr; + desc.Min = min; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_THRESHOLD, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression QuantizeLinear(Expression input, Expression scale, Expression zeroPoint, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + assert(detail::HasSameOwner({ input, scale, zeroPoint })); + + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc scaleTensor = scale.Impl()->GetOutputDesc(); + TensorDesc zeroPointTensor = zeroPoint.Impl()->GetOutputDesc(); + TensorDesc outputTensor(outputDataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_ELEMENT_WISE_QUANTIZE_LINEAR_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.ScaleTensor = scaleTensor.AsPtr(); + desc.ZeroPointTensor = zeroPointTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { input.Impl(), scale.Impl(), zeroPoint.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_QUANTIZE_LINEAR, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression DequantizeLinear(Expression input, Expression scale, Expression zeroPoint) + { + assert(detail::HasSameOwner({ input, scale, zeroPoint })); + + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc scaleTensor = scale.Impl()->GetOutputDesc(); + TensorDesc zeroPointTensor = zeroPoint.Impl()->GetOutputDesc(); + TensorDesc outputTensor(DML_TENSOR_DATA_TYPE_FLOAT32, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_ELEMENT_WISE_DEQUANTIZE_LINEAR_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.ScaleTensor = scaleTensor.AsPtr(); + desc.ZeroPointTensor = zeroPointTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { input.Impl(), scale.Impl(), zeroPoint.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_DEQUANTIZE_LINEAR, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Sign(Expression a) + { + return detail::ElementWiseUnary(a); + } + + inline Expression IsNaN(Expression input, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + return detail::ElementWiseUnary(input, outputDataType); + } + + inline Expression Erf(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Sinh(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Cosh(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression Tanh(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression ASinh(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression ACosh(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression ATanh(Expression input, const Optional& scaleBias = NullOpt) + { + return detail::ElementWiseUnary(input, scaleBias); + } + + inline Expression If(Expression condition, Expression a, Expression b) + { + assert(detail::HasSameOwner({ condition, a, b })); + assert(detail::HasSameDataType({ a, b })); + + detail::GraphBuilder* builder = condition.Impl()->GetGraphBuilder(); + + TensorDesc conditionTensor = condition.Impl()->GetOutputDesc(); + assert(conditionTensor.dataType == DML_TENSOR_DATA_TYPE_UINT8); + + TensorDesc aTensor = a.Impl()->GetOutputDesc(); + TensorDesc bTensor = b.Impl()->GetOutputDesc(); + TensorDesc outputTensor(aTensor.dataType, aTensor.sizes, builder->GetTensorPolicy()); + + DML_ELEMENT_WISE_IF_OPERATOR_DESC desc = {}; + desc.ConditionTensor = conditionTensor.AsPtr(); + desc.ATensor = aTensor.AsPtr(); + desc.BTensor = bTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { condition.Impl(), a.Impl(), b.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_IF, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression BitShiftLeft(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression BitShiftRight(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression BitAnd(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression BitOr(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression BitXor(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression BitNot(Expression a) + { + return detail::ElementWiseUnary(a); + } + + inline Expression BitCount(Expression a, DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + return detail::ElementWiseUnary(a, outputDataType); + } + + inline Expression Round(Expression input, DML_ROUNDING_MODE roundingMode = DML_ROUNDING_MODE_HALVES_TO_NEAREST_EVEN) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); // Same as input + + DML_ELEMENT_WISE_ROUND_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.RoundingMode = roundingMode; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_ROUND, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression IsInfinity( + Expression input, + DML_IS_INFINITY_MODE infinityMode = DML_IS_INFINITY_MODE_EITHER, + DML_TENSOR_DATA_TYPE outputDataType = DML_TENSOR_DATA_TYPE_UINT8) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(outputDataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_ELEMENT_WISE_IS_INFINITY_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.InfinityMode = infinityMode; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ELEMENT_WISE_IS_INFINITY, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression ModulusTruncate(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + + inline Expression ModulusFloor(Expression a, Expression b) + { + return detail::ElementWiseBinary(a, b); + } + +#pragma region detail +#define DMLX_ACTIVATION_IMPL(_name) \ + do { \ + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); \ + \ + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); \ + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); \ + \ + DML_##_name##_OPERATOR_DESC desc = {}; \ + desc.InputTensor = inputTensor.AsPtr(); \ + desc.OutputTensor = outputTensor.AsPtr(); \ + \ + detail::NodeOutput* const inputs[] = { input.Impl() }; \ + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_##_name, &desc, inputs); \ + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); \ + \ + return output; \ + } while(0) + +#define DMLX_ACTIVATION_IMPL_1(_name, _param1Name, _param1) \ + do { \ + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); \ + \ + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); \ + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); \ + \ + DML_##_name##_OPERATOR_DESC desc = {}; \ + desc.InputTensor = inputTensor.AsPtr(); \ + desc.OutputTensor = outputTensor.AsPtr(); \ + desc._param1Name = _param1; \ + \ + detail::NodeOutput* const inputs[] = { input.Impl() }; \ + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_##_name, &desc, inputs); \ + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); \ + \ + return output; \ + } while(0) + +#define DMLX_ACTIVATION_IMPL_2(_name, _param1Name, _param1, _param2Name, _param2) \ + do { \ + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); \ + \ + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); \ + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); \ + \ + DML_##_name##_OPERATOR_DESC desc = {}; \ + desc.InputTensor = inputTensor.AsPtr(); \ + desc.OutputTensor = outputTensor.AsPtr(); \ + desc._param1Name = _param1; \ + desc._param2Name = _param2; \ + \ + detail::NodeOutput* const inputs[] = { input.Impl() }; \ + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_##_name, &desc, inputs); \ + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); \ + \ + return output; \ + } while(0) +#pragma endregion + + inline Expression ActivationElu(Expression input, float alpha = 1.0f) + { + DMLX_ACTIVATION_IMPL_1(ACTIVATION_ELU, Alpha, alpha); + } + + inline Expression ActivationHardmax(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_HARDMAX); + } + + inline Expression ActivationHardSigmoid(Expression input, float alpha = 0.2f, float beta = 0.5f) + { + DMLX_ACTIVATION_IMPL_2(ACTIVATION_HARD_SIGMOID, Alpha, alpha, Beta, beta); + } + + inline Expression ActivationIdentity(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_IDENTITY); + } + + inline Expression ActivationLeakyRelu(Expression input, float alpha = 0.01f) + { + DMLX_ACTIVATION_IMPL_1(ACTIVATION_LEAKY_RELU, Alpha, alpha); + } + + inline Expression ActivationLinear(Expression input, float alpha, float beta) + { + DMLX_ACTIVATION_IMPL_2(ACTIVATION_LINEAR, Alpha, alpha, Beta, beta); + } + + inline Expression ActivationLogSoftmax(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_LOG_SOFTMAX); + } + + inline Expression ActivationParameterizedRelu(Expression input, Expression slope) + { + assert(detail::HasSameOwner({ input, slope })); + + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc slopeTensor = slope.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_ACTIVATION_PARAMETERIZED_RELU_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.SlopeTensor = slopeTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { input.Impl(), slope.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ACTIVATION_PARAMETERIZED_RELU, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression ActivationParametricSoftplus(Expression input, float alpha, float beta) + { + DMLX_ACTIVATION_IMPL_2(ACTIVATION_PARAMETRIC_SOFTPLUS, Alpha, alpha, Beta, beta); + } + + inline Expression ActivationRelu(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_RELU); + } + + inline Expression ActivationScaledElu(Expression input, float alpha = 1.67326319217681884765625f, float gamma = 1.05070102214813232421875f) + { + DMLX_ACTIVATION_IMPL_2(ACTIVATION_SCALED_ELU, Alpha, alpha, Gamma, gamma); + } + + inline Expression ActivationScaledTanh(Expression input, float alpha = 1.0f, float beta = 0.5f) + { + DMLX_ACTIVATION_IMPL_2(ACTIVATION_SCALED_TANH, Alpha, alpha, Beta, beta); + } + + inline Expression ActivationSigmoid(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_SIGMOID); + } + + inline Expression ActivationSoftmax(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_SOFTMAX); + } + + inline Expression ActivationSoftplus(Expression input, float steepness = 1.0f) + { + DMLX_ACTIVATION_IMPL_1(ACTIVATION_SOFTPLUS, Steepness, steepness); + } + + inline Expression ActivationSoftsign(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_SOFTSIGN); + } + + inline Expression ActivationTanh(Expression input) + { + DMLX_ACTIVATION_IMPL(ACTIVATION_TANH); + } + + inline Expression ActivationThresholdedRelu(Expression input, float alpha = 1.0f) + { + DMLX_ACTIVATION_IMPL_1(ACTIVATION_THRESHOLDED_RELU, Alpha, alpha); + } + + inline Expression ActivationShrink(Expression input, float bias = 0.0f, float threshold = 0.5f) + { + DMLX_ACTIVATION_IMPL_2(ACTIVATION_SHRINK, Bias, bias, Threshold, threshold); + } + + inline Expression ActivationCelu(Expression input, float alpha = 1.0f) + { + DMLX_ACTIVATION_IMPL_1(ACTIVATION_CELU, Alpha, alpha); + } + +#undef DMLX_ACTIVATION_IMPL +#undef DMLX_ACTIVATION_IMPL_1 +#undef DMLX_ACTIVATION_IMPL_2 + + // --------------------------------------------------------------------------------------------------------------- + + // If not specified, parameters are defaulted to the following values: + // Mode = DML_CONVOLUTION_MODE_CROSS_CORRELATION + // Direction = DML_CONVOLUTION_DIRECTION_FORWARD + // Strides = { 1, 1 } for 2D convolution, { 1, 1, 1 } for 3D convolution + // Dilations = { 1, 1 } for 2D convolution, { 1, 1, 1 } for 3D convolution + // StartPadding = { 0, 0 } for 2D convolution, { 0, 0, 0 } for 3D convolution + // EndPadding = { 0, 0 } for 2D convolution, { 0, 0, 0 } for 3D convolution + // OutputPadding = { 0, 0 } for 2D convolution, { 0, 0, 0 } for 3D convolution + // GroupCount = 1 + // FusedActivation = nullptr + // OutputSizes = computed from other parameters + inline Expression Convolution( + Expression input, + Expression filter, + Optional bias = NullOpt, + DML_CONVOLUTION_MODE mode = DML_CONVOLUTION_MODE_CROSS_CORRELATION, + DML_CONVOLUTION_DIRECTION direction = DML_CONVOLUTION_DIRECTION_FORWARD, + Span strides = {}, + Span dilations = {}, + Span startPadding = {}, + Span endPadding = {}, + Span outputPadding = {}, + uint32_t groupCount = 1, + FusedActivation fusedActivation = FusedActivation::None(), + TensorDimensions outputSizes = {}) + { + assert(detail::HasSameOwner({ input, filter })); + assert(!bias || detail::HasSameOwner({ input, *bias })); + + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc filterTensor = filter.Impl()->GetOutputDesc(); + TensorDesc biasTensor; + if (bias) + { + biasTensor = bias->Impl()->GetOutputDesc(); + } + + uint32_t dimensionCount = static_cast(inputTensor.sizes.size()); + + assert(dimensionCount == 4 || dimensionCount == 5); + uint32_t spatialDimensionCount = dimensionCount - 2; + + // If the spatial dimension count is 2, we'll just use the first two elements by setting + // DimensionCount = 2 in the desc + const uint32_t defaultStridesAndDilations[3] = { 1, 1, 1 }; + const uint32_t defaultPadding[3] = { 0, 0, 0 }; + + assert(strides.empty() || strides.size() == spatialDimensionCount); + assert(dilations.empty() || dilations.size() == spatialDimensionCount); + assert(startPadding.empty() || startPadding.size() == spatialDimensionCount); + assert(endPadding.empty() || endPadding.size() == spatialDimensionCount); + assert(outputPadding.empty() || outputPadding.size() == spatialDimensionCount); + assert(outputSizes.empty() || outputSizes.size() == inputTensor.sizes.size()); + + strides = strides.empty() ? Span{ defaultStridesAndDilations } : strides; + dilations = dilations.empty() ? Span{ defaultStridesAndDilations } : dilations; + startPadding = startPadding.empty() ? Span{ defaultPadding } : startPadding; + endPadding = endPadding.empty() ? Span{ defaultPadding } : endPadding; + outputPadding = outputPadding.empty() ? Span{ defaultPadding } : outputPadding; + + // Compute the output shapes + + if (outputSizes.empty()) + { + if (direction == DML_CONVOLUTION_DIRECTION_FORWARD) + { + outputSizes.push_back(inputTensor.sizes[0]); // output[N] = input[N] + outputSizes.push_back(filterTensor.sizes[0]); // output[C] = filter[N] + + for (uint32_t dim = 0; dim < spatialDimensionCount; ++dim) + { + uint32_t inputSize = inputTensor.sizes[dim + 2]; + uint32_t paddedSize = inputSize + startPadding[dim] + endPadding[dim]; + + uint32_t windowSize = filterTensor.sizes[dim + 2]; + uint32_t kernelSize = 1 + (windowSize - 1) * dilations[dim]; + + assert(kernelSize <= paddedSize); + assert(strides[dim] != 0); + + outputSizes.push_back(1 + (paddedSize - kernelSize) / strides[dim]); + } + } + else if (direction == DML_CONVOLUTION_DIRECTION_BACKWARD) + { + // TODO: implement me + assert(false); + } + else + { + assert(false); + DMLX_THROW(E_UNEXPECTED); + } + } + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + detail::FusedActivationStorage storage; + + DML_CONVOLUTION_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.FilterTensor = filterTensor.AsPtr(); + desc.BiasTensor = bias ? biasTensor.AsPtr() : nullptr; + desc.OutputTensor = outputTensor.AsPtr(); + desc.Mode = mode; + desc.Direction = direction; + desc.DimensionCount = spatialDimensionCount; + desc.Strides = strides.data(); + desc.Dilations = dilations.data(); + desc.StartPadding = startPadding.data(); + desc.EndPadding = endPadding.data(); + desc.OutputPadding = outputPadding.data(); + desc.GroupCount = groupCount; + desc.FusedActivation = detail::GetFusedActivationPtr(fusedActivation, &storage); + + detail::NodeOutput* const inputs[] = { input.Impl(), filter.Impl(), bias ? bias->Impl() : nullptr }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_CONVOLUTION, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + // Helper for setting parameters for the Convolution operator. Sample usage: + // + // auto conv = dml::ConvolutionBuilder(...) + // .StartPadding(...) + // .EndPadding(...) + // .Strides(...) + // .Build(); + // + // Parameters left unspecified will be defaulted with the same values as dml::Convolution(). + class ConvolutionBuilder + { + public: + ConvolutionBuilder(Expression input, Expression filter, Optional bias = NullOpt) + : m_input(input), m_filter(filter), m_bias(bias) + {} + + ConvolutionBuilder& Mode(DML_CONVOLUTION_MODE mode) { m_mode = mode; return *this; } + ConvolutionBuilder& Direction(DML_CONVOLUTION_DIRECTION direction) { m_direction = direction; return *this; } + ConvolutionBuilder& Strides(Span strides) { m_strides.assign(strides.begin(), strides.end()); return *this; } + ConvolutionBuilder& Dilations(Span dilations) { m_dilations.assign(dilations.begin(), dilations.end()); return *this; } + ConvolutionBuilder& StartPadding(Span startPadding) { m_startPadding.assign(startPadding.begin(), startPadding.end()); return *this; } + ConvolutionBuilder& EndPadding(Span endPadding) { m_endPadding.assign(endPadding.begin(), endPadding.end()); return *this; } + ConvolutionBuilder& OutputPadding(Span outputPadding) { m_outputPadding.assign(outputPadding.begin(), outputPadding.end()); return *this; } + ConvolutionBuilder& GroupCount(uint32_t groupCount) { m_groupCount = groupCount; return *this; } + ConvolutionBuilder& FusedActivation(FusedActivation fusedActivation) { m_fusedActivation = fusedActivation; return *this; } + ConvolutionBuilder& OutputSizes(TensorDimensions outputSizes) { m_outputSizes = std::move(outputSizes); return *this; } + + Expression Build() const + { + return Convolution( + m_input, + m_filter, + m_bias, + m_mode, + m_direction, + m_strides, + m_dilations, + m_startPadding, + m_endPadding, + m_outputPadding, + m_groupCount, + m_fusedActivation, + m_outputSizes); + } + + private: + Expression m_input; + Expression m_filter; + Optional m_bias; + DML_CONVOLUTION_MODE m_mode = DML_CONVOLUTION_MODE_CROSS_CORRELATION; + DML_CONVOLUTION_DIRECTION m_direction = DML_CONVOLUTION_DIRECTION_FORWARD; + SmallVector m_strides = {}; + SmallVector m_dilations = {}; + SmallVector m_startPadding = {}; + SmallVector m_endPadding = {}; + SmallVector m_outputPadding = {}; + uint32_t m_groupCount = 1; + dml::FusedActivation m_fusedActivation; + TensorDimensions m_outputSizes = {}; + }; + + // --------------------------------------------------------------------------------------------------------------- + + inline Expression Gemm( + Expression a, + Expression b, + Optional c = NullOpt, + DML_MATRIX_TRANSFORM transA = DML_MATRIX_TRANSFORM_NONE, + DML_MATRIX_TRANSFORM transB = DML_MATRIX_TRANSFORM_NONE, + float alpha = 1.0f, + float beta = 1.0f, + FusedActivation fusedActivation = FusedActivation::None()) + { + assert(detail::HasSameOwner({ a, b })); + assert(!c || detail::HasSameOwner({ a, *c })); + + detail::GraphBuilder* builder = a.Impl()->GetGraphBuilder(); + + TensorDesc aTensor = a.Impl()->GetOutputDesc(); + TensorDesc bTensor = b.Impl()->GetOutputDesc(); + TensorDesc cTensor; + if (c) + { + cTensor = c->Impl()->GetOutputDesc(); + } + + TensorDimensions outputSizes; + outputSizes.push_back(aTensor.sizes[0]); // output[N] = input[N] + outputSizes.push_back(aTensor.sizes[1]); // output[C] = input[C] + outputSizes.push_back(transA == DML_MATRIX_TRANSFORM_NONE ? aTensor.sizes[2] : aTensor.sizes[3]); + outputSizes.push_back(transB == DML_MATRIX_TRANSFORM_NONE ? bTensor.sizes[3] : bTensor.sizes[2]); + + TensorDesc outputTensor(aTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + detail::FusedActivationStorage storage; + + DML_GEMM_OPERATOR_DESC desc = {}; + desc.ATensor = aTensor.AsPtr(); + desc.BTensor = bTensor.AsPtr(); + desc.CTensor = c ? cTensor.AsPtr() : nullptr; + desc.OutputTensor = outputTensor.AsPtr(); + desc.TransA = transA; + desc.TransB = transB; + desc.Alpha = alpha; + desc.Beta = beta; + desc.FusedActivation = detail::GetFusedActivationPtr(fusedActivation, &storage); + + detail::NodeOutput* const inputs[] = { a.Impl(), b.Impl(), c ? c->Impl() : nullptr }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_GEMM, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + // Helper for setting parameters for the Gemm operator. Parameters left unspecified will be defaulted with the + // same values as dml::Gemm(). + class GemmBuilder + { + public: + GemmBuilder(Expression a, Expression b, Optional c = NullOpt) + : m_a(a), m_b(b), m_c(c) + {} + + GemmBuilder& TransA(DML_MATRIX_TRANSFORM transA) { m_transA = transA; return *this; } + GemmBuilder& TransB(DML_MATRIX_TRANSFORM transB) { m_transB = transB; return *this; } + GemmBuilder& Alpha(float alpha) { m_alpha = alpha; return *this; } + GemmBuilder& Beta(float beta) { m_beta = beta; return *this; } + GemmBuilder& FusedActivation(FusedActivation fusedActivation) { m_fusedActivation = fusedActivation; return *this; } + + Expression Build() const + { + return Gemm(m_a, m_b, m_c, m_transA, m_transB, m_alpha, m_beta, m_fusedActivation); + } + + private: + Expression m_a; + Expression m_b; + Optional m_c; + DML_MATRIX_TRANSFORM m_transA = DML_MATRIX_TRANSFORM_NONE; + DML_MATRIX_TRANSFORM m_transB = DML_MATRIX_TRANSFORM_NONE; + float m_alpha = 1.0f; + float m_beta = 1.0f; + dml::FusedActivation m_fusedActivation; + }; + + // --------------------------------------------------------------------------------------------------------------- + + // If `axes` is not specified, by default this reduces the entire tensor to single element. + inline Expression Reduce(Expression input, DML_REDUCE_FUNCTION function, Span axes = {}) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + uint32_t dimensionCount = static_cast(inputTensor.sizes.size()); + + SmallVector defaultAxes; + if (axes.empty()) + { + for (uint32_t i = 0; i < dimensionCount; ++i) + { + defaultAxes.push_back(i); + } + axes = defaultAxes; + } + + // Compute the output tensor dimensions + TensorDimensions outputSizes; + for (uint32_t i = 0; i < dimensionCount; ++i) + { + // If the dimension is to be reduced, this dimension in the output tensor has a size of 1, otherwise + // it matches the input tensor. + const bool dimensionIsReduced = std::find(axes.begin(), axes.end(), i) != axes.end(); + if (dimensionIsReduced) + { + outputSizes.push_back(1); + } + else + { + outputSizes.push_back(inputTensor.sizes[i]); + } + } + + // ARGMIN and ARGMAX reduction produce a UINT32 output; all other reductions produce an output with the same + // type as the input. + DML_TENSOR_DATA_TYPE outputDataType; + if (function == DML_REDUCE_FUNCTION_ARGMIN || function == DML_REDUCE_FUNCTION_ARGMAX) + { + outputDataType = DML_TENSOR_DATA_TYPE_UINT32; + } + else + { + outputDataType = inputTensor.dataType; + } + + TensorDesc outputTensor(outputDataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_REDUCE_OPERATOR_DESC desc = {}; + desc.Function = function; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.AxisCount = static_cast(axes.size()); + desc.Axes = axes.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_REDUCE, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression AveragePooling( + Expression input, + Span strides, + Span windowSizes, + Span startPadding, + Span endPadding, + bool includePadding, + TensorDimensions outputSizes = {}) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + assert(strides.size() == windowSizes.size()); + assert(strides.size() == startPadding.size()); + assert(strides.size() == endPadding.size()); + + // Calculate output size, if not explicitly provided + if (outputSizes.empty()) + { + outputSizes.push_back(inputTensor.sizes[0]); // N + outputSizes.push_back(inputTensor.sizes[1]); // C + for (size_t i = 0; i < windowSizes.size(); ++i) + { + uint32_t paddedInputSize = inputTensor.sizes[2 + i] + startPadding[i] + endPadding[i]; + uint32_t outputSize = (paddedInputSize - windowSizes[i]) / strides[i] + 1; + outputSizes.push_back(outputSize); + } + } + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_AVERAGE_POOLING_OPERATOR_DESC averagePoolDesc = {}; + averagePoolDesc.InputTensor = inputTensor.AsPtr(); + averagePoolDesc.OutputTensor = outputTensor.AsPtr(); + averagePoolDesc.DimensionCount = static_cast(windowSizes.size()); + averagePoolDesc.Strides = strides.data(); + averagePoolDesc.WindowSize = windowSizes.data(); + averagePoolDesc.StartPadding = startPadding.data(); + averagePoolDesc.EndPadding = endPadding.data(); + averagePoolDesc.IncludePadding = includePadding; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_AVERAGE_POOLING, &averagePoolDesc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + // + // TODO: LpPooling + // + + // --------------------------------------------------------------------------------------------------------------- + + struct MaxPoolingOutputs + { + Expression values; + Expression indices; // Only valid if outputIndices = true is supplied to MaxPooling() + }; + + // If not specified, parameters are defaulted to the following values: + // Strides = 1 for each spatial dimension + // StartPadding = 0 for each spatial dimension + // EndPadding = 0 for each spatial dimension + // Dilations = 1 for each spatial dimension + // OutputIndices = false + inline MaxPoolingOutputs MaxPooling( + Expression input, + Span windowSize, + Span strides = {}, + Span startPadding = {}, + Span endPadding = {}, + Span dilations = {}, + bool outputIndices = false, + TensorDimensions outputSizes = {}) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + // If the spatial dimension count is 2, we'll just use the first two elements by setting + // DimensionCount = 2 in the desc + const uint32_t defaultStridesAndDilations[3] = { 1, 1, 1 }; + const uint32_t defaultPadding[3] = { 0, 0, 0 }; + + assert(windowSize.size() == 2 || windowSize.size() == 3); + assert(strides.empty() || strides.size() == windowSize.size()); + assert(dilations.empty() || dilations.size() == windowSize.size()); + assert(startPadding.empty() || startPadding.size() == windowSize.size()); + assert(endPadding.empty() || endPadding.size() == windowSize.size()); + + strides = strides.empty() ? Span{ defaultStridesAndDilations } : strides; + dilations = dilations.empty() ? Span{ defaultStridesAndDilations } : dilations; + startPadding = startPadding.empty() ? Span{ defaultPadding } : startPadding; + endPadding = endPadding.empty() ? Span{ defaultPadding } : endPadding; + + // Calculate output size, if not explicitly provided + if (outputSizes.empty()) + { + outputSizes.push_back(inputTensor.sizes[0]); // N + outputSizes.push_back(inputTensor.sizes[1]); // C + for (size_t i = 0; i < windowSize.size(); i++) + { + uint32_t paddedInputSize = inputTensor.sizes[2 + i] + startPadding[i] + endPadding[i]; + uint32_t dilatedWindowSize = 1 + (windowSize[i] - 1) * dilations[i]; + uint32_t outputSize = (dilatedWindowSize >= paddedInputSize) ? 1 : (paddedInputSize - dilatedWindowSize) / strides[i] + 1; + outputSizes.push_back(outputSize); + } + } + + TensorDesc outputTensor(inputTensor.dataType, outputSizes, builder->GetTensorPolicy()); + TensorDesc outputIndicesTensor(DML_TENSOR_DATA_TYPE_UINT32, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_MAX_POOLING2_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.OutputIndicesTensor = outputIndices ? outputIndicesTensor.AsPtr() : nullptr; + desc.DimensionCount = static_cast(windowSize.size()); + desc.Strides = strides.data(); + desc.WindowSize = windowSize.data(); + desc.StartPadding = startPadding.data(); + desc.EndPadding = endPadding.data(); + desc.Dilations = dilations.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_MAX_POOLING2, &desc, inputs); + + detail::NodeOutput* outputExpr = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + if (outputIndices) + { + detail::NodeOutput* outputIndicesExpr = builder->CreateNodeOutput(node, 1, std::move(outputIndicesTensor)); + return { outputExpr, outputIndicesExpr }; + } + return { outputExpr, Expression() }; + } + + // Helper for setting parameters for the MaxPooling operator. Sample usage: + // + // auto [out, outIndices] = dml::MaxPoolingBuilder(...) + // .StartPadding(...) + // .EndPadding(...) + // .OutputIndices(...) + // .Build(); + // + // Parameters left unspecified will be defaulted with the same values as dml::MaxPooling(). + class MaxPoolingBuilder + { + public: + MaxPoolingBuilder(Expression input, Span windowSize) + : m_input(input), m_windowSize(windowSize.begin(), windowSize.end()) + {} + + MaxPoolingBuilder& Strides(Span strides) { m_strides.assign(strides.begin(), strides.end()); return *this; } + MaxPoolingBuilder& StartPadding(Span startPadding) { m_startPadding.assign(startPadding.begin(), startPadding.end()); return *this; } + MaxPoolingBuilder& EndPadding(Span endPadding) { m_endPadding.assign(endPadding.begin(), endPadding.end()); return *this; } + MaxPoolingBuilder& Dilations(Span dilations) { m_dilations.assign(dilations.begin(), dilations.end()); return *this; } + MaxPoolingBuilder& OutputIndices(bool outputIndices) { m_outputIndices = outputIndices; return *this; } + MaxPoolingBuilder& OutputSizes(TensorDimensions outputSizes) { m_outputSizes = std::move(outputSizes); return *this; } + + MaxPoolingOutputs Build() const + { + return MaxPooling( + m_input, + m_windowSize, + m_strides, + m_startPadding, + m_endPadding, + m_dilations, + m_outputIndices, + m_outputSizes); + } + + private: + Expression m_input; + SmallVector m_windowSize; + SmallVector m_strides = {}; + SmallVector m_startPadding = {}; + SmallVector m_endPadding = {}; + SmallVector m_dilations = {}; + bool m_outputIndices = false; + TensorDimensions m_outputSizes = {}; + }; + + // --------------------------------------------------------------------------------------------------------------- + + // + // TODO: MaxUnpooling + // + + // + // TODO: ROIPooling + // + + inline Expression Slice( + Expression input, + Span inputWindowOffsets, + Span inputWindowSizes, + Span inputWindowStrides) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDimensions outputSizes(inputTensor.sizes); + + assert(inputWindowOffsets.size() == outputSizes.size()); + assert(inputWindowOffsets.size() == inputWindowStrides.size()); + assert(inputWindowOffsets.size() == inputWindowSizes.size()); + + for (size_t i = 0; i < outputSizes.size(); i++) + { + uint32_t minimumInputSize = (inputWindowSizes[i] - 1) / abs(inputWindowStrides[i]) + 1; + outputSizes[i] = minimumInputSize; + } + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_SLICE1_OPERATOR_DESC sliceDesc = {}; + sliceDesc.InputTensor = inputTensor.AsPtr(); + sliceDesc.OutputTensor = outputTensor.AsPtr(); + sliceDesc.DimensionCount = static_cast(inputWindowOffsets.size()); + sliceDesc.InputWindowOffsets = inputWindowOffsets.data(); + sliceDesc.InputWindowSizes = inputWindowSizes.data(); + sliceDesc.InputWindowStrides = inputWindowStrides.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_SLICE1, &sliceDesc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Cast(Expression input, DML_TENSOR_DATA_TYPE targetDataType) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(targetDataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_CAST_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_CAST, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline std::vector Split( + Expression input, + uint32_t axis, + Span outputAxisSizes) + { + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + uint32_t axisSizeSum = 0; + + std::vector outputTensors; + outputTensors.reserve(outputAxisSizes.size()); + + std::vector outputDescs; + outputDescs.reserve(outputAxisSizes.size()); + + for (uint32_t outputAxisSize : outputAxisSizes) + { + TensorDimensions outputSizes = inputTensor.sizes; + outputSizes[axis] = outputAxisSize; + + TensorDesc tensorDesc(inputTensor.dataType, outputSizes, builder->GetTensorPolicy()); + outputTensors.push_back(std::move(tensorDesc)); + outputDescs.push_back(*outputTensors.back().AsPtr()); + + axisSizeSum += outputAxisSize; + } + + assert(axisSizeSum == inputTensor.sizes[axis]); + + DML_SPLIT_OPERATOR_DESC desc = {}; + desc.Axis = axis; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensors = outputDescs.data(); + desc.OutputCount = static_cast(outputAxisSizes.size()); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_SPLIT, &desc, inputs); + + std::vector outputs; + outputs.reserve(outputAxisSizes.size()); + + for (uint32_t i = 0; i < outputAxisSizes.size(); ++i) + { + outputs.push_back(builder->CreateNodeOutput(node, i, std::move(outputTensors[i]))); + } + + return outputs; + } + + inline Expression Join( + Span inputs, + uint32_t axis) + { + assert(!inputs.empty()); + + detail::GraphBuilder* builder = inputs[0].Impl()->GetGraphBuilder(); + DML_TENSOR_DATA_TYPE dataType = inputs[0].Impl()->GetOutputDesc().dataType; + + TensorDimensions outputSizes = inputs[0].Impl()->GetOutputDesc().sizes; + outputSizes[axis] = 0; + + std::vector inputTensors; + inputTensors.reserve(inputs.size()); + + std::vector inputDescs; + inputDescs.reserve(inputs.size()); + + std::vector inputNodes; + inputNodes.reserve(inputs.size()); + + for (Expression input : inputs) + { + inputTensors.push_back(input.Impl()->GetOutputDesc()); + TensorDesc& inputTensor = inputTensors.back(); + outputSizes[axis] += inputTensor.sizes[axis]; + inputDescs.push_back(*inputTensor.AsPtr()); + inputNodes.push_back(input.Impl()); + } + + TensorDesc outputTensor(dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_JOIN_OPERATOR_DESC desc = {}; + desc.Axis = axis; + desc.InputCount = static_cast(inputDescs.size()); + desc.InputTensors = inputDescs.data(); + desc.OutputTensor = outputTensor.AsPtr(); + + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_JOIN, &desc, inputNodes); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Padding( + Expression input, + DML_PADDING_MODE paddingMode, + float paddingValue, + Span startPadding, + Span endPadding) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDimensions outputSizes = inputTensor.sizes; + + assert(outputSizes.size() == startPadding.size()); + assert(outputSizes.size() == endPadding.size()); + + for (size_t i = 0; i < outputSizes.size(); i++) + { + outputSizes[i] += startPadding[i] + endPadding[i]; + } + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_PADDING_OPERATOR_DESC paddingDesc = {}; + paddingDesc.InputTensor = inputTensor.AsPtr(); + paddingDesc.OutputTensor = outputTensor.AsPtr(); + paddingDesc.PaddingMode = paddingMode; + paddingDesc.PaddingValue = paddingValue; + paddingDesc.DimensionCount = static_cast(startPadding.size()); + paddingDesc.StartPadding = startPadding.data(); + paddingDesc.EndPadding = endPadding.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_PADDING, &paddingDesc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression ValueScale2D( + Expression input, + float scale, + Span bias) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_VALUE_SCALE_2D_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Scale = scale; + desc.ChannelCount = static_cast(bias.size()); + desc.Bias = bias.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_VALUE_SCALE_2D, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Upsample2D(Expression input, DML_SIZE_2D scaleSize, DML_INTERPOLATION_MODE interpolationMode) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + assert(inputTensor.sizes.size() == 4 || inputTensor.sizes.size() == 5); + + uint32_t i = 0; + TensorDimensions outputSizes; + outputSizes.push_back(inputTensor.sizes[i++]); // output[N] = input[N] + outputSizes.push_back(inputTensor.sizes[i++]); // output[C] = input[C] + if (inputTensor.sizes.size() == 5) + { + outputSizes.push_back(inputTensor.sizes[i++]); // output[D] = input[D] + } + outputSizes.push_back(inputTensor.sizes[i++] * scaleSize.Height); // output[H] = input[H] * scaleH + outputSizes.push_back(inputTensor.sizes[i++] * scaleSize.Width); // output[W] = input[W] * scaleW + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_UPSAMPLE_2D_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.ScaleSize = scaleSize; + desc.InterpolationMode = interpolationMode; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_UPSAMPLE_2D, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Gather( + Expression input, + Expression indices, + uint32_t axis, + uint32_t indexDimensions) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc indicesTensor = indices.Impl()->GetOutputDesc(); + + uint32_t dimensionCount = static_cast(inputTensor.sizes.size()); + assert(indicesTensor.sizes.size() == dimensionCount); + assert(axis < dimensionCount); + assert(indexDimensions <= dimensionCount); + + TensorDimensions outputSizes(dimensionCount, 1); + + // All dimensions after the axis should be the same as the input + int outputDim = static_cast(dimensionCount) - 1; + for (; static_cast(outputDim) > axis; --outputDim) + { + outputSizes[outputDim] = inputTensor.sizes[outputDim]; + } + + // All dimensions within the range [axis - indexDimensions, axis] should be the same as the indices + int indexDim = static_cast(dimensionCount) - 1; + for (; outputDim > static_cast(axis) - static_cast(indexDimensions); --outputDim, --indexDim) + { + outputSizes[outputDim] = indicesTensor.sizes[indexDim]; + } + + // All dimensions before (axis - indexDimensions) should be the same as the input + int inputDim = axis - 1; + for (; outputDim >= 0 && inputDim >= 0; --outputDim, --inputDim) + { + outputSizes[outputDim] = inputTensor.sizes[inputDim]; + } + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_GATHER_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.IndicesTensor = indicesTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Axis = axis; + desc.IndexDimensions = indexDimensions; + + detail::NodeOutput* const inputs[] = { input.Impl(), indices.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_GATHER, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression GatherElements( + Expression input, + Expression indices, + uint32_t axis) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc indicesTensor = indices.Impl()->GetOutputDesc(); + + TensorDesc outputTensor(inputTensor.dataType, indicesTensor.sizes, builder->GetTensorPolicy()); + + DML_GATHER_ELEMENTS_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.IndicesTensor = indicesTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Axis = axis; + + detail::NodeOutput* const inputs[] = { input.Impl(), indices.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_GATHER_ELEMENTS, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression GatherND( + Expression input, + Expression indices, + uint32_t inputDimensionCount, + uint32_t indicesDimensionCount, + uint32_t batchDimensionCount) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc indicesTensor = indices.Impl()->GetOutputDesc(); + + assert(inputDimensionCount >= 1u && inputDimensionCount <= inputTensor.sizes.size()); + assert(indicesDimensionCount >= 1u && indicesDimensionCount <= indicesTensor.sizes.size()); + assert(batchDimensionCount < inputDimensionCount); + assert(batchDimensionCount < indicesDimensionCount); + + uint32_t numberOfCoordinatesPerIndex = indicesTensor.sizes.back(); + assert(numberOfCoordinatesPerIndex >= 1u && numberOfCoordinatesPerIndex <= inputDimensionCount - batchDimensionCount); + + uint32_t numberOfOutputDimensionsFromInput = inputDimensionCount - batchDimensionCount - numberOfCoordinatesPerIndex; + uint32_t outputPaddingAmount = static_cast(inputTensor.sizes.size()) - (indicesDimensionCount + numberOfOutputDimensionsFromInput - 1); + + TensorDimensions outputSizes(outputPaddingAmount, 1); + outputSizes.insert(outputSizes.end(), indicesTensor.sizes.end() - indicesDimensionCount, indicesTensor.sizes.end() - 1); + outputSizes.insert(outputSizes.end(), inputTensor.sizes.end() - numberOfOutputDimensionsFromInput, inputTensor.sizes.end()); + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_GATHER_ND1_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.IndicesTensor = indicesTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.InputDimensionCount = inputDimensionCount; + desc.IndicesDimensionCount = indicesDimensionCount; + desc.BatchDimensionCount = batchDimensionCount; + + detail::NodeOutput* const inputs[] = { input.Impl(), indices.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_GATHER_ND1, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression ScatterElements( + Expression input, + Expression indices, + Expression updates, + uint32_t axis) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc indicesTensor = indices.Impl()->GetOutputDesc(); + TensorDesc updatesTensor = updates.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_SCATTER_ELEMENTS_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.IndicesTensor = indicesTensor.AsPtr(); + desc.UpdatesTensor = updatesTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Axis = axis; + + detail::NodeOutput* const inputs[] = { input.Impl(), indices.Impl(), updates.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_SCATTER_ELEMENTS, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression ScatterND( + Expression input, + Expression indices, + Expression updates, + uint32_t inputDimensionCount, + uint32_t indicesDimensionCount) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc indicesTensor = indices.Impl()->GetOutputDesc(); + TensorDesc updatesTensor = updates.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_SCATTER_ND_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.IndicesTensor = indicesTensor.AsPtr(); + desc.UpdatesTensor = updatesTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.InputDimensionCount = inputDimensionCount; + desc.IndicesDimensionCount = indicesDimensionCount; + + detail::NodeOutput* const inputs[] = { input.Impl(), indices.Impl(), updates.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_SCATTER_ND, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression SpaceToDepth( + Expression input, + uint32_t blockSize, + DML_DEPTH_SPACE_ORDER order = DML_DEPTH_SPACE_ORDER_DEPTH_COLUMN_ROW) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + assert(inputTensor.sizes.size() == 4); + + dml::TensorDesc::Dimensions outputSizes = { + inputTensor.sizes[0], + inputTensor.sizes[1] * blockSize * blockSize, + inputTensor.sizes[2] / blockSize, + inputTensor.sizes[3] / blockSize + }; + + TensorDesc outputTensor(inputTensor.dataType, outputSizes, builder->GetTensorPolicy()); + + DML_SPACE_TO_DEPTH1_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.BlockSize = blockSize; + desc.Order = order; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_SPACE_TO_DEPTH1, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression DepthToSpace( + Expression input, + uint32_t blockSize, + DML_DEPTH_SPACE_ORDER order = DML_DEPTH_SPACE_ORDER_DEPTH_COLUMN_ROW) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + assert(inputTensor.sizes.size() == 4); + + dml::TensorDesc::Dimensions outputSizes = { + inputTensor.sizes[0], + inputTensor.sizes[1] / (blockSize * blockSize), + inputTensor.sizes[2] * blockSize, + inputTensor.sizes[3] * blockSize + }; + + TensorDesc outputTensor(inputTensor.dataType, outputSizes, builder->GetTensorPolicy()); + + DML_DEPTH_TO_SPACE1_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.BlockSize = blockSize; + desc.Order = order; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_DEPTH_TO_SPACE1, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression Tile(Expression input, Span repeats) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDimensions outputSizes = input.GetOutputDesc().sizes; + + assert(repeats.size() == outputSizes.size()); + + for (size_t i = 0; i < repeats.size(); ++i) + { + outputSizes[i] *= repeats[i]; + } + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_TILE_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.RepeatsCount = static_cast(repeats.size()); + desc.Repeats = repeats.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_TILE, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + struct TopKOutputs + { + Expression value; + Expression index; + }; + + inline TopKOutputs TopK(Expression input, uint32_t axis, uint32_t k, DML_AXIS_DIRECTION axisDirection) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + TensorDimensions outputSizes = inputTensor.sizes; + outputSizes.back() = k; + + TensorDesc outputValueTensor(inputTensor.dataType, outputSizes, builder->GetTensorPolicy()); + TensorDesc outputIndexTensor(DML_TENSOR_DATA_TYPE_UINT32, outputSizes, builder->GetTensorPolicy()); + + DML_TOP_K1_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputValueTensor = outputValueTensor.AsPtr(); + desc.OutputIndexTensor = outputIndexTensor.AsPtr(); + desc.Axis = axis; + desc.K = k; + desc.AxisDirection = axisDirection; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_TOP_K1, &desc, inputs); + detail::NodeOutput* outputValue = builder->CreateNodeOutput(node, 0, std::move(outputValueTensor)); + detail::NodeOutput* outputIndex = builder->CreateNodeOutput(node, 1, std::move(outputIndexTensor)); + + return { outputValue, outputIndex }; + } + + inline Expression BatchNormalization( + Expression input, + Expression mean, + Expression variance, + Expression scale, + Expression bias, + bool spatial, + float epsilon, + FusedActivation fusedActivation = FusedActivation::None()) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc meanTensor = mean.Impl()->GetOutputDesc(); + TensorDesc varianceTensor = variance.Impl()->GetOutputDesc(); + TensorDesc scaleTensor = scale.Impl()->GetOutputDesc(); + TensorDesc biasTensor = bias.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + detail::FusedActivationStorage storage; + + DML_BATCH_NORMALIZATION_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.MeanTensor = meanTensor.AsPtr(); + desc.VarianceTensor = varianceTensor.AsPtr(); + desc.ScaleTensor = scaleTensor.AsPtr(); + desc.BiasTensor = biasTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Spatial = spatial; + desc.Epsilon = epsilon; + desc.FusedActivation = detail::GetFusedActivationPtr(fusedActivation, &storage); + + detail::NodeOutput* const inputs[] = { input.Impl(), mean.Impl(), variance.Impl(), scale.Impl(), bias.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_BATCH_NORMALIZATION, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression BatchNormalization( + Expression input, + Expression mean, + Expression variance, + Expression scale, + Expression bias, + bool spatial, + float epsilon, + const DML_OPERATOR_DESC* fusedActivation) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc meanTensor = mean.Impl()->GetOutputDesc(); + TensorDesc varianceTensor = variance.Impl()->GetOutputDesc(); + TensorDesc scaleTensor = scale.Impl()->GetOutputDesc(); + TensorDesc biasTensor = bias.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_BATCH_NORMALIZATION_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.MeanTensor = meanTensor.AsPtr(); + desc.VarianceTensor = varianceTensor.AsPtr(); + desc.ScaleTensor = scaleTensor.AsPtr(); + desc.BiasTensor = biasTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Spatial = spatial; + desc.Epsilon = epsilon; + desc.FusedActivation = fusedActivation; + + detail::NodeOutput* const inputs[] = { input.Impl(), mean.Impl(), variance.Impl(), scale.Impl(), bias.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_BATCH_NORMALIZATION, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + +#if DML_TARGET_VERSION >= 0x3100 + + struct BatchNormalizationGradOutputs + { + Expression gradient; + Expression scaleGradient; + Expression biasGradient; + }; + + inline BatchNormalizationGradOutputs BatchNormalizationGrad( + Expression input, + Expression inputGradient, + Expression mean, + Expression variance, + Expression scale, + float epsilon) + { + dml::detail::GraphBuilder* builder = mean.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc inputGradientTensor = inputGradient.Impl()->GetOutputDesc(); + TensorDesc meanTensor = mean.Impl()->GetOutputDesc(); + TensorDesc varianceTensor = variance.Impl()->GetOutputDesc(); + TensorDesc scaleTensor = scale.Impl()->GetOutputDesc(); + TensorDesc outputGradientTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + TensorDesc outputScaleGradientTensor(meanTensor.dataType, meanTensor.sizes, builder->GetTensorPolicy()); + TensorDesc outputBiasGradientTensor(meanTensor.dataType, meanTensor.sizes, builder->GetTensorPolicy()); + + DML_BATCH_NORMALIZATION_GRAD_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.InputGradientTensor = inputGradientTensor.AsPtr(); + desc.MeanTensor = meanTensor.AsPtr(); + desc.VarianceTensor = varianceTensor.AsPtr(); + desc.ScaleTensor = scaleTensor.AsPtr(); + desc.Epsilon = epsilon; + + desc.OutputGradientTensor = outputGradientTensor.AsPtr(); + desc.OutputScaleGradientTensor = outputScaleGradientTensor.AsPtr(); + desc.OutputBiasGradientTensor = outputBiasGradientTensor.AsPtr(); + + dml::detail::NodeOutput* const inputs[] = { input.Impl(), inputGradient.Impl(), mean.Impl(), variance.Impl(), scale.Impl() }; + dml::detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_BATCH_NORMALIZATION_GRAD, &desc, inputs); + + BatchNormalizationGradOutputs outputValues; + outputValues.gradient = builder->CreateNodeOutput(node, 0, *desc.OutputGradientTensor); + outputValues.scaleGradient = builder->CreateNodeOutput(node, 1, *desc.OutputScaleGradientTensor); + outputValues.biasGradient = builder->CreateNodeOutput(node, 2, *desc.OutputBiasGradientTensor); + + return outputValues; + } + +#endif // DML_TARGET_VERSION >= 0x3100 + +#if DML_TARGET_VERSION >= 0x4100 + struct BatchNormalizationTrainingOutputs + { + Expression output; + Expression mean; + Expression variance; + }; + + inline BatchNormalizationTrainingOutputs BatchNormalizationTraining( + Expression input, + Expression scale, + Expression bias, + Optional fusedAdd, + float epsilon, + FusedActivation fusedActivation = FusedActivation::None()) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc scaleTensor = scale.Impl()->GetOutputDesc(); + TensorDesc biasTensor = bias.Impl()->GetOutputDesc(); + + TensorDesc fusedAddTensor; + if (fusedAdd) + { + fusedAddTensor = fusedAdd->Impl()->GetOutputDesc(); + } + + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + TensorDesc outputMeanTensor(inputTensor.dataType, scaleTensor.sizes, builder->GetTensorPolicy()); + TensorDesc outputVarianceTensor(inputTensor.dataType, scaleTensor.sizes, builder->GetTensorPolicy()); + + detail::FusedActivationStorage storage; + + DML_BATCH_NORMALIZATION_TRAINING_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.ScaleTensor = scaleTensor.AsPtr(); + desc.BiasTensor = biasTensor.AsPtr(); + desc.FusedAddTensor = fusedAdd.has_value() ? fusedAddTensor.AsPtr() : nullptr; + desc.OutputTensor = outputTensor.AsPtr(); + desc.OutputMeanTensor = outputMeanTensor.AsPtr(); + desc.OutputVarianceTensor = outputVarianceTensor.AsPtr(); + desc.Epsilon = epsilon; + desc.FusedActivation = detail::GetFusedActivationPtr(fusedActivation, &storage); + + detail::NodeOutput* const inputs[] = { input.Impl(), scale.Impl(), bias.Impl(), fusedAdd ? fusedAdd->Impl() : nullptr }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_BATCH_NORMALIZATION_TRAINING, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + detail::NodeOutput* outputMean = builder->CreateNodeOutput(node, 1, std::move(outputMeanTensor)); + detail::NodeOutput* outputVariance = builder->CreateNodeOutput(node, 2, std::move(outputVarianceTensor)); + + return {output, outputMean, outputVariance}; + } + + inline BatchNormalizationGradOutputs BatchNormalizationTrainingGrad( + Expression input, + Expression inputGradient, + Expression mean, + Expression variance, + Expression scale, + float epsilon) + { + dml::detail::GraphBuilder* builder = mean.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc inputGradientTensor = inputGradient.Impl()->GetOutputDesc(); + TensorDesc meanTensor = mean.Impl()->GetOutputDesc(); + TensorDesc varianceTensor = variance.Impl()->GetOutputDesc(); + TensorDesc scaleTensor = scale.Impl()->GetOutputDesc(); + TensorDesc outputGradientTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + TensorDesc outputScaleGradientTensor(meanTensor.dataType, meanTensor.sizes, builder->GetTensorPolicy()); + TensorDesc outputBiasGradientTensor(meanTensor.dataType, meanTensor.sizes, builder->GetTensorPolicy()); + + DML_BATCH_NORMALIZATION_TRAINING_GRAD_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.InputGradientTensor = inputGradientTensor.AsPtr(); + desc.MeanTensor = meanTensor.AsPtr(); + desc.VarianceTensor = varianceTensor.AsPtr(); + desc.ScaleTensor = scaleTensor.AsPtr(); + desc.Epsilon = epsilon; + + desc.OutputGradientTensor = outputGradientTensor.AsPtr(); + desc.OutputScaleGradientTensor = outputScaleGradientTensor.AsPtr(); + desc.OutputBiasGradientTensor = outputBiasGradientTensor.AsPtr(); + + dml::detail::NodeOutput* const inputs[] = { input.Impl(), inputGradient.Impl(), mean.Impl(), variance.Impl(), scale.Impl() }; + dml::detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_BATCH_NORMALIZATION_TRAINING_GRAD, &desc, inputs); + + BatchNormalizationGradOutputs outputValues; + outputValues.gradient = builder->CreateNodeOutput(node, 0, *desc.OutputGradientTensor); + outputValues.scaleGradient = builder->CreateNodeOutput(node, 1, *desc.OutputScaleGradientTensor); + outputValues.biasGradient = builder->CreateNodeOutput(node, 2, *desc.OutputBiasGradientTensor); + + return outputValues; + } +#endif // DML_TARGET_VERSION >= 0x4100 + + inline Expression MeanVarianceNormalization( + Expression input, + Optional scale, + Optional bias, + Span axes, + bool normalizeVariance, + float epsilon, + FusedActivation fusedActivation = FusedActivation::None()) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + TensorDesc scaleTensor; + TensorDesc biasTensor; + + if (scale) + { + scaleTensor = scale->Impl()->GetOutputDesc(); + } + if (bias) + { + biasTensor = bias->Impl()->GetOutputDesc(); + } + + detail::FusedActivationStorage storage; + + DML_MEAN_VARIANCE_NORMALIZATION1_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.ScaleTensor = scale ? scaleTensor.AsPtr() : nullptr; + desc.BiasTensor = bias ? biasTensor.AsPtr() : nullptr; + desc.OutputTensor = outputTensor.AsPtr(); + desc.AxisCount = static_cast(axes.size()); + desc.Axes = axes.data(); + desc.NormalizeVariance = normalizeVariance; + desc.Epsilon = epsilon; + desc.FusedActivation = detail::GetFusedActivationPtr(fusedActivation, &storage); + + detail::NodeOutput* const inputs[] = + { + input.Impl(), + scale ? scale->Impl() : nullptr, + bias ? bias->Impl() : nullptr + }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_MEAN_VARIANCE_NORMALIZATION1, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression LocalResponseNormalization( + Expression input, + bool crossChannel, + uint32_t localSize, + float alpha, + float beta, + float bias) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_LOCAL_RESPONSE_NORMALIZATION_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.CrossChannel = crossChannel; + desc.LocalSize = localSize; + desc.Alpha = alpha; + desc.Beta = beta; + desc.Bias = bias; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_LOCAL_RESPONSE_NORMALIZATION, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + // + // TODO: LpNormalization + // + + // + // TODO: RNN + // + + // + // TODO: LSTM + // + + enum class GRUOutputOptions + { + Both, + Sequence, + Single, + }; + + struct GRUOutputs + { + Expression sequence; + Expression single; + }; + + inline GRUOutputs GRU( + Expression input, + Expression weight, + Expression recurrence, + Optional bias, + Optional hiddenInit, + Optional sequenceLengths, + Span activationDescs, + DML_RECURRENT_NETWORK_DIRECTION direction, + bool linearBeforeReset, + GRUOutputOptions outputOptions) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc weightTensor = weight.Impl()->GetOutputDesc(); + TensorDesc recurrenceTensor = recurrence.Impl()->GetOutputDesc(); + TensorDesc biasTensor; + TensorDesc hiddenInitTensor; + TensorDesc sequenceLengthsTensor; + TensorDesc outputSequenceTensor; + TensorDesc outputSingleTensor; + if (bias) + { + biasTensor = bias->Impl()->GetOutputDesc(); + } + if (hiddenInit) + { + hiddenInitTensor = hiddenInit->Impl()->GetOutputDesc(); + } + if (sequenceLengths) + { + sequenceLengthsTensor = sequenceLengths->Impl()->GetOutputDesc(); + } + + TensorDesc::Dimensions outputSequenceSizes(4); + TensorDesc::Dimensions outputSingleSizes(4); + uint32_t directionCount = (direction == DML_RECURRENT_NETWORK_DIRECTION_BIDIRECTIONAL) ? 2 : 1; + if (outputOptions == GRUOutputOptions::Sequence || outputOptions == GRUOutputOptions::Both) + { + outputSequenceSizes[0] = inputTensor.sizes[1]; // SequenceLength + outputSequenceSizes[1] = directionCount; + outputSequenceSizes[2] = inputTensor.sizes[2]; // BatchSize + outputSequenceSizes[3] = recurrenceTensor.sizes[3]; // HiddenSize + outputSequenceTensor = TensorDesc(inputTensor.dataType, outputSequenceSizes, builder->GetTensorPolicy()); + } + if (outputOptions == GRUOutputOptions::Single || outputOptions == GRUOutputOptions::Both) + { + outputSingleSizes[0] = 1; + outputSingleSizes[1] = directionCount; + outputSingleSizes[2] = inputTensor.sizes[2]; // BatchSize + outputSingleSizes[3] = recurrenceTensor.sizes[3]; // HiddenSize + outputSingleTensor = TensorDesc(inputTensor.dataType, outputSingleSizes, builder->GetTensorPolicy()); + } + + uint32_t activationCount = static_cast(activationDescs.size()); + if (activationCount > 4) + { + DMLX_THROW(E_INVALIDARG); + } + + detail::FusedActivationStorage storage[4]; + DML_OPERATOR_DESC activationDescArray[4]; + for (uint32_t i = 0; i < activationCount; ++i) + { + activationDescArray[i] = *detail::GetFusedActivationPtr(activationDescs[i], &storage[i]); + } + + DML_GRU_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.WeightTensor = weightTensor.AsPtr(); + desc.RecurrenceTensor = recurrenceTensor.AsPtr(); + desc.BiasTensor = bias ? biasTensor.AsPtr() : nullptr; + desc.HiddenInitTensor = hiddenInit ? hiddenInitTensor.AsPtr() : nullptr; + desc.SequenceLengthsTensor = sequenceLengths ? sequenceLengthsTensor.AsPtr() : nullptr; + desc.OutputSequenceTensor = outputSequenceTensor.sizes.empty() ? nullptr : outputSequenceTensor.AsPtr(); + desc.OutputSingleTensor = outputSingleTensor.sizes.empty() ? nullptr : outputSingleTensor.AsPtr(); + desc.ActivationDescCount = activationCount; + desc.ActivationDescs = activationDescArray; + desc.Direction = direction; + desc.LinearBeforeReset = linearBeforeReset; + + detail::NodeOutput* const inputs[] = + { + input.Impl(), + weight.Impl(), + recurrence.Impl(), + bias ? bias->Impl() : nullptr, + hiddenInit ? hiddenInit->Impl() : nullptr, + sequenceLengths ? sequenceLengths->Impl() : nullptr + }; + + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_GRU, &desc, inputs); + + detail::NodeOutput* outputSequenceExpr = nullptr; + detail::NodeOutput* outputSingleExpr = nullptr; + if (outputOptions == GRUOutputOptions::Sequence || outputOptions == GRUOutputOptions::Both) + { + outputSequenceExpr = builder->CreateNodeOutput(node, 0, std::move(outputSequenceTensor)); + } + if (outputOptions == GRUOutputOptions::Single || outputOptions == GRUOutputOptions::Both) + { + outputSingleExpr = builder->CreateNodeOutput(node, 1, std::move(outputSingleTensor)); + } + return { outputSequenceExpr, outputSingleExpr }; + } + + // + // TODO: DiagonalMatrix + // + + inline Expression OneHot( + Expression indices, + Expression values, + uint32_t outputLength, + uint32_t axis) + { + detail::GraphBuilder* builder = indices.Impl()->GetGraphBuilder(); + TensorDesc indicesTensor = indices.Impl()->GetOutputDesc(); + TensorDesc valuesTensor = values.Impl()->GetOutputDesc(); + + assert(axis < static_cast(indicesTensor.sizes.size())); + + // The output and indices sizes must all match except for the active axis, which is supplied as outputLength. + TensorDimensions outputSizes = indicesTensor.sizes; + outputSizes[axis] = outputLength; + + TensorDesc outputTensor(valuesTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_ONE_HOT_OPERATOR_DESC desc = {}; + desc.IndicesTensor = indicesTensor.AsPtr(); + desc.ValuesTensor = valuesTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Axis = axis; + + detail::NodeOutput* const inputs[] = { indices.Impl(), values.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ONE_HOT, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + // If not specified, parameters are defaulted to the following values: + // Scales = computed by dividing the output sizes by the input sizes + // InputPixelOffsets = 0.5f for each dimension + // OutputPixelOffsets = -0.5f for each dimension + inline Expression Resample( + Expression input, + TensorDimensions outputSizes, + DML_INTERPOLATION_MODE mode, + Span scales = {}, + Span inputPixelOffsets = {}, + Span outputPixelOffsets = {}) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + uint32_t dimensionCount = static_cast(inputTensor.sizes.size()); + assert(outputSizes.size() == dimensionCount); + + SmallVector defaultScales; + if (scales.empty()) + { + for (uint32_t i = 0; i < dimensionCount; ++i) + { + defaultScales.push_back(static_cast(outputSizes[i]) / static_cast(inputTensor.sizes[i])); + } + scales = defaultScales; + } + + SmallVector defaultInputPixelOffsets; + if (inputPixelOffsets.empty()) + { + defaultInputPixelOffsets.assign(dimensionCount, 0.5f); + inputPixelOffsets = defaultInputPixelOffsets; + } + + SmallVector defaultOutputPixelOffsets; + if (outputPixelOffsets.empty()) + { + defaultOutputPixelOffsets.assign(dimensionCount, -0.5f); + outputPixelOffsets = defaultOutputPixelOffsets; + } + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_RESAMPLE1_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.InterpolationMode = mode; + desc.DimensionCount = static_cast(scales.size()); + desc.Scales = scales.data(); + desc.InputPixelOffsets = inputPixelOffsets.data(); + desc.OutputPixelOffsets = outputPixelOffsets.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_RESAMPLE1, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression FillValueConstant( + Graph& graph, + TensorDimensions outputSizes, + DML_TENSOR_DATA_TYPE valueDataType, + DML_SCALAR_UNION value) + { + detail::GraphBuilder* builder = graph.Impl(); + TensorDesc outputTensor(valueDataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_FILL_VALUE_CONSTANT_OPERATOR_DESC desc = {}; + desc.OutputTensor = outputTensor.AsPtr(); + desc.ValueDataType = valueDataType; + desc.Value = value; + + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_FILL_VALUE_CONSTANT, &desc, {}); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression FillValueSequence( + Graph& graph, + TensorDimensions outputSizes, + DML_TENSOR_DATA_TYPE valueDataType, + DML_SCALAR_UNION valueStart, + DML_SCALAR_UNION valueDelta) + { + detail::GraphBuilder* builder = graph.Impl(); + TensorDesc outputTensor(valueDataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_FILL_VALUE_SEQUENCE_OPERATOR_DESC desc = {}; + desc.OutputTensor = outputTensor.AsPtr(); + desc.ValueDataType = valueDataType; + desc.ValueStart = valueStart; + desc.ValueDelta = valueDelta; + + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_FILL_VALUE_SEQUENCE, &desc, {}); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression CumulativeSummation( + Expression input, + uint32_t axis, + DML_AXIS_DIRECTION axisDirection, + bool hasExclusiveSum) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_CUMULATIVE_SUMMATION_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Axis = axis; + desc.AxisDirection = axisDirection; + desc.HasExclusiveSum = hasExclusiveSum; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_CUMULATIVE_SUMMATION, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + +#if DML_TARGET_VERSION >= 0x3100 + + inline Expression CumulativeProduct( + Expression input, + uint32_t axis, + DML_AXIS_DIRECTION axisDirection, + bool hasExclusiveProduct) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_CUMULATIVE_PRODUCT_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.Axis = axis; + desc.AxisDirection = axisDirection; + desc.HasExclusiveProduct = hasExclusiveProduct; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_CUMULATIVE_PRODUCT, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + +#endif // DML_TARGET_VERSION >= 0x3100 + + inline Expression ReverseSubsequences( + Expression input, + Expression sequenceLengths, + uint32_t axis) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc sequenceLengthsTensor = sequenceLengths.Impl()->GetOutputDesc(); + TensorDesc outputTensor(inputTensor.dataType, inputTensor.sizes, builder->GetTensorPolicy()); + + DML_REVERSE_SUBSEQUENCES_OPERATOR_DESC reverseDesc = {}; + reverseDesc.InputTensor = inputTensor.AsPtr(); + reverseDesc.SequenceLengthsTensor = sequenceLengthsTensor.AsPtr(); + reverseDesc.OutputTensor = outputTensor.AsPtr(); + reverseDesc.Axis = axis; + + detail::NodeOutput* const inputs[] = { input.Impl(), sequenceLengths.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_REVERSE_SUBSEQUENCES, &reverseDesc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + // + // TODO: MatrixMultiplyInteger + // + + // + // TODO: QuantizedLinearMatrixMultiply + // + + // + // TODO: ConvolutionInteger + // + + // + // TODO: QuantizedLinearConvolution + // + + // + // TODO: ReluGrad + // + + // + // TODO: AveragePoolingGrad + // + + // + // TODO: MaxPoolingGrad + // + + struct RandomGeneratorOutputs + { + Expression values; + Expression state; // Only valid if outputState = true is supplied to RandomGenerator + }; + + inline RandomGeneratorOutputs RandomGenerator( + Expression inputState, + TensorDimensions outputSizes, + bool outputState = true, + DML_RANDOM_GENERATOR_TYPE type = DML_RANDOM_GENERATOR_TYPE_PHILOX_4X32_10) + { + detail::GraphBuilder* builder = inputState.Impl()->GetGraphBuilder(); + + TensorDesc inputStateTensor = inputState.Impl()->GetOutputDesc(); + TensorDesc outputTensor(DML_TENSOR_DATA_TYPE_UINT32, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_RANDOM_GENERATOR_OPERATOR_DESC desc = {}; + desc.Type = type; + desc.InputStateTensor = inputStateTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + if (outputState) + { + // Input and output state have the same TensorDesc. + desc.OutputStateTensor = inputStateTensor.AsPtr(); + } + + RandomGeneratorOutputs out; + + detail::NodeOutput* const inputs[] = { inputState.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_RANDOM_GENERATOR, &desc, inputs); + out.values = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + if (outputState) + { + TensorDesc outputStateTensor = inputStateTensor; + out.state = builder->CreateNodeOutput(node, 1, std::move(outputStateTensor)); + } + + return out; + } + + struct NonZeroCoordinatesOutputs + { + Expression count; + Expression coordinates; + }; + inline NonZeroCoordinatesOutputs NonZeroCoordinates(Expression input) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + const auto& inputTensorSizes = inputTensor.sizes; + uint32_t dimensionCount = static_cast(inputTensorSizes.size()); + + TensorDimensions outputCountSizes = {1}; + uint32_t totalElements = 1; + for (uint32_t i = 0; i < dimensionCount; ++i) + { + totalElements *= inputTensorSizes[i]; + } + TensorDesc outputCountTensor(DML_TENSOR_DATA_TYPE_UINT32, outputCountSizes, builder->GetTensorPolicy()); + TensorDesc outputCoordinatesTensor(DML_TENSOR_DATA_TYPE_UINT32, {totalElements, dimensionCount}, builder->GetTensorPolicy()); + + DML_NONZERO_COORDINATES_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.OutputCountTensor = outputCountTensor.AsPtr(); + desc.OutputCoordinatesTensor = outputCoordinatesTensor.AsPtr(); + + NonZeroCoordinatesOutputs output; + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_NONZERO_COORDINATES, &desc, inputs); + output.count = builder->CreateNodeOutput(node, 0, std::move(outputCountTensor)); + output.coordinates = builder->CreateNodeOutput(node, 1, std::move(outputCoordinatesTensor)); + return output; + } + + // If not specified, parameters are defaulted to the following values: + // Scales = computed by dividing the input sizes by the output sizes + // InputPixelOffsets = 0.5f for each dimension + // OutputPixelOffsets = -0.5f for each dimension + inline Expression ResampleGrad( + Expression input, + TensorDimensions outputSizes, + DML_INTERPOLATION_MODE mode, + Span scales = {}, + Span inputPixelOffsets = {}, + Span outputPixelOffsets = {}) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + uint32_t dimensionCount = static_cast(inputTensor.sizes.size()); + assert(outputSizes.size() == dimensionCount); + + SmallVector defaultScales; + if (scales.empty()) + { + for (uint32_t i = 0; i < dimensionCount; ++i) + { + defaultScales.push_back(static_cast(inputTensor.sizes[i]) / static_cast(outputSizes[i])); + } + scales = defaultScales; + } + + SmallVector defaultInputPixelOffsets; + if (inputPixelOffsets.empty()) + { + defaultInputPixelOffsets.assign(dimensionCount, 0.5f); + inputPixelOffsets = defaultInputPixelOffsets; + } + + SmallVector defaultOutputPixelOffsets; + if (outputPixelOffsets.empty()) + { + defaultOutputPixelOffsets.assign(dimensionCount, -0.5f); + outputPixelOffsets = defaultOutputPixelOffsets; + } + + TensorDesc outputTensor(inputTensor.dataType, std::move(outputSizes), builder->GetTensorPolicy()); + + DML_RESAMPLE_GRAD_OPERATOR_DESC desc = {}; + desc.InputGradientTensor = inputTensor.AsPtr(); + desc.OutputGradientTensor = outputTensor.AsPtr(); + desc.InterpolationMode = mode; + desc.DimensionCount = static_cast(scales.size()); + desc.Scales = scales.data(); + desc.InputPixelOffsets = inputPixelOffsets.data(); + desc.OutputPixelOffsets = outputPixelOffsets.data(); + + detail::NodeOutput* const inputs[] = { input.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_RESAMPLE_GRAD, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + + inline Expression SliceGrad( + Expression inputGradient, + TensorDimensions outputGradientSizes, + Span inputWindowOffsets, + Span inputWindowSizes, + Span inputWindowStrides) + { + detail::GraphBuilder* builder = inputGradient.Impl()->GetGraphBuilder(); + + TensorDesc inputGradientTensor = inputGradient.Impl()->GetOutputDesc(); + + assert(inputWindowOffsets.size() == inputGradientTensor.sizes.size()); + assert(inputWindowOffsets.size() == outputGradientSizes.size()); + assert(inputWindowOffsets.size() == inputWindowStrides.size()); + assert(inputWindowOffsets.size() == inputWindowSizes.size()); + + TensorDesc outputGradientTensor(inputGradientTensor.dataType, std::move(outputGradientSizes), builder->GetTensorPolicy()); + + DML_SLICE_GRAD_OPERATOR_DESC sliceGradDesc = {}; + sliceGradDesc.InputGradientTensor = inputGradientTensor.AsPtr(); + sliceGradDesc.OutputGradientTensor = outputGradientTensor.AsPtr(); + sliceGradDesc.DimensionCount = static_cast(inputWindowOffsets.size()); + sliceGradDesc.InputWindowOffsets = inputWindowOffsets.data(); + sliceGradDesc.InputWindowSizes = inputWindowSizes.data(); + sliceGradDesc.InputWindowStrides = inputWindowStrides.data(); + + detail::NodeOutput* const inputs[] = { inputGradient.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_SLICE_GRAD, &sliceGradDesc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputGradientTensor)); + + return output; + } + + // + // TODO: AdamOptimizer + // + + // + // TODO: Argmin + // + + // + // TODO: Argmax + // + +#if DML_TARGET_VERSION >= 0x4000 + + inline Expression RoiAlign( + Expression input, + Expression roi, + Expression batchIndices, + DML_REDUCE_FUNCTION reductionFunction, + DML_INTERPOLATION_MODE interpolationMode, + float spatialScaleX, + float spatialScaleY, + float inputPixelOffset, + float outputPixelOffset, + float outOfBoundsInputValue, + uint32_t minimumSamplesPerOutput, + uint32_t maximumSamplesPerOutput, + bool alignRegionsToCorners, + uint32_t outputHeight, + uint32_t outputWidth) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc roiTensor = roi.Impl()->GetOutputDesc(); + TensorDesc batchIndicesTensor = batchIndices.Impl()->GetOutputDesc(); + + uint32_t channelCount = inputTensor.sizes[1]; + uint32_t roiCount = roiTensor.sizes.size() < 2 ? 1u : roiTensor.sizes[roiTensor.sizes.size() - 2]; + + TensorDesc::Dimensions outputSizes({ + roiCount, + channelCount, + outputHeight, + outputWidth, + }); + + TensorDesc outputTensor(inputTensor.dataType, outputSizes, builder->GetTensorPolicy()); + + DML_ROI_ALIGN1_OPERATOR_DESC desc = {}; + desc.InputTensor = inputTensor.AsPtr(); + desc.ROITensor = roiTensor.AsPtr(); + desc.BatchIndicesTensor = batchIndicesTensor.AsPtr(); + desc.OutputTensor = outputTensor.AsPtr(); + desc.ReductionFunction = reductionFunction; + desc.InterpolationMode = interpolationMode; + desc.SpatialScaleX = spatialScaleX; + desc.SpatialScaleY = spatialScaleY; + desc.InputPixelOffset = inputPixelOffset; + desc.OutputPixelOffset = outputPixelOffset; + desc.OutOfBoundsInputValue = outOfBoundsInputValue; + desc.MinimumSamplesPerOutput = minimumSamplesPerOutput; + desc.MaximumSamplesPerOutput = maximumSamplesPerOutput; + desc.AlignRegionsToCorners = alignRegionsToCorners; + + detail::NodeOutput* const inputs[] = { input.Impl(), roi.Impl(), batchIndices.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_ROI_ALIGN1, &desc, inputs); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(outputTensor)); + + return output; + } + +#endif // DML_TARGET_VERSION >= 0x4000 + +#if DML_TARGET_VERSION >= 0x4100 + struct RoiAlignGradOutputs + { + Expression outputGradient; + Expression outputROIGradient; + }; + + inline RoiAlignGradOutputs RoiAlignGrad( + Optional input, + Expression inputGradient, + Expression roi, + Expression batchIndices, + DML_REDUCE_FUNCTION reductionFunction, + DML_INTERPOLATION_MODE interpolationMode, + float spatialScaleX, + float spatialScaleY, + float inputPixelOffset, + float outputPixelOffset, + uint32_t minimumSamplesPerOutput, + uint32_t maximumSamplesPerOutput, + bool alignRegionsToCorners, + uint32_t batchSize, + uint32_t imageHeight, + uint32_t imageWidth, + bool computeOutputGradient, + bool computeOutputROIGradient) + { + detail::GraphBuilder* builder = inputGradient.Impl()->GetGraphBuilder(); + + TensorDesc inputTensor = input.has_value() ? input->Impl()->GetOutputDesc() : TensorDesc(); + TensorDesc inputGradientTensor = inputGradient.Impl()->GetOutputDesc(); + TensorDesc roiTensor = roi.Impl()->GetOutputDesc(); + TensorDesc batchIndicesTensor = batchIndices.Impl()->GetOutputDesc(); + + assert(computeOutputGradient || computeOutputROIGradient); + assert(inputGradientTensor.sizes.size() > 1); + + TensorDesc outputGradientTensor; + if (computeOutputGradient) + { + TensorDesc::Dimensions outputGradientSizes({ + batchSize, + inputGradientTensor.sizes[1], + imageHeight, + imageWidth, + }); + + outputGradientTensor = TensorDesc(inputGradientTensor.dataType, outputGradientSizes, builder->GetTensorPolicy()); + } + + TensorDesc outputROIGradientTensor = computeOutputROIGradient ? TensorDesc(roiTensor.dataType, roiTensor.sizes, builder->GetTensorPolicy()) : TensorDesc(); + assert(!computeOutputROIGradient || outputROIGradientTensor.sizes == roiTensor.sizes); + + DML_ROI_ALIGN_GRAD_OPERATOR_DESC desc = {}; + desc.InputTensor = input ? inputTensor.AsPtr() : nullptr; + desc.InputGradientTensor = inputGradientTensor.AsPtr(); + desc.ROITensor = roiTensor.AsPtr(); + desc.BatchIndicesTensor = batchIndicesTensor.AsPtr(); + desc.OutputGradientTensor = computeOutputGradient ? outputGradientTensor.AsPtr() : nullptr; + desc.OutputROIGradientTensor = computeOutputROIGradient ? outputROIGradientTensor.AsPtr() : nullptr; + desc.ReductionFunction = reductionFunction; + desc.InterpolationMode = interpolationMode; + desc.SpatialScaleX = spatialScaleX; + desc.SpatialScaleY = spatialScaleY; + desc.InputPixelOffset = inputPixelOffset; + desc.OutputPixelOffset = outputPixelOffset; + desc.MinimumSamplesPerOutput = minimumSamplesPerOutput; + desc.MaximumSamplesPerOutput = maximumSamplesPerOutput; + desc.AlignRegionsToCorners = alignRegionsToCorners; + + detail::NodeOutput* const inputs[] = { input ? input->Impl() : nullptr, inputGradient.Impl(), roi.Impl(), batchIndices.Impl() }; + detail::NodeID node = builder->CreateOperatorNode(static_cast(DML_OPERATOR_ROI_ALIGN_GRAD), &desc, inputs); + + RoiAlignGradOutputs outputs {}; + + if (computeOutputGradient) + { + outputs.outputGradient = builder->CreateNodeOutput(node, 0, std::move(outputGradientTensor)); + } + + if (computeOutputROIGradient) + { + outputs.outputROIGradient = builder->CreateNodeOutput(node, 1, std::move(outputROIGradientTensor)); + } + + return outputs; + } +#endif + + // Reinterprets the memory of a tensor with a different type and dimensions (analogously to using + // reinterpret_cast to access raw bits). Note that this is different to the DML Cast operator, which performs + // a type cast on the contents of a tensor (analogously to static_cast). The total tensor size of the output + // (which depends on the supplied type/sizes/strides) must match the input. + inline Expression Reinterpret( + Expression input, + DML_TENSOR_DATA_TYPE newType, + TensorDimensions newSizes, + Optional newStrides) + { + detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder(); + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + TensorDesc newTensor( + newType, + inputTensor.flags, + std::move(newSizes), + std::move(newStrides), + inputTensor.totalTensorSizeInBytes, + inputTensor.guaranteedBaseOffsetAlignment); + + detail::NodeID node = builder->CreateReinterpretNode(input.Impl()); + detail::NodeOutput* output = builder->CreateNodeOutput(node, 0, std::move(newTensor)); + + return output; + } + + // Same as Reinterpret above, but only adjusts tensor dimensions without affecting type. + inline Expression Reinterpret( + Expression input, + TensorDimensions newSizes, + Optional newStrides) + { + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + return Reinterpret(input, inputTensor.dataType, std::move(newSizes), std::move(newStrides)); + } + + // Same as Reinterpret above, but only adjusts tensor type without affecting sizes or strides. + inline Expression Reinterpret(Expression input, DML_TENSOR_DATA_TYPE newType) + { + TensorDesc inputTensor = input.Impl()->GetOutputDesc(); + + return Reinterpret(input, newType, inputTensor.sizes, inputTensor.strides); + } + + // Operator overloads for convenience, which merely map to one of the functions above + inline Expression operator+(Expression a, Expression b) { return dml::Add(a, b); } + inline Expression operator-(Expression a, Expression b) { return dml::Subtract(a, b); } + inline Expression operator*(Expression a, Expression b) { return dml::Multiply(a, b); } + inline Expression operator/(Expression a, Expression b) { return dml::Divide(a, b); } + inline Expression operator%(Expression a, Expression b) { return dml::ModulusTruncate(a, b); } + inline Expression operator&(Expression a, Expression b) { return dml::BitAnd(a, b); } + inline Expression operator|(Expression a, Expression b) { return dml::BitOr(a, b); } + inline Expression operator^(Expression a, Expression b) { return dml::BitXor(a, b); } + inline Expression operator<<(Expression a, Expression b) { return dml::BitShiftLeft(a, b); } + inline Expression operator>>(Expression a, Expression b) { return dml::BitShiftRight(a, b); } + inline Expression& operator+=(Expression& a, Expression b) { a = a + b; return a; } + inline Expression& operator-=(Expression& a, Expression b) { a = a - b; return a; } + inline Expression& operator*=(Expression& a, Expression b) { a = a * b; return a; } + inline Expression& operator/=(Expression& a, Expression b) { a = a / b; return a; } + inline Expression& operator%=(Expression& a, Expression b) { a = a % b; return a; } + inline Expression& operator&=(Expression& a, Expression b) { a = a & b; return a; } + inline Expression& operator|=(Expression& a, Expression b) { a = a | b; return a; } + inline Expression& operator^=(Expression& a, Expression b) { a = a ^ b; return a; } + inline Expression& operator<<=(Expression& a, Expression b) { a = a << b; return a; } + inline Expression& operator>>=(Expression& a, Expression b) { a = a >> b; return a; } + + // Operations involving scalars can be reduced to elementwise identity + inline Expression operator+(Expression a, float b) { return dml::Identity(a, DML_SCALE_BIAS{ 1.0f, b }); } + inline Expression operator-(Expression a, float b) { return dml::Identity(a, DML_SCALE_BIAS{ 1.0f, -b }); } + inline Expression operator*(Expression a, float b) { return dml::Identity(a, DML_SCALE_BIAS{ b, 0.0f }); } + inline Expression operator/(Expression a, float b) { return dml::Identity(a, DML_SCALE_BIAS{ 1.0f / b, 0.0f }); } + inline Expression operator+(float a, Expression b) { return dml::Identity(b, DML_SCALE_BIAS{ 1.0f, a }); } + inline Expression operator-(float a, Expression b) { return dml::Identity(b, DML_SCALE_BIAS{ -1.0f, a }); } + inline Expression operator*(float a, Expression b) { return dml::Identity(b, DML_SCALE_BIAS{ a, 0.0f }); } + inline Expression operator/(float a, Expression b) { return dml::Recip(b, DML_SCALE_BIAS{ a, 0.0f }); } + inline Expression& operator+=(Expression& a, float b) { a = a + b; return a; } + inline Expression& operator-=(Expression& a, float b) { a = a - b; return a; } + inline Expression& operator*=(Expression& a, float b) { a = a * b; return a; } + inline Expression& operator/=(Expression& a, float b) { a = a / b; return a; } + + // Unary + inline Expression operator~(Expression input) { return dml::BitNot(input); } + inline Expression operator+(Expression input) { return dml::Identity(input); } + inline Expression operator-(Expression input) { return dml::Identity(input, DML_SCALE_BIAS{ -1.0f, 0.0f }); } + + // Logical + inline Expression operator!(Expression a) { return dml::LogicalNot(a); } + inline Expression operator&&(Expression a, Expression b) { return dml::LogicalAnd(a, b); } + inline Expression operator||(Expression a, Expression b) { return dml::LogicalOr(a, b); } + inline Expression operator>(Expression a, Expression b) { return dml::GreaterThan(a, b); } + inline Expression operator<(Expression a, Expression b) { return dml::LessThan(a, b); } + inline Expression operator==(Expression a, Expression b) { return dml::Equals(a, b); } + inline Expression operator!=(Expression a, Expression b) { return !(a == b); } + inline Expression operator>=(Expression a, Expression b) { return dml::GreaterThanOrEqual(a, b); } + inline Expression operator<=(Expression a, Expression b) { return dml::LessThanOrEqual(a, b); } + + // GraphBuilder implementation details + namespace detail + { + inline NodeID GraphBuilder::CreateOperatorNode( + DML_OPERATOR_TYPE type, + const void* desc, + Span inputs) + { + DML_OPERATOR_DESC opDesc = { type, desc }; + + Microsoft::WRL::ComPtr op; + DMLX_THROW_IF_FAILED(m_device->CreateOperator(&opDesc, IID_PPV_ARGS(&op))); + + OperatorNode node = {}; + node.op = std::move(op); + node.inputs.assign(inputs.begin(), inputs.end()); + + uint32_t index = static_cast(m_operatorNodes.size()); + m_operatorNodes.push_back(std::move(node)); + + return { NodeType::Operator, index }; + } + + inline NodeID GraphBuilder::CreateInputNode(uint32_t inputIndex) + { + uint32_t index = static_cast(m_inputNodes.size()); + m_inputNodes.push_back(InputNode{ inputIndex }); + return { NodeType::Input, index }; + } + + inline NodeID GraphBuilder::CreateReinterpretNode(NodeOutput* input) + { + uint32_t index = static_cast(m_reinterpretNodes.size()); + m_reinterpretNodes.push_back(ReinterpretNode{ input }); + return { NodeType::Reinterpret, index }; + } + + inline NodeOutput* GraphBuilder::CreateNodeOutput(NodeID node, uint32_t outputIndex, TensorDesc tensorDesc) + { + // Construct the object in the deque, which doesn't invalidate references to elements as it grows + m_nodeOutputs.emplace_back(this, node, outputIndex, std::move(tensorDesc)); + + return &m_nodeOutputs.back(); + } + + inline GraphDesc GraphBuilder::GetGraphDesc(Span outputs) const + { + GraphDesc desc = {}; + desc.inputCount = static_cast(m_inputNodes.size()); + desc.outputCount = static_cast(outputs.size()); + + for (const OperatorNode& node : m_operatorNodes) + { + uint32_t nodeIndex = static_cast(desc.nodes.size()); + desc.nodes.push_back(DML_OPERATOR_GRAPH_NODE_DESC{ node.op.Get() }); + + // Walk through each of this node's inputs and add it as an edge + const uint32_t inputCount = static_cast(node.inputs.size()); + for (uint32_t inputIndex = 0; inputIndex < inputCount; ++inputIndex) + { + NodeOutput* input = node.inputs[inputIndex]; + if (input == nullptr) + { + continue; + } + NodeID inputNode = input->GetNode(); + + // Reinterpret nodes aren't "real" nodes, they're just used to modify TensorDescs across + // edges. So we follow this node backwards until it hits a real node. + while (inputNode.type == NodeType::Reinterpret) + { + input = m_reinterpretNodes[inputNode.index].input; + inputNode = input->GetNode(); + } + + if (inputNode.type == NodeType::Input) + { + DML_INPUT_GRAPH_EDGE_DESC inputEdge = {}; + inputEdge.GraphInputIndex = m_inputNodes[inputNode.index].inputIndex; + inputEdge.ToNodeIndex = nodeIndex; + inputEdge.ToNodeInputIndex = inputIndex; + + desc.inputEdges.push_back(inputEdge); + } + else if (inputNode.type == NodeType::Operator) + { + DML_INTERMEDIATE_GRAPH_EDGE_DESC intermediateEdge = {}; + intermediateEdge.FromNodeIndex = inputNode.index; + intermediateEdge.FromNodeOutputIndex = input->GetOutputIndex(); + intermediateEdge.ToNodeIndex = nodeIndex; + intermediateEdge.ToNodeInputIndex = inputIndex; + + desc.intermediateEdges.push_back(intermediateEdge); + } + else + { + assert(false); // Invalid node type + DMLX_THROW(E_UNEXPECTED); + } + } + } + + // Add output edges + for (uint32_t outputIndex = 0; outputIndex < desc.outputCount; ++outputIndex) + { + NodeOutput* output = outputs[outputIndex].Impl(); + if (output == nullptr) + { + continue; + } + NodeID outputNode = output->GetNode(); + + // Reinterpret nodes are meaningless on outputs (they're no-ops), so just follow them back until we + // get to a real operator node. + while (outputNode.type == NodeType::Reinterpret) + { + output = m_reinterpretNodes[outputNode.index].input; + outputNode = output->GetNode(); + } + + if (outputNode.type == NodeType::Input) + { + // It's not valid to connect an output of the graph directly to an input without an intervening + // node. If this behavior is desired, it should instead be accomplished with a copy e.g. using + // the elementwise identity operator. + DMLX_THROW(E_INVALIDARG); + } + + assert(outputNode.type == NodeType::Operator); + + DML_OUTPUT_GRAPH_EDGE_DESC outputEdge = {}; + outputEdge.FromNodeIndex = output->GetNode().index; + outputEdge.FromNodeOutputIndex = output->GetOutputIndex(); + outputEdge.GraphOutputIndex = outputIndex; + + desc.outputEdges.push_back(outputEdge); + } + + // Sanity + assert(desc.nodes.size() == m_operatorNodes.size()); + assert(desc.outputEdges.size() == desc.outputCount); + assert(desc.outputCount == outputs.size()); + + return desc; + } + } // namespace detail + +} // namespace dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphDescBuilder.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphDescBuilder.cpp index 0f65c11945..ef0ade7a47 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphDescBuilder.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphDescBuilder.cpp @@ -140,8 +140,9 @@ namespace Dml::GraphDescBuilder node, constantCpuNodeInputGetter, executionHandle, - &graphNodeInfo + /*out*/ &graphNodeInfo ); + ORT_THROW_HR_IF(E_UNEXPECTED, !graphNodeInfo.desc); uint32_t nodeIndex = gsl::narrow_cast(graphNodes.size()); AbstractOperatorDesc opDesc = *graphNodeInfo.desc; // Make a copy diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp index ad52d4056c..ed30d9540e 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/GraphPartitioner.cpp @@ -532,7 +532,7 @@ namespace Dml partitionNodePropsMap.insert(std::make_pair( GraphDescBuilder::GetUniqueNodeName(*node), std::move(graphNodePropertyMap[node]))); } - + #ifdef PRINT_PARTITON_INFO printf("\n"); #endif @@ -540,7 +540,7 @@ namespace Dml auto fused_kernel_func = [partitionNodePropsMap, transferredInitializerMap](onnxruntime::FuncManager& func_mgr, const onnxruntime::OpKernelInfo& info, std::unique_ptr& out) mutable ->onnxruntime::Status { out.reset(CreateFusedGraphKernel(info, partitionNodePropsMap, *transferredInitializerMap)); - return Status::OK(); + return Status::OK(); }; // build the kernel definition on the fly, and register it to the fused_kernel_regisitry. diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp index 6c1b670502..cfa5ba8a2b 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.cpp @@ -18,2034 +18,2283 @@ namespace Windows::AI::MachineLearning::Adapter #pragma warning(push) #pragma warning(disable:4702) -size_t AttributeValue::ElementCount() const { - switch (type) { - case MLOperatorAttributeType::Float: - ML_CHECK_BOOL(floats.size() == 1); - return 1; + size_t AttributeValue::ElementCount() const + { + switch (type) + { + case MLOperatorAttributeType::Float: + ML_CHECK_BOOL(floats.size() == 1); + return 1; - case MLOperatorAttributeType::Int: - ML_CHECK_BOOL(ints.size() == 1); - return 1; + case MLOperatorAttributeType::Int: + ML_CHECK_BOOL(ints.size() == 1); + return 1; - case MLOperatorAttributeType::String: - ML_CHECK_BOOL(strings.size() == 1); - return 1; + case MLOperatorAttributeType::String: + ML_CHECK_BOOL(strings.size() == 1); + return 1; - case MLOperatorAttributeType::FloatArray: - return floats.size(); + case MLOperatorAttributeType::FloatArray: + return floats.size(); - case MLOperatorAttributeType::IntArray: - return ints.size(); + case MLOperatorAttributeType::IntArray: + return ints.size(); - case MLOperatorAttributeType::StringArray: - return strings.size(); + case MLOperatorAttributeType::StringArray: + return strings.size(); - default: - // The type is validated when default attributes are registered - assert(false); - ORT_THROW_HR(E_FAIL); - return 0; - } - #pragma warning(pop) -} + default: + // The type is validated when default attributes are registered + assert(false); + ORT_THROW_HR(E_FAIL); + return 0; + } +#pragma warning(pop) + } -void AttributeValue::GetAttribute( - MLOperatorAttributeType attributeType, - uint32_t elementCount, - size_t elementByteSize, - void* value) const { - switch (attributeType) { - case MLOperatorAttributeType::Float: - ML_CHECK_BOOL(floats.size() == 1); - __fallthrough; - case MLOperatorAttributeType::FloatArray: - ML_CHECK_BOOL(floats.size() == elementCount); - ML_CHECK_BOOL(elementByteSize == sizeof(float)); - std::copy(floats.begin(), floats.end(), static_cast(value)); - break; + void AttributeValue::GetAttribute( + MLOperatorAttributeType attributeType, + uint32_t elementCount, + size_t elementByteSize, + void* value) const + { + switch (attributeType) + { + case MLOperatorAttributeType::Float: + ML_CHECK_BOOL(floats.size() == 1); + __fallthrough; + case MLOperatorAttributeType::FloatArray: + ML_CHECK_BOOL(floats.size() == elementCount); + ML_CHECK_BOOL(elementByteSize == sizeof(float)); + std::copy(floats.begin(), floats.end(), static_cast(value)); + break; - case MLOperatorAttributeType::Int: - ML_CHECK_BOOL(ints.size() == 1); - __fallthrough; - case MLOperatorAttributeType::IntArray: - ML_CHECK_BOOL(ints.size() == elementCount); - ML_CHECK_BOOL(elementByteSize == sizeof(int64_t)); - std::copy(ints.begin(), ints.end(), static_cast(value)); - break; + case MLOperatorAttributeType::Int: + ML_CHECK_BOOL(ints.size() == 1); + __fallthrough; + case MLOperatorAttributeType::IntArray: + ML_CHECK_BOOL(ints.size() == elementCount); + ML_CHECK_BOOL(elementByteSize == sizeof(int64_t)); + std::copy(ints.begin(), ints.end(), static_cast(value)); + break; - default: - ORT_THROW_HR(E_INVALIDARG); - } -} + default: + ORT_THROW_HR(E_INVALIDARG); + } + } -const std::string* AttributeValue::GetStringAttribute( - _In_z_ const char* attributeName, - uint32_t elementIndex) const { - ML_CHECK_BOOL((type == MLOperatorAttributeType::String && elementIndex == 0 && strings.size() == 1) || - (type == MLOperatorAttributeType::StringArray && elementIndex < strings.size())); + const std::string* AttributeValue::GetStringAttribute( + _In_z_ const char* attributeName, + uint32_t elementIndex) const + { + ML_CHECK_BOOL((type == MLOperatorAttributeType::String && elementIndex == 0 && strings.size() == 1) || + (type == MLOperatorAttributeType::StringArray && elementIndex < strings.size())); - return &strings.data()[elementIndex]; -} + return &strings.data()[elementIndex]; + } -bool IsAllocationInterface(const ::OrtMemoryInfo& info) { - return strcmp(info.name, onnxruntime::CPU) && !(info.mem_type == ::OrtMemType::OrtMemTypeCPUOutput || info.mem_type == ::OrtMemType::OrtMemTypeCPUInput); -} + bool IsAllocationInterface(const ::OrtMemoryInfo& info) + { + return strcmp(info.name, onnxruntime::CPU) && !(info.mem_type == ::OrtMemType::OrtMemTypeCPUOutput || info.mem_type == ::OrtMemType::OrtMemTypeCPUInput); + } -// Translate the data object stored in a tensor to the type which will be returned through -// the ABI. The translation is determined by the provider and based on options with which the -// kernels are registered. -void TranslateAllocationDataToAbi( - IWinmlExecutionProvider* winmlProvider, - bool isInternalOperator, - const ::OrtMemoryInfo& allocInfo, - IUnknown* allocation, - IUnknown** abiAllocation) { - if (winmlProvider) { - winmlProvider->GetABIDataInterface(isInternalOperator, allocation, abiAllocation); - } else { - ComPtr tmp = allocation; - *abiAllocation = tmp.Detach(); - } -} + // Translate the data object stored in a tensor to the type which will be returned through + // the ABI. The translation is determined by the provider and based on options with which the + // kernels are registered. + void TranslateAllocationDataToAbi( + IWinmlExecutionProvider* winmlProvider, + bool isInternalOperator, + const ::OrtMemoryInfo& allocInfo, + IUnknown* allocation, + IUnknown** abiAllocation) + { + if (winmlProvider) + { + winmlProvider->GetABIDataInterface(isInternalOperator, allocation, abiAllocation); + } + else + { + ComPtr tmp = allocation; + *abiAllocation = tmp.Detach(); + } + } -// -// Traits for numeric attribute types -// -template -struct MLAttributeTypeTraits { -}; + // + // Traits for numeric attribute types + // + template + struct MLAttributeTypeTraits { + }; -template <> -struct MLAttributeTypeTraits { - using Type = float; - static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_FLOAT; - static const bool IsPrimitiveAttributeType = true; - static const bool IsArray = false; -}; + template <> + struct MLAttributeTypeTraits { + using Type = float; + static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_FLOAT; + static const bool IsPrimitiveAttributeType = true; + static const bool IsArray = false; + }; -template <> -struct MLAttributeTypeTraits { - using Type = int64_t; - static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_INT; - static const bool IsPrimitiveAttributeType = true; - static const bool IsArray = false; -}; + template <> + struct MLAttributeTypeTraits { + using Type = int64_t; + static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_INT; + static const bool IsPrimitiveAttributeType = true; + static const bool IsArray = false; + }; -template <> -struct MLAttributeTypeTraits { - static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_STRING; - static const bool IsPrimitiveAttributeType = true; - static const bool IsArray = false; -}; + template <> + struct MLAttributeTypeTraits { + static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_STRING; + static const bool IsPrimitiveAttributeType = true; + static const bool IsArray = false; + }; -template <> -struct MLAttributeTypeTraits { - static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_TENSOR; - static const bool IsPrimitiveAttributeType = false; - static const bool IsArray = false; -}; + template <> + struct MLAttributeTypeTraits { + static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_TENSOR; + static const bool IsPrimitiveAttributeType = false; + static const bool IsArray = false; + }; -template <> -struct MLAttributeTypeTraits { - using Type = float; - static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_FLOATS; - static const bool IsPrimitiveAttributeType = true; - static const bool IsArray = true; -}; + template <> + struct MLAttributeTypeTraits { + using Type = float; + static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_FLOATS; + static const bool IsPrimitiveAttributeType = true; + static const bool IsArray = true; + }; -template <> -struct MLAttributeTypeTraits { - using Type = int64_t; - static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_INTS; - static const bool IsPrimitiveAttributeType = true; - static const bool IsArray = true; -}; + template <> + struct MLAttributeTypeTraits { + using Type = int64_t; + static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_INTS; + static const bool IsPrimitiveAttributeType = true; + static const bool IsArray = true; + }; -template <> -struct MLAttributeTypeTraits { - static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_STRINGS; - static const bool IsPrimitiveAttributeType = true; - static const bool IsArray = true; -}; + template <> + struct MLAttributeTypeTraits { + static const onnx::AttributeProto_AttributeType ProtoType = onnx::AttributeProto_AttributeType_STRINGS; + static const bool IsPrimitiveAttributeType = true; + static const bool IsArray = true; + }; #define ML_ATTR_TO_PROTO_CASE(x) \ case MLOperatorAttributeType::x: \ return MLAttributeTypeTraits::ProtoType; -onnx::AttributeProto_AttributeType ToProto(MLOperatorAttributeType type) { - switch (type) { - case MLOperatorAttributeType::Float: - return MLAttributeTypeTraits::ProtoType; - case MLOperatorAttributeType::Int: - return MLAttributeTypeTraits::ProtoType; - case MLOperatorAttributeType::FloatArray: - return MLAttributeTypeTraits::ProtoType; - #pragma warning(suppress:4063) - case MLOperatorAttributeTypeTensor: - return MLAttributeTypeTraits::ProtoType; - case MLOperatorAttributeType::IntArray: - return MLAttributeTypeTraits::ProtoType; - case MLOperatorAttributeType::String: - return MLAttributeTypeTraits::ProtoType; - case MLOperatorAttributeType::StringArray: - return MLAttributeTypeTraits::ProtoType; - default: - return onnx::AttributeProto_AttributeType_UNDEFINED; - } -} + onnx::AttributeProto_AttributeType ToProto(MLOperatorAttributeType type) + { + switch (type) + { + case MLOperatorAttributeType::Float: + return MLAttributeTypeTraits::ProtoType; + case MLOperatorAttributeType::Int: + return MLAttributeTypeTraits::ProtoType; + case MLOperatorAttributeType::FloatArray: + return MLAttributeTypeTraits::ProtoType; +#pragma warning(suppress:4063) + case MLOperatorAttributeTypeTensor: + return MLAttributeTypeTraits::ProtoType; + case MLOperatorAttributeType::IntArray: + return MLAttributeTypeTraits::ProtoType; + case MLOperatorAttributeType::String: + return MLAttributeTypeTraits::ProtoType; + case MLOperatorAttributeType::StringArray: + return MLAttributeTypeTraits::ProtoType; + default: + return onnx::AttributeProto_AttributeType_UNDEFINED; + } + } -bool IsPrimitiveAttributeType(MLOperatorAttributeType type) { - switch (type) { - case MLOperatorAttributeType::Float: - return MLAttributeTypeTraits::IsPrimitiveAttributeType; - case MLOperatorAttributeType::Int: - return MLAttributeTypeTraits::IsPrimitiveAttributeType; - case MLOperatorAttributeType::FloatArray: - return MLAttributeTypeTraits::IsPrimitiveAttributeType; - case MLOperatorAttributeType::IntArray: - return MLAttributeTypeTraits::IsPrimitiveAttributeType; - case MLOperatorAttributeType::String: - return MLAttributeTypeTraits::IsPrimitiveAttributeType; - case MLOperatorAttributeType::StringArray: - return MLAttributeTypeTraits::IsPrimitiveAttributeType; - default: - return false; // Including other types like tensor and graph... - } -} + bool IsPrimitiveAttributeType(MLOperatorAttributeType type) + { + switch (type) + { + case MLOperatorAttributeType::Float: + return MLAttributeTypeTraits::IsPrimitiveAttributeType; + case MLOperatorAttributeType::Int: + return MLAttributeTypeTraits::IsPrimitiveAttributeType; + case MLOperatorAttributeType::FloatArray: + return MLAttributeTypeTraits::IsPrimitiveAttributeType; + case MLOperatorAttributeType::IntArray: + return MLAttributeTypeTraits::IsPrimitiveAttributeType; + case MLOperatorAttributeType::String: + return MLAttributeTypeTraits::IsPrimitiveAttributeType; + case MLOperatorAttributeType::StringArray: + return MLAttributeTypeTraits::IsPrimitiveAttributeType; + default: + return false; // Including other types like tensor and graph... + } + } #define ML_TENSOR_TYPE_CASE(x) \ - if (onnxruntime::utils::IsPrimitiveDataType(type)) { \ + if (onnxruntime::utils::IsPrimitiveDataType(type)) \ + { \ return MLTypeTraits::TensorType; \ } #pragma warning(push) #pragma warning(disable:4702) -::MLOperatorTensorDataType ToMLTensorDataType(onnxruntime::MLDataType type) { - if (onnxruntime::utils::IsDataTypeString(type)) { - return MLOperatorTensorDataType::String; - } + ::MLOperatorTensorDataType ToMLTensorDataType(onnxruntime::MLDataType type) + { + if (onnxruntime::utils::IsDataTypeString(type)) + { + return MLOperatorTensorDataType::String; + } - ML_TENSOR_TYPE_CASE(float); - ML_TENSOR_TYPE_CASE(uint8_t); - ML_TENSOR_TYPE_CASE(int8_t); - ML_TENSOR_TYPE_CASE(uint16_t); - ML_TENSOR_TYPE_CASE(int16_t); - ML_TENSOR_TYPE_CASE(int32_t); - ML_TENSOR_TYPE_CASE(int64_t); - ML_TENSOR_TYPE_CASE(bool); - ML_TENSOR_TYPE_CASE(double); - ML_TENSOR_TYPE_CASE(uint32_t); - ML_TENSOR_TYPE_CASE(uint64_t); - ML_TENSOR_TYPE_CASE(onnxruntime::MLFloat16); + ML_TENSOR_TYPE_CASE(float); + ML_TENSOR_TYPE_CASE(uint8_t); + ML_TENSOR_TYPE_CASE(int8_t); + ML_TENSOR_TYPE_CASE(uint16_t); + ML_TENSOR_TYPE_CASE(int16_t); + ML_TENSOR_TYPE_CASE(int32_t); + ML_TENSOR_TYPE_CASE(int64_t); + ML_TENSOR_TYPE_CASE(bool); + ML_TENSOR_TYPE_CASE(double); + ML_TENSOR_TYPE_CASE(uint32_t); + ML_TENSOR_TYPE_CASE(uint64_t); + ML_TENSOR_TYPE_CASE(onnxruntime::MLFloat16); - ORT_THROW_HR(E_NOTIMPL); - return MLOperatorTensorDataType::Undefined; - #pragma warning(pop) -} + ORT_THROW_HR(E_NOTIMPL); + return MLOperatorTensorDataType::Undefined; +#pragma warning(pop) + } #undef ML_TENSOR_TYPE_CASE #define ML_TENSOR_TYPE_CASE(x) \ - if (type == MLTypeTraits::TensorType) { \ + if (type == MLTypeTraits::TensorType) \ + { \ return onnxruntime::DataTypeImpl::GetTensorType(); \ } #pragma warning(push) #pragma warning(disable:4702) -onnxruntime::MLDataType ToTensorDataType(::MLOperatorTensorDataType type) { - if (type == MLOperatorTensorDataType::String) - return onnxruntime::DataTypeImpl::GetTensorType(); + onnxruntime::MLDataType ToTensorDataType(::MLOperatorTensorDataType type) + { + if (type == MLOperatorTensorDataType::String) + return onnxruntime::DataTypeImpl::GetTensorType(); - ML_TENSOR_TYPE_CASE(float); - ML_TENSOR_TYPE_CASE(uint8_t); - ML_TENSOR_TYPE_CASE(int8_t); - ML_TENSOR_TYPE_CASE(uint16_t); - ML_TENSOR_TYPE_CASE(int16_t); - ML_TENSOR_TYPE_CASE(int32_t); - ML_TENSOR_TYPE_CASE(int64_t); - ML_TENSOR_TYPE_CASE(bool); - ML_TENSOR_TYPE_CASE(double); - ML_TENSOR_TYPE_CASE(uint32_t); - ML_TENSOR_TYPE_CASE(uint64_t); - ML_TENSOR_TYPE_CASE(onnxruntime::MLFloat16); + ML_TENSOR_TYPE_CASE(float); + ML_TENSOR_TYPE_CASE(uint8_t); + ML_TENSOR_TYPE_CASE(int8_t); + ML_TENSOR_TYPE_CASE(uint16_t); + ML_TENSOR_TYPE_CASE(int16_t); + ML_TENSOR_TYPE_CASE(int32_t); + ML_TENSOR_TYPE_CASE(int64_t); + ML_TENSOR_TYPE_CASE(bool); + ML_TENSOR_TYPE_CASE(double); + ML_TENSOR_TYPE_CASE(uint32_t); + ML_TENSOR_TYPE_CASE(uint64_t); + ML_TENSOR_TYPE_CASE(onnxruntime::MLFloat16); - ORT_THROW_HR(E_NOTIMPL); - return onnxruntime::DataTypeImpl::GetTensorType(); + ORT_THROW_HR(E_NOTIMPL); + return onnxruntime::DataTypeImpl::GetTensorType(); #pragma warning(pop) -} + } #pragma warning(push) #pragma warning(disable:4702) -::MLOperatorTensorDataType ToMLTensorDataType(onnx::TensorProto_DataType type) { - switch (type) { - case onnx::TensorProto_DataType_FLOAT: - return MLOperatorTensorDataType::Float; + ::MLOperatorTensorDataType ToMLTensorDataType(onnx::TensorProto_DataType type) + { + switch (type) + { + case onnx::TensorProto_DataType_FLOAT: + return MLOperatorTensorDataType::Float; - case onnx::TensorProto_DataType_UINT8: - return MLOperatorTensorDataType::UInt8; + case onnx::TensorProto_DataType_UINT8: + return MLOperatorTensorDataType::UInt8; - case onnx::TensorProto_DataType_INT8: - return MLOperatorTensorDataType::Int8; + case onnx::TensorProto_DataType_INT8: + return MLOperatorTensorDataType::Int8; - case onnx::TensorProto_DataType_UINT16: - return MLOperatorTensorDataType::UInt16; + case onnx::TensorProto_DataType_UINT16: + return MLOperatorTensorDataType::UInt16; - case onnx::TensorProto_DataType_INT16: - return MLOperatorTensorDataType::Int16; + case onnx::TensorProto_DataType_INT16: + return MLOperatorTensorDataType::Int16; - case onnx::TensorProto_DataType_INT32: - return MLOperatorTensorDataType::Int32; + case onnx::TensorProto_DataType_INT32: + return MLOperatorTensorDataType::Int32; - case onnx::TensorProto_DataType_INT64: - return MLOperatorTensorDataType::Int64; + case onnx::TensorProto_DataType_INT64: + return MLOperatorTensorDataType::Int64; - case onnx::TensorProto_DataType_STRING: - return MLOperatorTensorDataType::String; + case onnx::TensorProto_DataType_STRING: + return MLOperatorTensorDataType::String; - case onnx::TensorProto_DataType_BOOL: - return MLOperatorTensorDataType::Bool; + case onnx::TensorProto_DataType_BOOL: + return MLOperatorTensorDataType::Bool; - case onnx::TensorProto_DataType_FLOAT16: - return MLOperatorTensorDataType::Float16; + case onnx::TensorProto_DataType_FLOAT16: + return MLOperatorTensorDataType::Float16; - case onnx::TensorProto_DataType_DOUBLE: - return MLOperatorTensorDataType::Double; + case onnx::TensorProto_DataType_DOUBLE: + return MLOperatorTensorDataType::Double; - case onnx::TensorProto_DataType_UINT32: - return MLOperatorTensorDataType::UInt32; + case onnx::TensorProto_DataType_UINT32: + return MLOperatorTensorDataType::UInt32; - case onnx::TensorProto_DataType_UINT64: - return MLOperatorTensorDataType::UInt64; + case onnx::TensorProto_DataType_UINT64: + return MLOperatorTensorDataType::UInt64; - case onnx::TensorProto_DataType_COMPLEX64: - return MLOperatorTensorDataType::Complex64; + case onnx::TensorProto_DataType_COMPLEX64: + return MLOperatorTensorDataType::Complex64; - case onnx::TensorProto_DataType_COMPLEX128: - return MLOperatorTensorDataType::Complex128; + case onnx::TensorProto_DataType_COMPLEX128: + return MLOperatorTensorDataType::Complex128; - default: - ORT_THROW_HR(E_NOTIMPL); - return MLOperatorTensorDataType::Undefined; - } + default: + ORT_THROW_HR(E_NOTIMPL); + return MLOperatorTensorDataType::Undefined; + } #pragma warning(pop) -} - -::MLOperatorEdgeDescription ToMLEdgeDesc(const onnx::TypeProto* type) { - // Initialized to undefined class and data type - MLOperatorEdgeDescription ret = {}; - - ML_CHECK_BOOL(type->value_case() == onnx::TypeProto::kTensorType || - type->value_case() == onnx::TypeProto::VALUE_NOT_SET); - - if (type->value_case() == onnx::TypeProto::kTensorType) { - ret.edgeType = MLOperatorEdgeType::Tensor; - const onnx::TypeProto_Tensor tensorType = type->tensor_type(); - if (tensorType.has_elem_type()) { - ret.tensorDataType = ToMLTensorDataType(onnx::TensorProto_DataType(tensorType.elem_type())); } - } - return ret; -} + ::MLOperatorEdgeDescription ToMLEdgeDesc(const onnx::TypeProto* type) + { + // Initialized to undefined class and data type + MLOperatorEdgeDescription ret = {}; + + ML_CHECK_BOOL(type->value_case() == onnx::TypeProto::kTensorType || + type->value_case() == onnx::TypeProto::VALUE_NOT_SET); + + if (type->value_case() == onnx::TypeProto::kTensorType) + { + ret.edgeType = MLOperatorEdgeType::Tensor; + const onnx::TypeProto_Tensor tensorType = type->tensor_type(); + if (tensorType.has_elem_type()) + { + ret.tensorDataType = ToMLTensorDataType(onnx::TensorProto_DataType(tensorType.elem_type())); + } + } + + return ret; + } #pragma warning(push) #pragma warning(disable:4702) -std::string ToTypeString(MLOperatorEdgeDescription desc) { - if (desc.edgeType != MLOperatorEdgeType::Tensor) { - ORT_THROW_HR(E_NOTIMPL); - } + std::string ToTypeString(MLOperatorEdgeDescription desc) + { + if (desc.edgeType != MLOperatorEdgeType::Tensor) + { + ORT_THROW_HR(E_NOTIMPL); + } - switch (desc.tensorDataType) { - case MLOperatorTensorDataType::Float: - return "tensor(float)"; + switch (desc.tensorDataType) + { + case MLOperatorTensorDataType::Float: + return "tensor(float)"; - case MLOperatorTensorDataType::UInt8: - return "tensor(uint8)"; + case MLOperatorTensorDataType::UInt8: + return "tensor(uint8)"; - case MLOperatorTensorDataType::Int8: - return "tensor(int8)"; + case MLOperatorTensorDataType::Int8: + return "tensor(int8)"; - case MLOperatorTensorDataType::UInt16: - return "tensor(uint16)"; + case MLOperatorTensorDataType::UInt16: + return "tensor(uint16)"; - case MLOperatorTensorDataType::Int16: - return "tensor(int16)"; + case MLOperatorTensorDataType::Int16: + return "tensor(int16)"; - case MLOperatorTensorDataType::Int32: - return "tensor(int32)"; + case MLOperatorTensorDataType::Int32: + return "tensor(int32)"; - case MLOperatorTensorDataType::Int64: - return "tensor(int64)"; + case MLOperatorTensorDataType::Int64: + return "tensor(int64)"; - case MLOperatorTensorDataType::String: - return "tensor(string)"; + case MLOperatorTensorDataType::String: + return "tensor(string)"; - case MLOperatorTensorDataType::Bool: - return "tensor(bool)"; + case MLOperatorTensorDataType::Bool: + return "tensor(bool)"; - case MLOperatorTensorDataType::Float16: - return "tensor(float16)"; + case MLOperatorTensorDataType::Float16: + return "tensor(float16)"; - case MLOperatorTensorDataType::Double: - return "tensor(double)"; + case MLOperatorTensorDataType::Double: + return "tensor(double)"; - case MLOperatorTensorDataType::UInt32: - return "tensor(uint32)"; + case MLOperatorTensorDataType::UInt32: + return "tensor(uint32)"; - case MLOperatorTensorDataType::UInt64: - return "tensor(uint64)"; + case MLOperatorTensorDataType::UInt64: + return "tensor(uint64)"; - case MLOperatorTensorDataType::Complex64: - return "tensor(complext64)"; + case MLOperatorTensorDataType::Complex64: + return "tensor(complext64)"; - case MLOperatorTensorDataType::Complex128: - return "tensor(complext128)"; + case MLOperatorTensorDataType::Complex128: + return "tensor(complext128)"; - default: - ORT_THROW_HR(E_NOTIMPL); - return ""; - } + default: + ORT_THROW_HR(E_NOTIMPL); + return ""; + } #pragma warning(pop) -} - -OpKernelInfoWrapper::OpKernelInfoWrapper( - const onnxruntime::OpKernelInfo* kerneInfo, - IUnknown* abiExecutionObject, - const EdgeShapes* inputShapeOverrides, - const EdgeShapes* inferredOutputShapes, - bool allowInputShapeQuery, - bool allowOutputShapeQuery, - bool isInternalOperator, - const AttributeMap* defaultAttributes, - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter) : OpNodeInfoWrapper(kerneInfo, inputShapeOverrides, defaultAttributes, requiredConstantCpuInputs, constantInputGetter), - m_inferredOutputShapes(inferredOutputShapes), - m_allowInputShapeQuery(allowInputShapeQuery), - m_allowOutputShapeQuery(allowOutputShapeQuery), - m_internalOperator(isInternalOperator), - m_impl(kerneInfo), - m_abiExecutionObject(abiExecutionObject) { - const void* executionHandle = kerneInfo->GetExecutionProvider()->GetExecutionHandle(); - if (executionHandle) { - // We assume the execution object inherits IUnknown as its first base - ComPtr providerExecutionObject = const_cast(static_cast(executionHandle)); - providerExecutionObject.As(&m_winmlProvider); - } - - assert(allowInputShapeQuery || !allowOutputShapeQuery); - - // The input may be exposed using non-overridden sizes. Exposing output shapes requires - // those shapes be provided here. - assert(!allowOutputShapeQuery || (inferredOutputShapes != nullptr)); -} - -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetAttributeElementCount( - _In_z_ const char* name, - MLOperatorAttributeType type, - uint32_t* elementCount) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *elementCount = 0; - - if (IsPrimitiveAttributeType(type)) { - *elementCount = m_impl->GetPrimitiveAttrElementCount(ToProto(type), std::string(name)); - } else { - // ONNX runtime does not implement OpNodeProtoHelper::GetPrimitiveAttrElementCount for tensors. - // So we need to test presence a different way. - - const onnx::AttributeProto* attributeProto = m_impl->TryGetAttribute(std::string(name)); - *elementCount = attributeProto ? 1 : 0; - } - - // Look for a value in the kernel's registered defaults if one was not found - if (*elementCount == 0 && m_defaultAttributes) { - auto defaultAttr = m_defaultAttributes->find(name); - if (defaultAttr != m_defaultAttributes->end()) { - *elementCount = static_cast(defaultAttr->second.ElementCount()); - } - } - - return S_OK; } - ORT_CATCH_RETURN -} -template -template -HRESULT OpNodeInfoWrapper::GetAttributeArrayHelper( - _In_z_ const char* name, - uint32_t elementCount, - uint32_t elementByteSize, - void* values) const { - using elementType_t = typename MLAttributeTypeTraits::Type; - static_assert(MLAttributeTypeTraits::IsArray, "This function only works with array types."); - ML_CHECK_BOOL(sizeof(elementType_t) == elementByteSize); - - THROW_IF_NOT_OK(m_impl->GetAttrs(name, gsl::span(static_cast::Type*>(values), elementCount))); - return S_OK; -} - -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetAttribute( - _In_z_ const char* name, - MLOperatorAttributeType type, - uint32_t elementCount, - size_t elementByteSize, - /*out*/void* attributeValue) const noexcept -{ - ORT_TRY + OpKernelInfoWrapper::OpKernelInfoWrapper( + const onnxruntime::OpKernelInfo* kerneInfo, + IUnknown* abiExecutionObject, + const EdgeShapes* inputShapeOverrides, + const EdgeShapes* inferredOutputShapes, + bool allowInputShapeQuery, + bool allowOutputShapeQuery, + bool isInternalOperator, + const AttributeMap* defaultAttributes, + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& constantInputGetter + ) + : OpNodeInfoWrapper(kerneInfo, inputShapeOverrides, defaultAttributes, requiredConstantCpuInputs, constantInputGetter), + m_inferredOutputShapes(inferredOutputShapes), + m_allowInputShapeQuery(allowInputShapeQuery), + m_allowOutputShapeQuery(allowOutputShapeQuery), + m_internalOperator(isInternalOperator), + m_impl(kerneInfo), + m_abiExecutionObject(abiExecutionObject) { - VerifyNotClosed(); - - // Look for a value in the kernel's registered defaults if one does not exist otherwise - if (m_impl->GetPrimitiveAttrElementCount(ToProto(type), name) == 0) { - if (!m_defaultAttributes) { - ORT_THROW_HR(E_FAIL); + const void* executionHandle = kerneInfo->GetExecutionProvider()->GetExecutionHandle(); + if (executionHandle) + { + // We assume the execution object inherits IUnknown as its first base + ComPtr providerExecutionObject = const_cast(static_cast(executionHandle)); + providerExecutionObject.As(&m_winmlProvider); } - auto defaultAttr = m_defaultAttributes->find(name); - if (defaultAttr == m_defaultAttributes->end()) { - ORT_THROW_HR(E_FAIL); - } + assert(allowInputShapeQuery || !allowOutputShapeQuery); - defaultAttr->second.GetAttribute(type, elementCount, elementByteSize, /*out*/attributeValue); - } else { - switch (type) { - case MLOperatorAttributeType::Float: - ML_CHECK_BOOL(elementCount == 1); - return GetAttributeHelper(name, static_cast(elementByteSize), /*out*/attributeValue); - - case MLOperatorAttributeType::Int: - ML_CHECK_BOOL(elementCount == 1); - return GetAttributeHelper(name, static_cast(elementByteSize), /*out*/attributeValue); - - case MLOperatorAttributeType::FloatArray: - return GetAttributeArrayHelper(name, elementCount, static_cast(elementByteSize), /*out*/attributeValue); - - case MLOperatorAttributeType::IntArray: - return GetAttributeArrayHelper(name, elementCount, static_cast(elementByteSize), /*out*/attributeValue); - - default: - ML_CHECK_BOOL(false); - break; - } - } - - return S_OK; - } - ORT_CATCH_RETURN -} - -template -const std::string* OpNodeInfoWrapper::GetStringAttribute( - _In_z_ const char* name, - uint32_t elementIndex) const { - // Get the proto attribute - const onnx::AttributeProto* attr = m_impl->TryGetAttribute(std::string(name)); - - // Look for a value in the kernel's registered defaults if one was not found - if (!attr) { - if (!m_defaultAttributes) { - ORT_THROW_HR(E_FAIL); + // The input may be exposed using non-overridden sizes. Exposing output shapes requires + // those shapes be provided here. + assert(!allowOutputShapeQuery || (inferredOutputShapes != nullptr)); } - auto defaultAttr = m_defaultAttributes->find(name); - if (defaultAttr == m_defaultAttributes->end()) { - ORT_THROW_HR(E_FAIL); - } - - return defaultAttr->second.GetStringAttribute(name, elementIndex); - } else { - // Get the string vector from the attribute - if (attr->has_s()) { - return &attr->s(); - } else { - // Check the size of the vector - ML_CHECK_BOOL(attr->strings_size() > 0); - ML_CHECK_BOOL(elementIndex < static_cast(attr->strings_size())); - - return &attr->strings(elementIndex); - } - } -} - -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetStringAttributeElementLength( - _In_z_ const char* name, - uint32_t elementIndex, - uint32_t* attributeElementByteLength) const noexcept -{ - ORT_TRY + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetAttributeElementCount( + _In_z_ const char* name, + MLOperatorAttributeType type, + uint32_t* elementCount) const noexcept { - VerifyNotClosed(); + ORT_TRY + { + VerifyNotClosed(); - *attributeElementByteLength = 0; - const std::string* protoString = GetStringAttribute(name, elementIndex); + *elementCount = 0; - // Check for overflow and casting safety - ML_CHECK_BOOL(protoString->size() < protoString->size() + 1); - ML_CHECK_BOOL(protoString->size() + 1 <= std::numeric_limits::max()); + if (IsPrimitiveAttributeType(type)) + { + *elementCount = m_impl->GetPrimitiveAttrElementCount(ToProto(type), std::string(name)); + } + else + { + // ONNX runtime does not implement OpNodeProtoHelper::GetPrimitiveAttrElementCount for tensors. + // So we need to test presence a different way. - // Set the length including null termination - *attributeElementByteLength = static_cast(protoString->size() + 1); - return S_OK; - } - ORT_CATCH_RETURN -} + const onnx::AttributeProto* attributeProto = m_impl->TryGetAttribute(std::string(name)); + *elementCount = attributeProto ? 1 : 0; + } -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetStringAttributeElement( - _In_z_ const char* name, - uint32_t elementIndex, - uint32_t attributeElementByteLength, - char* attributeElement) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); + // Look for a value in the kernel's registered defaults if one was not found + if (*elementCount == 0 && m_defaultAttributes) + { + auto defaultAttr = m_defaultAttributes->find(name); + if (defaultAttr != m_defaultAttributes->end()) + { + *elementCount = static_cast(defaultAttr->second.ElementCount()); + } + } - const std::string* protoString = GetStringAttribute(name, elementIndex); - - size_t stringLength = protoString->size(); - ML_CHECK_BOOL(stringLength < attributeElementByteLength); - memcpy(attributeElement, protoString->c_str(), stringLength + 1); - - return S_OK; - } - ORT_CATCH_RETURN -} - -template -template -HRESULT OpNodeInfoWrapper::GetAttributeHelper( - _In_z_ const char* name, - uint32_t elementByteSize, - void* value) const { - using elementType_t = typename MLAttributeTypeTraits::Type; - static_assert(!MLAttributeTypeTraits::IsArray, "This function only works for simple non-array types."); - ML_CHECK_BOOL(sizeof(elementType_t) == elementByteSize); - THROW_IF_NOT_OK(m_impl->template GetAttr(name, static_cast(value))); - return S_OK; -} - -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetTensorAttribute( - _In_z_ const char* name, - _Outptr_ IMLOperatorTensor** tensor) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *tensor = nullptr; - - // Read the tensor if present, and wrap it in a IMLOperatorTensor. - const onnx::AttributeProto* attributeProto = m_impl->TryGetAttribute(std::string(name)); - if (attributeProto) { - if (attributeProto->has_t()) { - const onnx::TensorProto* tensorProto = &attributeProto->t(); - Microsoft::WRL::ComPtr tensorWrapper = wil::MakeOrThrow(const_cast(tensorProto)); - *tensor = tensorWrapper.Detach(); return S_OK; } - } - - return E_INVALIDARG; // The argument has no valid matching attribute. + ORT_CATCH_RETURN } - ORT_CATCH_RETURN -} -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputEdgeDescription(uint32_t inputIndex, MLOperatorEdgeDescription* edgeDesc) const noexcept -{ - ORT_TRY + template + template + HRESULT OpNodeInfoWrapper::GetAttributeArrayHelper( + _In_z_ const char* name, + uint32_t elementCount, + uint32_t elementByteSize, + void* values + ) const { - VerifyNotClosed(); + using elementType_t = typename MLAttributeTypeTraits::Type; + static_assert(MLAttributeTypeTraits::IsArray, "This function only works with array types."); + ML_CHECK_BOOL(sizeof(elementType_t) == elementByteSize); - memset(edgeDesc, 0, sizeof(*edgeDesc)); - const onnx::TypeProto* type = m_impl->GetInputType(inputIndex); - ML_CHECK_BOOL(type != nullptr); - *edgeDesc = ToMLEdgeDesc(type); - - assert(edgeDesc->edgeType != MLOperatorEdgeType::Undefined); - assert((edgeDesc->edgeType != MLOperatorEdgeType::Tensor /*&& edgeDesc->edgeType != MLOperatorEdgeType::TensorSequence*/) || - edgeDesc->tensorDataType != MLOperatorTensorDataType::Undefined); - - return S_OK; + THROW_IF_NOT_OK(m_impl->GetAttrs(name, gsl::span(static_cast::Type*>(values), elementCount))); + return S_OK; } - ORT_CATCH_RETURN -} -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetOutputEdgeDescription(uint32_t outputIndex, MLOperatorEdgeDescription* edgeDesc) const noexcept -{ - ORT_TRY + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetAttribute( + _In_z_ const char* name, + MLOperatorAttributeType type, + uint32_t elementCount, + size_t elementByteSize, + /*out*/void* attributeValue) const noexcept { - VerifyNotClosed(); - - memset(edgeDesc, 0, sizeof(*edgeDesc)); - const onnx::TypeProto* type = m_impl->GetOutputType(outputIndex); - ML_CHECK_BOOL(type != nullptr); - *edgeDesc = ToMLEdgeDesc(type); - - return S_OK; - } - ORT_CATCH_RETURN -} - -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputTensorShape(uint32_t inputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - memset(dimensions, 0, dimensionCount * sizeof(dimensions[0])); - if (inputIndex >= GetInputCount()) { - return E_INVALIDARG; - } - - // Input shapes are determined either from the override or from the underlying proto - if (m_inputShapesOverride) { - if (m_inputShapesOverride->GetShape(inputIndex).size() != dimensionCount) { - return E_INVALIDARG; - } - - for (uint32_t i = 0; i < dimensionCount; ++i) { - dimensions[i] = m_inputShapesOverride->GetShape(inputIndex)[i]; - } - } else { - const auto* inputType = m_impl->GetInputType(inputIndex); - ML_CHECK_BOOL(inputType->has_tensor_type()); - for (uint32_t i = 0; i < dimensionCount; ++i) { - // Shape inference is only done when all dimensions of all inputs have known values, - // so the input tensors will always have shapes at this point. - assert(inputType->tensor_type().shape().dim(i).has_dim_value()); - dimensions[i] = static_cast(inputType->tensor_type().shape().dim(i).dim_value()); - } - } - - return S_OK; - } - ORT_CATCH_RETURN -} - -template -bool STDMETHODCALLTYPE OpNodeInfoWrapper::IsInputValid(uint32_t inputIndex) const noexcept { - if (IsClosed()) { - return false; - } - - return (GetInputCount() > inputIndex) && !!m_impl->GetInputType(inputIndex); -} - -template -bool STDMETHODCALLTYPE OpNodeInfoWrapper::IsOutputValid(uint32_t outputIndex) const noexcept { - if (IsClosed()) { - return false; - } - - return (GetOutputCount() > outputIndex) && !!m_impl->GetOutputType(outputIndex); -} - -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputTensorDimensionCount(uint32_t inputIndex, uint32_t* dimensionCount) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *dimensionCount = 0; - - if (inputIndex >= GetInputCount()) { - return E_INVALIDARG; - } - - // Input shapes are determined either from the override or from the underlying proto - if (m_inputShapesOverride) { - *dimensionCount = gsl::narrow_cast(m_inputShapesOverride->GetShape(inputIndex).size()); - } else { - const auto* inputType = m_impl->GetInputType(inputIndex); - ML_CHECK_BOOL(inputType->has_tensor_type()); - - // Shape inference is only done when all dimensions of all inputs have known values, - // so the input tensors will always have shapes at this point. - assert(inputType->tensor_type().has_shape()); - - *dimensionCount = inputType->tensor_type().shape().dim_size(); - } - - return S_OK; - } - ORT_CATCH_RETURN -} - -template -HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetConstantInputTensor(uint32_t inputIndex, IMLOperatorTensor** tensor) const noexcept -{ - ORT_TRY - { - bool inputRequiredAsConstant = std::find( - m_requiredConstantCpuInputs.begin(), - m_requiredConstantCpuInputs.end(), - inputIndex) != m_requiredConstantCpuInputs.end(); - - ORT_THROW_HR_IF(E_INVALIDARG, !inputRequiredAsConstant); - - ComPtr tensorWrapper = m_constantInputGetter(inputIndex); - - if (tensorWrapper == nullptr) { - // This shouldn't happen since kernel creation is deferred and repeated when required constant inputs are not present. - return E_UNEXPECTED; - } - - *tensor = tensorWrapper.Detach(); - - return S_OK; - } - ORT_CATCH_RETURN -} - -HRESULT STDMETHODCALLTYPE OpKernelInfoWrapper::GetOutputTensorShape(uint32_t outputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - memset(dimensions, 0, dimensionCount * sizeof(dimensions[0])); - - if (!HasOutputShapeDescription()) { - return E_FAIL; - } - - if (outputIndex >= GetOutputCount()) { - return E_INVALIDARG; - } - - if (m_inferredOutputShapes->GetShape(outputIndex).size() != dimensionCount) { - return E_INVALIDARG; - } - - for (uint32_t i = 0; i < dimensionCount; ++i) { - dimensions[i] = m_inferredOutputShapes->GetShape(outputIndex)[i]; - } - - return S_OK; - } - ORT_CATCH_RETURN -} - -HRESULT STDMETHODCALLTYPE OpKernelInfoWrapper::GetOutputTensorDimensionCount(uint32_t outputIndex, uint32_t* dimensionCount) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *dimensionCount = 0; - - if (!HasOutputShapeDescription()) { - return E_FAIL; - } - - if (outputIndex >= GetOutputCount()) { - return E_INVALIDARG; - } - - *dimensionCount = gsl::narrow_cast(m_inferredOutputShapes->GetShape(outputIndex).size()); - - return S_OK; - } - ORT_CATCH_RETURN -} - -bool STDMETHODCALLTYPE OpKernelInfoWrapper::HasTensorShapeDescription() const noexcept { - return m_allowInputShapeQuery; -} - -HRESULT STDMETHODCALLTYPE OpKernelInfoWrapper::GetTensorShapeDescription(IMLOperatorTensorShapeDescription** shapeInfo) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *shapeInfo = nullptr; - - if (!HasTensorShapeDescription()) { - *shapeInfo = nullptr; - return E_FAIL; - //return MLStatus::REQUIREMENT_NOT_REGISTERED; - } - - ComPtr ret = const_cast(this); - *shapeInfo = ret.Detach(); - return S_OK; - } - ORT_CATCH_RETURN -} - -void STDMETHODCALLTYPE OpKernelInfoWrapper::GetExecutionInterface(IUnknown** executionInterface) const noexcept { - m_abiExecutionObject.CopyTo(executionInterface); -} - -template -uint32_t STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputCount() const noexcept { - if (IsClosed()) { - return 0; - } - - return m_impl->GetInputCount(); -} - -template -uint32_t STDMETHODCALLTYPE OpNodeInfoWrapper::GetOutputCount() const noexcept { - if (IsClosed()) { - return 0; - } - - return m_impl->GetOutputCount(); -} - -bool STDMETHODCALLTYPE OpKernelInfoWrapper::HasOutputShapeDescription() const noexcept { - return m_allowOutputShapeQuery; -} - -DmlGraphOpKernelInfoWrapper::DmlGraphOpKernelInfoWrapper( - const onnxruntime::OpNodeProtoHelper* protoHelper, - const void* executionHandle, - bool isInternalOperator, - const EdgeShapes* inferredOutputShapes, - const AttributeMap* defaultAttributes, - DmlGraphNodeCreateInfo* graphNodeCreateInfo, - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter) : OpNodeInfoWrapper(protoHelper, nullptr, defaultAttributes, requiredConstantCpuInputs, constantInputGetter), - m_inferredOutputShapes(inferredOutputShapes), - m_internalOperator(isInternalOperator), - m_graphNodeCreateInfo(graphNodeCreateInfo) { - // We assume the execution object inherits IUnknown as its first base - m_abiExecutionObject = const_cast(static_cast(executionHandle)); - m_abiExecutionObject.As(&m_winmlProvider); -} - -HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetOutputTensorShape(uint32_t outputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - memset(dimensions, 0, dimensionCount * sizeof(dimensions[0])); - - if (!HasOutputShapeDescription()) { - return E_FAIL; - } - - if (outputIndex >= GetOutputCount()) { - return E_INVALIDARG; - } - - if (m_inferredOutputShapes->GetShape(outputIndex).size() != dimensionCount) { - return E_INVALIDARG; - } - - for (uint32_t i = 0; i < dimensionCount; ++i) { - dimensions[i] = m_inferredOutputShapes->GetShape(outputIndex)[i]; - } - - return S_OK; - } - ORT_CATCH_RETURN -} - -HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetOutputTensorDimensionCount(uint32_t outputIndex, uint32_t* dimensionCount) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *dimensionCount = 0; - - if (!HasOutputShapeDescription()) { - return E_FAIL; - } - - if (outputIndex >= GetOutputCount()) { - return E_INVALIDARG; - } - - *dimensionCount = gsl::narrow_cast(m_inferredOutputShapes->GetShape(outputIndex).size()); - - return S_OK; - } - ORT_CATCH_RETURN -} - -bool STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::HasTensorShapeDescription() const noexcept { - return true; -} - -HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetTensorShapeDescription(IMLOperatorTensorShapeDescription** shapeInfo) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *shapeInfo = nullptr; - - if (!HasTensorShapeDescription()) { - *shapeInfo = nullptr; - return E_FAIL; - //return MLStatus::REQUIREMENT_NOT_REGISTERED; - } - - ComPtr ret = const_cast(this); - *shapeInfo = ret.Detach(); - return S_OK; - } - ORT_CATCH_RETURN -} - -void STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetExecutionInterface(IUnknown** executionInterface) const noexcept { - m_abiExecutionObject.CopyTo(executionInterface); -} - -bool STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::HasOutputShapeDescription() const noexcept { - // DML kernels are only used in graph in graph partitions when shapes are static - return true; -} - -bool STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::IsDmlGraphNode() const noexcept { - return (m_graphNodeCreateInfo != nullptr); -} - -void DmlGraphOpKernelInfoWrapper::SetDmlProperties(_In_ const MLOperatorKernelDmlProperties* dmlProperties) const { - // Populate the mappings between DML in/outs and kernel in/outs. By default they are the same. - if (dmlProperties && dmlProperties->kernelInputIndices) { - m_graphNodeCreateInfo->kernelInputIndices.insert( - m_graphNodeCreateInfo->kernelInputIndices.begin(), - dmlProperties->kernelInputIndices, - dmlProperties->kernelInputIndices + dmlProperties->dmlInputCount); - } else { - m_graphNodeCreateInfo->kernelInputIndices.resize(dmlProperties ? dmlProperties->dmlInputCount : GetInputCount()); - std::iota(m_graphNodeCreateInfo->kernelInputIndices.begin(), m_graphNodeCreateInfo->kernelInputIndices.end(), 0); - } - - if (dmlProperties && dmlProperties->kernelOutputIndices) { - m_graphNodeCreateInfo->kernelOutputIndices.insert( - m_graphNodeCreateInfo->kernelOutputIndices.begin(), - dmlProperties->kernelOutputIndices, - dmlProperties->kernelOutputIndices + dmlProperties->dmlOutputCount); - } else { - m_graphNodeCreateInfo->kernelOutputIndices.resize(dmlProperties ? dmlProperties->dmlOutputCount : GetOutputCount()); - std::iota(m_graphNodeCreateInfo->kernelOutputIndices.begin(), m_graphNodeCreateInfo->kernelOutputIndices.end(), 0); - } - - m_graphNodeCreateInfo->allowHalfPrecisionComputation = dmlProperties ? dmlProperties->allowHalfPrecisionComputation : true; -} - -HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::SetDmlOperator( - IDMLOperator* op, - _In_ const DML_OPERATOR_DESC* desc, - _In_opt_ const MLOperatorKernelDmlProperties* dmlProperties) const noexcept -{ - ORT_TRY - { - ML_CHECK_BOOL(op != nullptr); - ML_CHECK_BOOL(dmlProperties != nullptr); - - m_graphNodeCreateInfo->initialized = true; - - SetDmlProperties(dmlProperties); - - m_graphNodeCreateInfo->op = op; - AbstractOperatorDesc abstractDesc = SchemaHelpers::ConvertOperatorDesc(*desc); - m_graphNodeCreateInfo->desc = std::make_unique(std::move(abstractDesc)); - - return S_OK; - } - ORT_CATCH_RETURN -} - -OnnxTensorWrapper::OnnxTensorWrapper(onnx::TensorProto* impl) : m_impl(impl) { - // The tensor may be stored as raw data or in typed fields. - if (impl->has_raw_data()) { - m_dataPtr = reinterpret_cast(impl->mutable_raw_data()->data()); - m_tensorByteSize = impl->raw_data().size(); - } else { - std::tie(m_unpackedTensor, m_tensorByteSize) = UnpackTensor(*impl); - m_dataPtr = m_unpackedTensor.get(); - } -} - -uint32_t STDMETHODCALLTYPE OnnxTensorWrapper::GetDimensionCount() const noexcept { - if (IsClosed()) { - return 0; - } - - return gsl::narrow_cast(m_impl->dims().size()); -} - -HRESULT STDMETHODCALLTYPE OnnxTensorWrapper::GetShape( - uint32_t dimensionCount, - uint32_t* dimensions) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - std::fill(dimensions, dimensions + dimensionCount, 0u); - - uint32_t count = static_cast(m_impl->dims().size()); - ML_CHECK_BOOL(dimensionCount == count); - - for (uint32_t i = 0; i < dimensionCount; ++i) { - dimensions[i] = static_cast(m_impl->dims()[i]); - } - - return S_OK; - } - ORT_CATCH_RETURN -} - -MLOperatorTensorDataType STDMETHODCALLTYPE OnnxTensorWrapper::GetTensorDataType() const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - return ToMLTensorDataType(static_cast(m_impl->data_type())); - } - ORT_CATCH_GENERIC - { - return MLOperatorTensorDataType::Undefined; - } -} - -bool STDMETHODCALLTYPE OnnxTensorWrapper::IsCpuData() const noexcept { - return true; -} - -bool STDMETHODCALLTYPE OnnxTensorWrapper::IsDataInterface() const noexcept { - return false; -} - -void* STDMETHODCALLTYPE OnnxTensorWrapper::GetData() noexcept { - if (IsClosed()) { - return nullptr; - } - - return m_dataPtr; -} - -void STDMETHODCALLTYPE OnnxTensorWrapper::GetDataInterface(IUnknown** dataInterface) noexcept { - *dataInterface = nullptr; -} - -TensorWrapper::TensorWrapper(onnxruntime::Tensor* impl, bool isDataInterface, IWinmlExecutionProvider* provider, bool isInternalOperator) : m_impl(impl), m_winmlExecutionProvider(provider), m_internalOperator(isInternalOperator), m_isDataInterface(isDataInterface) { - if (impl) { - if (isDataInterface) { - // We assume that all data handles derive from IUnknown as their first base. - m_dataInterface = static_cast(m_impl->MutableDataRaw()); - - if (m_dataInterface) { - if (m_winmlExecutionProvider) { - // The resource may require conversion to the layout expected according to the kernel options. - // This will return either the original object or a shadow copy which uses a different layout. - // This pattern assumes that Lotus is not re-using tensor allocations, so each output is - // a fresh allocation which will not trigger a conversion in the provider. - m_winmlExecutionProvider->GetShadowCopyIfRequired(m_internalOperator, m_dataInterface.Get(), m_dataInterfaceOrShadowCopy.GetAddressOf()); - - // Get the actual object to be returned from the ABI, which varies for internal and external - // kernels (i.e. ID3D12Resource, versus something that tracks the layout). - TranslateAllocationDataToAbi( - m_winmlExecutionProvider.Get(), - m_internalOperator, - m_impl->Location(), - m_dataInterfaceOrShadowCopy ? m_dataInterfaceOrShadowCopy.Get() : m_dataInterface.Get(), - m_abiDataInterface.GetAddressOf()); - } else { - m_abiDataInterface = m_dataInterface; - } - } - } else { - m_tensorData = m_impl->MutableDataRaw(); - } - } -} - -uint32_t STDMETHODCALLTYPE TensorWrapper::GetDimensionCount() const noexcept { - if (IsClosed()) { - return 0; - } - - return gsl::narrow_cast(m_impl->Shape().NumDimensions()); -} - -HRESULT STDMETHODCALLTYPE TensorWrapper::GetShape( - uint32_t dimensionCount, - uint32_t* dimensions) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - std::fill(dimensions, dimensions + dimensionCount, 0u); - - uint32_t count = static_cast(m_impl->Shape().NumDimensions()); - ML_CHECK_BOOL(dimensionCount == count); - - for (size_t i = 0; i < dimensionCount; ++i) { - dimensions[i] = static_cast(m_impl->Shape()[i]); - } - - return S_OK; - } - ORT_CATCH_RETURN -} - -MLOperatorTensorDataType STDMETHODCALLTYPE TensorWrapper::GetTensorDataType() const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - return ToMLTensorDataType(m_impl->DataType()); - } - ORT_CATCH_GENERIC - { - return MLOperatorTensorDataType::Undefined; - } -} - -bool STDMETHODCALLTYPE TensorWrapper::IsCpuData() const noexcept { - if (IsClosed()) { - return true; - } - - // tells caller whether this tensor is in CPU memory - return !strcmp(m_impl->Location().name, onnxruntime::CPU) || m_impl->Location().mem_type == ::OrtMemType::OrtMemTypeCPUOutput || m_impl->Location().mem_type == ::OrtMemType::OrtMemTypeCPUInput; -} - -bool STDMETHODCALLTYPE TensorWrapper::IsDataInterface() const noexcept { - if (IsClosed()) { - return false; - } - - return m_isDataInterface; -} - -void* STDMETHODCALLTYPE TensorWrapper::GetData() noexcept { - if (IsClosed()) { - return nullptr; - } - - return m_isDataInterface ? nullptr : m_tensorData; -} - -void STDMETHODCALLTYPE TensorWrapper::GetDataInterface(IUnknown** dataInterface) noexcept { - if (!m_isDataInterface) { - VerifyNotClosed(); - *dataInterface = nullptr; - } else { - m_abiDataInterface.CopyTo(dataInterface); - } -} - -void OpKernelContextWrapper::TransitionResourcesForOperatorIfRequired(bool isBeforeOp) { - if (m_winmlProvider->TransitionsRequiredForOperator(m_internalOperator)) { - std::vector resourcesToTransition; - resourcesToTransition.reserve(m_inputTensors.size() + m_outputTensors.size() + m_temporaryAllocations.size()); - - for (uint32_t i = 0; i < m_inputTensors.size(); ++i) { - ComPtr tensor; - ORT_THROW_IF_FAILED(GetInputTensor(i, tensor.GetAddressOf())); - - ComPtr resource; - tensor->GetDataInterface(resource.GetAddressOf()); - if (resource) { - resourcesToTransition.push_back(resource.Get()); - } - } - - for (uint32_t i = 0; i < m_outputTensors.size(); ++i) { - ComPtr tensor; - ORT_THROW_IF_FAILED(GetOutputTensor(i, tensor.GetAddressOf())); - - ComPtr resource; - tensor->GetDataInterface(resource.GetAddressOf()); - if (resource) { - resourcesToTransition.push_back(resource.Get()); - } - } - - for (auto& tempAlloc : m_temporaryAbiAllocations) { - resourcesToTransition.push_back(tempAlloc.Get()); - } - - m_winmlProvider->TransitionResourcesForOperator( - isBeforeOp, - gsl::narrow_cast(resourcesToTransition.size()), - resourcesToTransition.data()); - } -} - -OpKernelContextWrapper::OpKernelContextWrapper( - onnxruntime::OpKernelContext* context, - const onnxruntime::IExecutionProvider* provider, - bool isInternalOperator, - const EdgeShapes* outputShapes) : m_impl(context), m_outputShapes(outputShapes), m_provider(provider), m_internalOperator(isInternalOperator) { - // Pre-size tensor arrays. Member methods return pointers to these which - // are stored in these arrays, which would become stale if the vectors reallocate - // their internal storage. - m_inputTensors.resize(context->InputCount()); - m_outputTensors.resize(context->OutputCount()); - - const void* executionHandle = m_provider->GetExecutionHandle(); - if (executionHandle) { - // We assume the execution object inherits IUnknown as its first base - m_providerExecutionObject = const_cast(static_cast(executionHandle)); - m_providerExecutionObject.As(&m_winmlProvider); - - // Query the actual object to return through the ABI, based on options registered - // with the kernel - m_abiExecutionObject = m_providerExecutionObject; - if (m_winmlProvider) { - m_winmlProvider->GetABIExecutionInterface(isInternalOperator, m_abiExecutionObject.ReleaseAndGetAddressOf()); - } - - TransitionResourcesForOperatorIfRequired(true); - } -} - -OpKernelContextWrapper::~OpKernelContextWrapper() { - ClearTempAllocations(); -} - -void OpKernelContextWrapper::ClearTempAllocations() { - if (m_winmlProvider) { - m_temporaryAllocations.clear(); - m_temporaryAbiAllocations.clear(); - } -} - -void OpKernelContextWrapper::Close() { - if (m_winmlProvider && m_winmlProvider->TransitionsRequiredForOperator(m_internalOperator)) { - TransitionResourcesForOperatorIfRequired(false); - } - - for (auto& tensor : m_inputTensors) { - if (tensor) { - tensor->Close(); - } - } - - for (auto& tensor : m_outputTensors) { - if (tensor) { - tensor->Close(); - } - } - - ClearTempAllocations(); - - Closable::Close(); -} - -HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::GetInputTensor(uint32_t inputIndex, IMLOperatorTensor** tensor) const noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - *tensor = nullptr; - - ML_CHECK_BOOL(inputIndex < m_inputTensors.size()); - - if (m_inputTensors[inputIndex]->GetInterface() == nullptr) { - auto inputTensor = m_impl->Input(inputIndex); - - ComPtr tensorWrapper = wil::MakeOrThrow( - const_cast(inputTensor), - IsAllocationInterface(inputTensor->Location()), - m_winmlProvider.Get(), - m_internalOperator); - - const_cast(this)->m_inputTensors[inputIndex] = tensorWrapper; - } - - const_cast(this)->m_inputTensors[inputIndex].CopyTo(tensor); - - return S_OK; - } - ORT_CATCH_RETURN -} - -HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::GetOutputTensor(uint32_t outputIndex, IMLOperatorTensor** tensor) noexcept -{ - ORT_TRY - { - VerifyNotClosed(); - - *tensor = nullptr; - - ML_CHECK_BOOL(outputIndex < m_outputTensors.size()); - - // GetOutputTensor must be called unless a kernel provides shape inferencing, - // in which case m_outputShapes will be valid here. - if (!m_outputShapes) { - return E_FAIL; - //return MLStatus::SHAPE_INFERENCE_NOT_REGISTERED; - } - - uint32_t dimensionCount = gsl::narrow_cast(m_outputShapes->GetShape(outputIndex).size()); - return GetOutputTensor(outputIndex, dimensionCount, m_outputShapes->GetShape(outputIndex).data(), tensor); - } - ORT_CATCH_RETURN -} - -HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::GetOutputTensor(uint32_t outputIndex, uint32_t dimensions, const uint32_t* dimensionSizes, IMLOperatorTensor** tensor) noexcept -{ ORT_TRY + { + VerifyNotClosed(); + + // Look for a value in the kernel's registered defaults if one does not exist otherwise + if (m_impl->GetPrimitiveAttrElementCount(ToProto(type), name) == 0) + { + if (!m_defaultAttributes) + { + ORT_THROW_HR(E_FAIL); + } + + auto defaultAttr = m_defaultAttributes->find(name); + if (defaultAttr == m_defaultAttributes->end()) + { + ORT_THROW_HR(E_FAIL); + } + + defaultAttr->second.GetAttribute(type, elementCount, elementByteSize, /*out*/attributeValue); + } + else + { + switch (type) + { + case MLOperatorAttributeType::Float: + ML_CHECK_BOOL(elementCount == 1); + return GetAttributeHelper(name, static_cast(elementByteSize), /*out*/attributeValue); + + case MLOperatorAttributeType::Int: + ML_CHECK_BOOL(elementCount == 1); + return GetAttributeHelper(name, static_cast(elementByteSize), /*out*/attributeValue); + + case MLOperatorAttributeType::FloatArray: + return GetAttributeArrayHelper(name, elementCount, static_cast(elementByteSize), /*out*/attributeValue); + + case MLOperatorAttributeType::IntArray: + return GetAttributeArrayHelper(name, elementCount, static_cast(elementByteSize), /*out*/attributeValue); + + default: + ML_CHECK_BOOL(false); + break; + } + } + + return S_OK; + } + ORT_CATCH_RETURN + } + + template + const std::string* OpNodeInfoWrapper::GetStringAttribute( + _In_z_ const char* name, + uint32_t elementIndex) const { - VerifyNotClosed(); - *tensor = nullptr; + // Get the proto attribute + const onnx::AttributeProto* attr = m_impl->TryGetAttribute(std::string(name)); - ML_CHECK_BOOL(outputIndex < m_outputTensors.size()); + // Look for a value in the kernel's registered defaults if one was not found + if (!attr) + { + if (!m_defaultAttributes) + { + ORT_THROW_HR(E_FAIL); + } - // Verify that the provided shape matches the shape determined using the kernel's shape inference function. - if (m_outputTensors[outputIndex]->GetInterface() == nullptr) { - if (m_outputShapes) { - if ((m_outputShapes->GetShape(outputIndex).size() != dimensions || - memcmp(dimensionSizes, m_outputShapes->GetShape(outputIndex).data(), dimensions * sizeof(*dimensionSizes)))) { + auto defaultAttr = m_defaultAttributes->find(name); + if (defaultAttr == m_defaultAttributes->end()) + { + ORT_THROW_HR(E_FAIL); + } + + return defaultAttr->second.GetStringAttribute(name, elementIndex); + } + else + { + // Get the string vector from the attribute + if (attr->has_s()) + { + return &attr->s(); + } + else + { + // Check the size of the vector + ML_CHECK_BOOL(attr->strings_size() > 0); + ML_CHECK_BOOL(elementIndex < static_cast(attr->strings_size())); + + return &attr->strings(elementIndex); + } + } + } + + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetStringAttributeElementLength( + _In_z_ const char* name, + uint32_t elementIndex, + uint32_t* attributeElementByteLength) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + *attributeElementByteLength = 0; + const std::string* protoString = GetStringAttribute(name, elementIndex); + + // Check for overflow and casting safety + ML_CHECK_BOOL(protoString->size() < protoString->size() + 1); + ML_CHECK_BOOL(protoString->size() + 1 <= std::numeric_limits::max()); + + // Set the length including null termination + *attributeElementByteLength = static_cast(protoString->size() + 1); + return S_OK; + } + ORT_CATCH_RETURN + } + + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetStringAttributeElement( + _In_z_ const char* name, + uint32_t elementIndex, + uint32_t attributeElementByteLength, + char* attributeElement) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + const std::string* protoString = GetStringAttribute(name, elementIndex); + + size_t stringLength = protoString->size(); + ML_CHECK_BOOL(stringLength < attributeElementByteLength); + memcpy(attributeElement, protoString->c_str(), stringLength + 1); + + return S_OK; + } + ORT_CATCH_RETURN + } + + template + template + HRESULT OpNodeInfoWrapper::GetAttributeHelper( + _In_z_ const char* name, + uint32_t elementByteSize, + void* value) const + { + using elementType_t = typename MLAttributeTypeTraits::Type; + static_assert(!MLAttributeTypeTraits::IsArray, "This function only works for simple non-array types."); + ML_CHECK_BOOL(sizeof(elementType_t) == elementByteSize); + THROW_IF_NOT_OK(m_impl->template GetAttr(name, static_cast(value))); + return S_OK; + } + + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetTensorAttribute( + _In_z_ const char* name, + _Outptr_ IMLOperatorTensor** tensor) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + *tensor = nullptr; + + // Read the tensor if present, and wrap it in a IMLOperatorTensor. + const onnx::AttributeProto* attributeProto = m_impl->TryGetAttribute(std::string(name)); + if (attributeProto) + { + if (attributeProto->has_t()) + { + const onnx::TensorProto* tensorProto = &attributeProto->t(); + Microsoft::WRL::ComPtr tensorWrapper = wil::MakeOrThrow(const_cast(tensorProto)); + *tensor = tensorWrapper.Detach(); + return S_OK; + } + } + + return E_INVALIDARG; // The argument has no valid matching attribute. + } + ORT_CATCH_RETURN + } + + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputEdgeDescription(uint32_t inputIndex, MLOperatorEdgeDescription* edgeDesc) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + memset(edgeDesc, 0, sizeof(*edgeDesc)); + const onnx::TypeProto* type = m_impl->GetInputType(inputIndex); + ML_CHECK_BOOL(type != nullptr); + *edgeDesc = ToMLEdgeDesc(type); + + assert(edgeDesc->edgeType != MLOperatorEdgeType::Undefined); + assert((edgeDesc->edgeType != MLOperatorEdgeType::Tensor /*&& edgeDesc->edgeType != MLOperatorEdgeType::TensorSequence*/) || + edgeDesc->tensorDataType != MLOperatorTensorDataType::Undefined); + + return S_OK; + } + ORT_CATCH_RETURN + } + + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetOutputEdgeDescription(uint32_t outputIndex, MLOperatorEdgeDescription* edgeDesc) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + memset(edgeDesc, 0, sizeof(*edgeDesc)); + const onnx::TypeProto* type = m_impl->GetOutputType(outputIndex); + ML_CHECK_BOOL(type != nullptr); + *edgeDesc = ToMLEdgeDesc(type); + + return S_OK; + } + ORT_CATCH_RETURN + } + + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputTensorShape(uint32_t inputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + memset(dimensions, 0, dimensionCount * sizeof(dimensions[0])); + if (inputIndex >= GetInputCount()) + { return E_INVALIDARG; } + + // Input shapes are determined either from the override or from the underlying proto + if (m_inputShapesOverride) + { + if (m_inputShapesOverride->GetShape(inputIndex).size() != dimensionCount) + { + return E_INVALIDARG; + } + + for (uint32_t i = 0; i < dimensionCount; ++i) + { + dimensions[i] = m_inputShapesOverride->GetShape(inputIndex)[i]; + } + } + else + { + const auto* inputType = m_impl->GetInputType(inputIndex); + ML_CHECK_BOOL(inputType->has_tensor_type()); + for (uint32_t i = 0; i < dimensionCount; ++i) + { + // Shape inference is only done when all dimensions of all inputs have known values, + // so the input tensors will always have shapes at this point. + assert(inputType->tensor_type().shape().dim(i).has_dim_value()); + dimensions[i] = static_cast(inputType->tensor_type().shape().dim(i).dim_value()); + } + } + + return S_OK; } - std::vector convertedSizes(dimensions); - for (size_t i = 0; i < dimensions; ++i) { - convertedSizes[i] = dimensionSizes[i]; - } - - onnxruntime::TensorShape shape(convertedSizes.data(), dimensions); - auto outputTensor = m_impl->Output(outputIndex, shape); - if (outputTensor) { - ComPtr tensorWrapper = wil::MakeOrThrow( - const_cast(outputTensor), - IsAllocationInterface(outputTensor->Location()), - m_winmlProvider.Get(), - m_internalOperator); - - const_cast(this)->m_outputTensors[outputIndex] = tensorWrapper; - } - } - - m_outputTensors[outputIndex].CopyTo(tensor); - - return S_OK; + ORT_CATCH_RETURN } - ORT_CATCH_RETURN -} -HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::AllocateTemporaryData(size_t size, IUnknown** abiAllocation) const -{ - ORT_TRY + template + bool STDMETHODCALLTYPE OpNodeInfoWrapper::IsInputValid(uint32_t inputIndex) const noexcept { - uint64_t allocId; - return AllocateTemporaryData(size, abiAllocation, &allocId); - } - ORT_CATCH_RETURN -} - -HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::AllocateTemporaryData(size_t size, IUnknown** abiAllocation, uint64_t* allocId) const -{ - ORT_TRY - { - VerifyNotClosed(); - - *abiAllocation = nullptr; - onnxruntime::AllocatorPtr alloc; - THROW_IF_NOT_OK(m_impl->GetTempSpaceAllocator(&alloc)); - - if (!IsAllocationInterface(alloc->Info())) { - return E_FAIL; - } - - ComPtr allocation; - allocation.Attach(static_cast(alloc->Alloc(size))); - - *allocId = m_winmlProvider->TryGetPooledAllocationId(allocation.Get(), 0); - - TranslateAllocationDataToAbi(m_winmlProvider.Get(), m_internalOperator, alloc->Info(), allocation.Get(), abiAllocation); - - if (m_winmlProvider->TransitionsRequiredForOperator(m_internalOperator)) { - m_winmlProvider->TransitionResourcesForOperator(true, 1, abiAllocation); - } - - // Ensure the allocation is freed and transitioned when the context destructs - m_temporaryAllocations.push_back(allocation); - m_temporaryAbiAllocations.push_back(*abiAllocation); - - return S_OK; - } - ORT_CATCH_RETURN -} - -void STDMETHODCALLTYPE OpKernelContextWrapper::GetExecutionInterface(IUnknown** executionInterface) const noexcept { - m_abiExecutionObject.CopyTo(executionInterface); -} - -std::vector OpKernelContextWrapper::GetInputTensors() { - std::vector ret; - ret.reserve(m_inputTensors.size()); - - for (int i = 0; i < m_impl->InputCount(); ++i) { - ComPtr tensor; - ORT_THROW_IF_FAILED(GetInputTensor(i, tensor.GetAddressOf())); - ret.push_back(m_inputTensors[i].Get()); - } - - return ret; -} - -std::vector OpKernelContextWrapper::GetOutputTensors(const EdgeShapes& outputShapes) { - std::vector ret; - ret.reserve(m_outputTensors.size()); - - ORT_THROW_HR_IF(E_INVALIDARG, static_cast(m_impl->OutputCount()) != outputShapes.EdgeCount()); - - for (int i = 0; i < m_impl->OutputCount(); ++i) { - ComPtr tensor; - ORT_THROW_IF_FAILED(GetOutputTensor( - i, - static_cast(outputShapes.GetShape(i).size()), - outputShapes.GetShape(i).data(), - tensor.GetAddressOf())); - - ret.push_back(m_outputTensors[i].Get()); - } - - return ret; -} - -AbiOpKernel::AbiOpKernel( - IMLOperatorKernelFactory* operatorFactory, - const onnxruntime::OpKernelInfo& kerneInfo, - bool requiresInputShapesAtCreation, - bool requiresOutputShapesAtCreation, - bool isInternalOperator, - gsl::span requiredConstantCpuInputs, - IMLOperatorShapeInferrer* shapeInferrer, - const AttributeMap* defaultAttributes) : OpKernel(kerneInfo), - m_requiresInputShapesAtCreation(requiresInputShapesAtCreation), - m_requiresOutputShapesAtCreation(requiresOutputShapesAtCreation), - m_shapeInferrer(shapeInferrer), - m_internalOperator(isInternalOperator), - m_defaultAttributes(defaultAttributes) { - assert(requiresInputShapesAtCreation || !requiresOutputShapesAtCreation); - - m_requiredConstantCpuInputs.assign(requiredConstantCpuInputs.begin(), requiredConstantCpuInputs.end()); - - const void* executionHandle = kerneInfo.GetExecutionProvider()->GetExecutionHandle(); - if (executionHandle) { - // We assume the execution object inherits IUnknown as its first base - ComPtr providerExecutionObject = const_cast(static_cast(executionHandle)); - m_abiExecutionObject = providerExecutionObject; - - // Get the WinML-specific execution provider interface from the execution object. - providerExecutionObject.As(&m_winmlProvider); - - if (m_winmlProvider) { - // Get the particular object to return to a isInternalOperator based on the registration of that kernel. - m_winmlProvider->GetABIExecutionInterface(isInternalOperator, m_abiExecutionObject.ReleaseAndGetAddressOf()); - } - } - - bool requiredConstantCpuInputsAvailable = true; - for (uint32_t index : requiredConstantCpuInputs) { - const onnxruntime::Tensor* tensor = nullptr; - if (!kerneInfo.TryGetConstantInput(index, &tensor) || !tensor) { - requiredConstantCpuInputsAvailable = false; - break; - } - } - - // If input sizes are either available or not required at creation, no need to delay kernel creation. - if (requiredConstantCpuInputsAvailable && (!m_requiresInputShapesAtCreation || InputTensorShapesDefined())) { - auto winmlProviderCapture = m_winmlProvider; - auto internalOpCapture = m_internalOperator; - - MLOperatorTensorGetter constantInputGetter = [kerneInfo, winmlProviderCapture, internalOpCapture](uint32_t index) { - Microsoft::WRL::ComPtr tensorWrapper = nullptr; - const onnxruntime::Tensor* tensor = nullptr; - if (kerneInfo.TryGetConstantInput(index, &tensor)) { - tensorWrapper = wil::MakeOrThrow( - const_cast(tensor), - IsAllocationInterface(tensor->Location()), - winmlProviderCapture.Get(), - internalOpCapture); - } - - return tensorWrapper; - }; - - // If the output size is not dynamic, infer it using the kernel. Then if the output size was predicted - // by schema, verify consistency. The result of inference is stored in m_inferredOutputShapes. - if (m_requiresOutputShapesAtCreation) { - // Use the same list of required inputs for the shape inferrer and the kernel. - InferAndVerifyOutputSizes(m_requiredConstantCpuInputs, constantInputGetter, nullptr, m_inferredOutputShapes); - } - - // Create the kernel while allowing input shape and output shape queries according to options - ComPtr kernelInfoWrapper = wil::MakeOrThrow( - &kerneInfo, - m_abiExecutionObject.Get(), - nullptr, - m_requiresOutputShapesAtCreation ? &m_inferredOutputShapes : nullptr, - m_requiresInputShapesAtCreation, - m_requiresOutputShapesAtCreation, - isInternalOperator, - m_defaultAttributes, - m_requiredConstantCpuInputs, - constantInputGetter); - - ORT_THROW_IF_FAILED(operatorFactory->CreateKernel(kernelInfoWrapper.Get(), m_kernel.GetAddressOf())); - kernelInfoWrapper->Close(); - - // Ensure that scheduled work, if any, is completed before freeing the kernel if the execution - // provider requires this. - if (m_winmlProvider) { - m_winmlProvider->QueueReference(m_kernel.Get()); - } - } else { - m_operatorFactory = operatorFactory; - } -} - -onnxruntime::Status AbiOpKernel::Compute(onnxruntime::OpKernelContext* context) const { - auto winmlProviderCapture = m_winmlProvider; - auto internalOpCapture = m_internalOperator; - - MLOperatorTensorGetter constantInputGetter = [context, winmlProviderCapture, internalOpCapture](uint32_t index) { - Microsoft::WRL::ComPtr tensorWrapper = nullptr; - const onnxruntime::Tensor* tensor = context->Input(static_cast(index)); - if (tensor != nullptr) - { - tensorWrapper = wil::MakeOrThrow( - const_cast(tensor), - tensor ? IsAllocationInterface(tensor->Location()) : false, - winmlProviderCapture.Get(), - internalOpCapture); - } - - return tensorWrapper; - }; - - auto inferShapesAndCreateKernel = [&](const EdgeShapes& inputShapes, EdgeShapes& outputShapes) -> ComPtr { - // If the output size is not dynamic, infer it using the kernel. The result of inference is stored in m_inferredOutputShapes. - if (m_requiresOutputShapesAtCreation) { - // Use the same list of required inputs for the shape inferrer and the kernel. - InferAndVerifyOutputSizes(m_requiredConstantCpuInputs, constantInputGetter, &inputShapes, outputShapes); - } - - // Create the kernel while allowing input shape and output shape queries according to options - ComPtr kernelInfoWrapper = wil::MakeOrThrow( - &Info(), - m_abiExecutionObject.Get(), - &inputShapes, - m_requiresInputShapesAtCreation ? &outputShapes : nullptr, - m_requiresInputShapesAtCreation, - m_requiresOutputShapesAtCreation, - m_internalOperator, - m_defaultAttributes, - m_requiredConstantCpuInputs, - constantInputGetter); - - ComPtr ret; - ORT_THROW_IF_FAILED(m_operatorFactory->CreateKernel(kernelInfoWrapper.Get(), ret.GetAddressOf())); - kernelInfoWrapper->Close(); - - return ret; - }; - - // The kernel creation may have been delayed because input shapes were required but not inferred by schema. - if (RequiresLazyInitialization()) { - std::lock_guard lock(m_mutex); - - if (RequiresLazyInitialization()) { - m_inputShapesOfKernelInference = GetInputShapes(context); - - m_constantInputTensorContentsOfKernel.resize(context->InputCount()); - for (uint32_t index : m_requiredConstantCpuInputs) { - const onnxruntime::Tensor* weakTensor = context->Input(static_cast(index)); - - // Skip optional constant tensors. - if (weakTensor != nullptr) + if (IsClosed()) { - MLOperatorTensor tensor = MLOperatorTensor(constantInputGetter(index).Get()); + return false; + } - if (index >= static_cast(context->InputCount())) { - continue; + return (GetInputCount() > inputIndex) && !!m_impl->GetInputType(inputIndex); + } + + template + bool STDMETHODCALLTYPE OpNodeInfoWrapper::IsOutputValid(uint32_t outputIndex) const noexcept + { + if (IsClosed()) + { + return false; + } + + return (GetOutputCount() > outputIndex) && !!m_impl->GetOutputType(outputIndex); + } + + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputTensorDimensionCount(uint32_t inputIndex, uint32_t* dimensionCount) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + *dimensionCount = 0; + + if (inputIndex >= GetInputCount()) + { + return E_INVALIDARG; } - m_constantInputTensorContentsOfKernel[index].isValid = (tensor.GetInterface() != nullptr); - if (tensor.GetInterface() != nullptr) { - m_constantInputTensorContentsOfKernel[index].shape = tensor.GetShape(); - m_constantInputTensorContentsOfKernel[index].type = tensor.GetTensorDataType(); - m_constantInputTensorContentsOfKernel[index].data.resize(tensor.GetUnalignedTensorByteSize()); + // Input shapes are determined either from the override or from the underlying proto + if (m_inputShapesOverride) + { + *dimensionCount = gsl::narrow_cast(m_inputShapesOverride->GetShape(inputIndex).size()); } - m_constantInputTensorContentsOfKernel[index].data.assign( - reinterpret_cast(tensor.GetByteData()), - reinterpret_cast(tensor.GetByteData()) + tensor.GetUnalignedTensorByteSize()); + else + { + const auto* inputType = m_impl->GetInputType(inputIndex); + ML_CHECK_BOOL(inputType->has_tensor_type()); + + // Shape inference is only done when all dimensions of all inputs have known values, + // so the input tensors will always have shapes at this point. + assert(inputType->tensor_type().has_shape()); + + *dimensionCount = inputType->tensor_type().shape().dim_size(); + } + + return S_OK; } - } - - m_kernel = inferShapesAndCreateKernel(m_inputShapesOfKernelInference, m_inferredOutputShapes); - SetLazyInitialized(); + ORT_CATCH_RETURN } - } else if (m_inputShapesOfKernelInference.EdgeCount() > 0) { - EdgeShapes local_input_shapes = GetInputShapes(context); - bool requiredCpuInputsChanged = false; - for (uint32_t index : m_requiredConstantCpuInputs) { - if (index >= m_constantInputTensorContentsOfKernel.size()) { - continue; - } + template + HRESULT STDMETHODCALLTYPE OpNodeInfoWrapper::GetConstantInputTensor(uint32_t inputIndex, IMLOperatorTensor** tensor) const noexcept + { + ORT_TRY + { + bool inputRequiredAsConstant = std::find( + m_requiredConstantCpuInputs.begin(), + m_requiredConstantCpuInputs.end(), + inputIndex) != m_requiredConstantCpuInputs.end(); - const TensorContent& lastValue = m_constantInputTensorContentsOfKernel[index]; - MLOperatorTensor currentValue(constantInputGetter(index).Get()); + ORT_THROW_HR_IF(E_INVALIDARG, !inputRequiredAsConstant); - if (lastValue.isValid != (currentValue.GetInterface() != nullptr)) { - break; - } + ComPtr tensorWrapper = m_constantInputGetter(inputIndex); - if (lastValue.isValid) { - if (lastValue.shape != currentValue.GetShape() || - lastValue.type != currentValue.GetTensorDataType() || - currentValue.GetUnalignedTensorByteSize() != lastValue.data.size() || - (memcmp(lastValue.data.data(), currentValue.GetByteData(), lastValue.data.size()) != 0)) { - requiredCpuInputsChanged = true; - break; + if (tensorWrapper == nullptr) + { + // This shouldn't happen since kernel creation is deferred and repeated when required constant inputs are not present. + return E_UNEXPECTED; + } + + *tensor = tensorWrapper.Detach(); + + return S_OK; } - } + ORT_CATCH_RETURN } - // In the edge case that the input size is changing across invocations and the kernel requires - // its input size at construction, use a local instance of the kernel. - if (local_input_shapes != m_inputShapesOfKernelInference || requiredCpuInputsChanged) { - EdgeShapes localInferredOutputShapes; - ComPtr localKernel = inferShapesAndCreateKernel(local_input_shapes, localInferredOutputShapes); + HRESULT STDMETHODCALLTYPE OpKernelInfoWrapper::GetOutputTensorShape(uint32_t outputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); - ComPtr kernelContextWrapper = wil::MakeOrThrow( - context, - Info().GetExecutionProvider(), - m_internalOperator, - m_requiresOutputShapesAtCreation ? &localInferredOutputShapes : nullptr); + memset(dimensions, 0, dimensionCount * sizeof(dimensions[0])); - ORT_THROW_IF_FAILED(localKernel->Compute(kernelContextWrapper.Get())); - kernelContextWrapper->Close(); + if (!HasOutputShapeDescription()) + { + return E_FAIL; + } - // Ensure that scheduled work, if any, is completed before freeing the kernel if the execution - // provider requires this. - if (m_winmlProvider) { - m_winmlProvider->QueueReference(localKernel.Get()); - } - return onnxruntime::Status(); - } - } + if (outputIndex >= GetOutputCount()) + { + return E_INVALIDARG; + } - ComPtr kernelContextWrapper = wil::MakeOrThrow( - context, - Info().GetExecutionProvider(), - m_internalOperator, - m_requiresOutputShapesAtCreation ? &m_inferredOutputShapes : nullptr); + if (m_inferredOutputShapes->GetShape(outputIndex).size() != dimensionCount) + { + return E_INVALIDARG; + } - ORT_THROW_IF_FAILED(m_kernel->Compute(kernelContextWrapper.Get())); - kernelContextWrapper->Close(); + for (uint32_t i = 0; i < dimensionCount; ++i) + { + dimensions[i] = m_inferredOutputShapes->GetShape(outputIndex)[i]; + } - // Ensure that scheduled work, if any, is completed before freeing the kernel if the execution - // provider requires this. - if (m_winmlProvider) { - m_winmlProvider->QueueReference(m_kernel.Get()); - } - - return onnxruntime::Status(); -} - -bool AbiOpKernel::InputTensorShapesDefined() const { - onnxruntime::ProtoHelperNodeContext protoContext(Node()); - onnxruntime::OpNodeProtoHelper info(&protoContext); - - return InputTensorShapesDefinedOnNode(info); -} - -EdgeShapes AbiOpKernel::GetInputShapes(onnxruntime::OpKernelContext* context) const { - EdgeShapes ret(context->InputCount()); - - for (size_t i = 0; i < ret.EdgeCount(); ++i) { - // The input type is null if unused - auto inputType = context->InputType(static_cast(i)); - if (inputType != nullptr && inputType->IsTensorType()) { - const onnxruntime::Tensor* tensor = context->Input(static_cast(i)); - if (tensor) { - ret.GetMutableShape(i).resize(tensor->Shape().GetDims().size()); - for (size_t j = 0; j < ret.GetMutableShape(i).size(); ++j) { - ret.GetMutableShape(i)[j] = gsl::narrow_cast(tensor->Shape().GetDims()[j]); + return S_OK; } - } - } - } - - return ret; -} - -void AbiOpKernel::InferAndVerifyOutputSizes( - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter, - const EdgeShapes* inputShapes, - EdgeShapes& outputShapes) const -{ - // call the non member function (below) - Windows::AI::MachineLearning::Adapter::InferAndVerifyOutputSizes( - Node(), - m_defaultAttributes, - m_shapeInferrer.Get(), - requiredConstantCpuInputs, - constantInputGetter, - inputShapes, - outputShapes - ); -} - -void InferAndVerifyOutputSizes( - const onnxruntime::Node& node, - const AttributeMap* defaultAttributes, - IMLOperatorShapeInferrer* shapeInferrer, - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter, - const EdgeShapes* inputShapes, - EdgeShapes& outputShapes) { - onnxruntime::ProtoHelperNodeContext protoContext(node); - onnxruntime::OpNodeProtoHelper info(&protoContext); - - ComPtr inferenceContext = wil::MakeOrThrow(&info, inputShapes, outputShapes, defaultAttributes, requiredConstantCpuInputs, constantInputGetter); - - outputShapes.Reset(info.GetOutputCount()); - - ORT_THROW_IF_FAILED(shapeInferrer->InferOutputShapes(inferenceContext.Get())); - inferenceContext->Close(); - - for (size_t outputIndex = 0; outputIndex < outputShapes.EdgeCount(); ++outputIndex) { - const onnx::TypeProto* outputProto = info.GetOutputType(outputIndex); - - // Skip this output if it is not valid. - if (outputProto == nullptr) { - continue; + ORT_CATCH_RETURN } - if (outputProto->value_case() != onnx::TypeProto::kTensorType) { - ML_CHECK_BOOL(outputShapes.GetShape(outputIndex).empty()); - continue; - } + HRESULT STDMETHODCALLTYPE OpKernelInfoWrapper::GetOutputTensorDimensionCount(uint32_t outputIndex, uint32_t* dimensionCount) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); - const auto& tensorType = outputProto->tensor_type(); + *dimensionCount = 0; - if (tensorType.has_shape()) { - const auto& shape = tensorType.shape(); - ML_CHECK_BOOL(static_cast(shape.dim_size()) == outputShapes.GetShape(outputIndex).size()); + if (!HasOutputShapeDescription()) + { + return E_FAIL; + } - for (uint32_t output_dim = 0; output_dim < outputShapes.GetShape(outputIndex).size(); ++output_dim) { - if (shape.dim(output_dim).has_dim_value()) { - int64_t expected_size = shape.dim(output_dim).dim_value(); - int64_t actual_size = outputShapes.GetShape(outputIndex)[output_dim]; - ML_CHECK_BOOL(expected_size == actual_size); + if (outputIndex >= GetOutputCount()) + { + return E_INVALIDARG; + } + + *dimensionCount = gsl::narrow_cast(m_inferredOutputShapes->GetShape(outputIndex).size()); + + return S_OK; } - } + ORT_CATCH_RETURN } - } -} -ComPtr MLSchemaInferenceContext::Create(onnxruntime::OpNodeProtoHelper* info, + bool STDMETHODCALLTYPE OpKernelInfoWrapper::HasTensorShapeDescription() const noexcept + { + return m_allowInputShapeQuery; + } + + HRESULT STDMETHODCALLTYPE OpKernelInfoWrapper::GetTensorShapeDescription(IMLOperatorTensorShapeDescription** shapeInfo) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + *shapeInfo = nullptr; + + if (!HasTensorShapeDescription()) + { + *shapeInfo = nullptr; + return E_FAIL; + //return MLStatus::REQUIREMENT_NOT_REGISTERED; + } + + ComPtr ret = const_cast(this); + *shapeInfo = ret.Detach(); + return S_OK; + } + ORT_CATCH_RETURN + } + + void STDMETHODCALLTYPE OpKernelInfoWrapper::GetExecutionInterface(IUnknown** executionInterface) const noexcept + { + m_abiExecutionObject.CopyTo(executionInterface); + } + + template + uint32_t STDMETHODCALLTYPE OpNodeInfoWrapper::GetInputCount() const noexcept + { + if (IsClosed()) + { + return 0; + } + + return m_impl->GetInputCount(); + } + + template + uint32_t STDMETHODCALLTYPE OpNodeInfoWrapper::GetOutputCount() const noexcept + { + if (IsClosed()) + { + return 0; + } + + return m_impl->GetOutputCount(); + } + + bool STDMETHODCALLTYPE OpKernelInfoWrapper::HasOutputShapeDescription() const noexcept + { + return m_allowOutputShapeQuery; + } + + DmlGraphOpKernelInfoWrapper::DmlGraphOpKernelInfoWrapper( + const onnxruntime::OpNodeProtoHelper* protoHelper, + const void* executionHandle, + bool isInternalOperator, + const EdgeShapes* inferredOutputShapes, + const AttributeMap* defaultAttributes, + DmlGraphNodeCreateInfo* graphNodeCreateInfo, + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& constantInputGetter + ) + : OpNodeInfoWrapper(protoHelper, nullptr, defaultAttributes, requiredConstantCpuInputs, constantInputGetter), + m_inferredOutputShapes(inferredOutputShapes), + m_internalOperator(isInternalOperator), + m_graphNodeCreateInfo(graphNodeCreateInfo) + { + // We assume the execution object inherits IUnknown as its first base + m_abiExecutionObject = const_cast(static_cast(executionHandle)); + m_abiExecutionObject.As(&m_winmlProvider); + } + + HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetOutputTensorShape(uint32_t outputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + memset(dimensions, 0, dimensionCount * sizeof(dimensions[0])); + + if (!HasOutputShapeDescription()) + { + return E_FAIL; + } + + if (outputIndex >= GetOutputCount()) + { + return E_INVALIDARG; + } + + if (m_inferredOutputShapes->GetShape(outputIndex).size() != dimensionCount) + { + return E_INVALIDARG; + } + + for (uint32_t i = 0; i < dimensionCount; ++i) + { + dimensions[i] = m_inferredOutputShapes->GetShape(outputIndex)[i]; + } + + return S_OK; + } + ORT_CATCH_RETURN + } + + HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetOutputTensorDimensionCount(uint32_t outputIndex, uint32_t* dimensionCount) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + *dimensionCount = 0; + + if (!HasOutputShapeDescription()) + { + return E_FAIL; + } + + if (outputIndex >= GetOutputCount()) + { + return E_INVALIDARG; + } + + *dimensionCount = gsl::narrow_cast(m_inferredOutputShapes->GetShape(outputIndex).size()); + + return S_OK; + } + ORT_CATCH_RETURN + } + + bool STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::HasTensorShapeDescription() const noexcept + { + return true; + } + + HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetTensorShapeDescription(IMLOperatorTensorShapeDescription** shapeInfo) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + *shapeInfo = nullptr; + + if (!HasTensorShapeDescription()) + { + *shapeInfo = nullptr; + return E_FAIL; + //return MLStatus::REQUIREMENT_NOT_REGISTERED; + } + + ComPtr ret = const_cast(this); + *shapeInfo = ret.Detach(); + return S_OK; + } + ORT_CATCH_RETURN + } + + void STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::GetExecutionInterface(IUnknown** executionInterface) const noexcept + { + m_abiExecutionObject.CopyTo(executionInterface); + } + + bool STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::HasOutputShapeDescription() const noexcept + { + // DML kernels are only used in graph in graph partitions when shapes are static + return true; + } + + bool STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::IsDmlGraphNode() const noexcept + { + return (m_graphNodeCreateInfo != nullptr); + } + + void DmlGraphOpKernelInfoWrapper::SetDmlProperties(_In_ const MLOperatorKernelDmlProperties* dmlProperties) const + { + // Populate the mappings between DML in/outs and kernel in/outs. By default they are the same. + if (dmlProperties && dmlProperties->kernelInputIndices) + { + m_graphNodeCreateInfo->kernelInputIndices.insert( + m_graphNodeCreateInfo->kernelInputIndices.begin(), + dmlProperties->kernelInputIndices, + dmlProperties->kernelInputIndices + dmlProperties->dmlInputCount); + } + else + { + m_graphNodeCreateInfo->kernelInputIndices.resize(dmlProperties ? dmlProperties->dmlInputCount : GetInputCount()); + std::iota(m_graphNodeCreateInfo->kernelInputIndices.begin(), m_graphNodeCreateInfo->kernelInputIndices.end(), 0); + } + + if (dmlProperties && dmlProperties->kernelOutputIndices) + { + m_graphNodeCreateInfo->kernelOutputIndices.insert( + m_graphNodeCreateInfo->kernelOutputIndices.begin(), + dmlProperties->kernelOutputIndices, + dmlProperties->kernelOutputIndices + dmlProperties->dmlOutputCount); + } + else + { + m_graphNodeCreateInfo->kernelOutputIndices.resize(dmlProperties ? dmlProperties->dmlOutputCount : GetOutputCount()); + std::iota(m_graphNodeCreateInfo->kernelOutputIndices.begin(), m_graphNodeCreateInfo->kernelOutputIndices.end(), 0); + } + + m_graphNodeCreateInfo->allowHalfPrecisionComputation = dmlProperties ? dmlProperties->allowHalfPrecisionComputation : true; + } + + HRESULT STDMETHODCALLTYPE DmlGraphOpKernelInfoWrapper::SetDmlOperator( + IDMLOperator* op, + _In_ const DML_OPERATOR_DESC* desc, + _In_opt_ const MLOperatorKernelDmlProperties* dmlProperties) const noexcept + { + ORT_TRY + { + ML_CHECK_BOOL(op != nullptr); + ML_CHECK_BOOL(dmlProperties != nullptr); + + m_graphNodeCreateInfo->initialized = true; + + SetDmlProperties(dmlProperties); + + m_graphNodeCreateInfo->op = op; + AbstractOperatorDesc abstractDesc = SchemaHelpers::ConvertOperatorDesc(*desc); + m_graphNodeCreateInfo->desc = std::make_unique(std::move(abstractDesc)); + + return S_OK; + } + ORT_CATCH_RETURN + } + + OnnxTensorWrapper::OnnxTensorWrapper(onnx::TensorProto* impl) : m_impl(impl) + { + // The tensor may be stored as raw data or in typed fields. + if (impl->has_raw_data()) + { + m_dataPtr = reinterpret_cast(impl->mutable_raw_data()->data()); + m_tensorByteSize = impl->raw_data().size(); + } + else + { + std::tie(m_unpackedTensor, m_tensorByteSize) = UnpackTensor(*impl); + m_dataPtr = m_unpackedTensor.get(); + } + } + + uint32_t STDMETHODCALLTYPE OnnxTensorWrapper::GetDimensionCount() const noexcept + { + if (IsClosed()) + { + return 0; + } + + return gsl::narrow_cast(m_impl->dims().size()); + } + + HRESULT STDMETHODCALLTYPE OnnxTensorWrapper::GetShape( + uint32_t dimensionCount, + uint32_t* dimensions) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + std::fill(dimensions, dimensions + dimensionCount, 0u); + + uint32_t count = static_cast(m_impl->dims().size()); + ML_CHECK_BOOL(dimensionCount == count); + + for (uint32_t i = 0; i < dimensionCount; ++i) + { + dimensions[i] = static_cast(m_impl->dims()[i]); + } + + return S_OK; + } + ORT_CATCH_RETURN + } + + MLOperatorTensorDataType STDMETHODCALLTYPE OnnxTensorWrapper::GetTensorDataType() const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + return ToMLTensorDataType(static_cast(m_impl->data_type())); + } + ORT_CATCH_GENERIC + { + return MLOperatorTensorDataType::Undefined; + } + } + + bool STDMETHODCALLTYPE OnnxTensorWrapper::IsCpuData() const noexcept + { + return true; + } + + bool STDMETHODCALLTYPE OnnxTensorWrapper::IsDataInterface() const noexcept + { + return false; + } + + void* STDMETHODCALLTYPE OnnxTensorWrapper::GetData() noexcept + { + if (IsClosed()) + { + return nullptr; + } + + return m_dataPtr; + } + + void STDMETHODCALLTYPE OnnxTensorWrapper::GetDataInterface(IUnknown** dataInterface) noexcept + { + *dataInterface = nullptr; + } + + TensorWrapper::TensorWrapper(onnxruntime::Tensor* impl, bool isDataInterface, IWinmlExecutionProvider* provider, bool isInternalOperator) + : m_impl(impl), + m_winmlExecutionProvider(provider), + m_internalOperator(isInternalOperator), + m_isDataInterface(isDataInterface) + { + if (impl) + { + if (isDataInterface) + { + // We assume that all data handles derive from IUnknown as their first base. + m_dataInterface = static_cast(m_impl->MutableDataRaw()); + + if (m_dataInterface) + { + if (m_winmlExecutionProvider) + { + // The resource may require conversion to the layout expected according to the kernel options. + // This will return either the original object or a shadow copy which uses a different layout. + // This pattern assumes that Lotus is not re-using tensor allocations, so each output is + // a fresh allocation which will not trigger a conversion in the provider. + m_winmlExecutionProvider->GetShadowCopyIfRequired(m_internalOperator, m_dataInterface.Get(), m_dataInterfaceOrShadowCopy.GetAddressOf()); + + // Get the actual object to be returned from the ABI, which varies for internal and external + // kernels (i.e. ID3D12Resource, versus something that tracks the layout). + TranslateAllocationDataToAbi( + m_winmlExecutionProvider.Get(), + m_internalOperator, + m_impl->Location(), + m_dataInterfaceOrShadowCopy ? m_dataInterfaceOrShadowCopy.Get() : m_dataInterface.Get(), + m_abiDataInterface.GetAddressOf()); + } + else + { + m_abiDataInterface = m_dataInterface; + } + } + } + else + { + m_tensorData = m_impl->MutableDataRaw(); + } + } + } + + uint32_t STDMETHODCALLTYPE TensorWrapper::GetDimensionCount() const noexcept + { + if (IsClosed()) + { + return 0; + } + + return gsl::narrow_cast(m_impl->Shape().NumDimensions()); + } + + HRESULT STDMETHODCALLTYPE TensorWrapper::GetShape( + uint32_t dimensionCount, + uint32_t* dimensions) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + std::fill(dimensions, dimensions + dimensionCount, 0u); + + uint32_t count = static_cast(m_impl->Shape().NumDimensions()); + ML_CHECK_BOOL(dimensionCount == count); + + for (size_t i = 0; i < dimensionCount; ++i) + { + dimensions[i] = static_cast(m_impl->Shape()[i]); + } + + return S_OK; + } + ORT_CATCH_RETURN + } + + MLOperatorTensorDataType STDMETHODCALLTYPE TensorWrapper::GetTensorDataType() const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + return ToMLTensorDataType(m_impl->DataType()); + } + ORT_CATCH_GENERIC + { + return MLOperatorTensorDataType::Undefined; + } + } + + bool STDMETHODCALLTYPE TensorWrapper::IsCpuData() const noexcept + { + if (IsClosed()) + { + return true; + } + + // tells caller whether this tensor is in CPU memory + return !strcmp(m_impl->Location().name, onnxruntime::CPU) || m_impl->Location().mem_type == ::OrtMemType::OrtMemTypeCPUOutput || m_impl->Location().mem_type == ::OrtMemType::OrtMemTypeCPUInput; + } + + bool STDMETHODCALLTYPE TensorWrapper::IsDataInterface() const noexcept + { + if (IsClosed()) + { + return false; + } + + return m_isDataInterface; + } + + void* STDMETHODCALLTYPE TensorWrapper::GetData() noexcept + { + if (IsClosed()) + { + return nullptr; + } + + return m_isDataInterface ? nullptr : m_tensorData; + } + + void STDMETHODCALLTYPE TensorWrapper::GetDataInterface(IUnknown** dataInterface) noexcept + { + if (!m_isDataInterface) + { + VerifyNotClosed(); + *dataInterface = nullptr; + } + else + { + m_abiDataInterface.CopyTo(dataInterface); + } + } + + void OpKernelContextWrapper::TransitionResourcesForOperatorIfRequired(bool isBeforeOp) + { + if (m_winmlProvider->TransitionsRequiredForOperator(m_internalOperator)) + { + std::vector resourcesToTransition; + resourcesToTransition.reserve(m_inputTensors.size() + m_outputTensors.size() + m_temporaryAllocations.size()); + + for (uint32_t i = 0; i < m_inputTensors.size(); ++i) + { + ComPtr tensor; + ORT_THROW_IF_FAILED(GetInputTensor(i, tensor.GetAddressOf())); + + ComPtr resource; + tensor->GetDataInterface(resource.GetAddressOf()); + if (resource) + { + resourcesToTransition.push_back(resource.Get()); + } + } + + for (uint32_t i = 0; i < m_outputTensors.size(); ++i) + { + ComPtr tensor; + ORT_THROW_IF_FAILED(GetOutputTensor(i, tensor.GetAddressOf())); + + ComPtr resource; + tensor->GetDataInterface(resource.GetAddressOf()); + if (resource) + { + resourcesToTransition.push_back(resource.Get()); + } + } + + for (auto& tempAlloc : m_temporaryAbiAllocations) + { + resourcesToTransition.push_back(tempAlloc.Get()); + } + + m_winmlProvider->TransitionResourcesForOperator( + isBeforeOp, + gsl::narrow_cast(resourcesToTransition.size()), + resourcesToTransition.data()); + } + } + + OpKernelContextWrapper::OpKernelContextWrapper( + onnxruntime::OpKernelContext* context, + const onnxruntime::IExecutionProvider* provider, + bool isInternalOperator, + const EdgeShapes* outputShapes + ) + : m_impl(context), m_outputShapes(outputShapes), m_provider(provider), m_internalOperator(isInternalOperator) + { + // Pre-size tensor arrays. Member methods return pointers to these which + // are stored in these arrays, which would become stale if the vectors reallocate + // their internal storage. + m_inputTensors.resize(context->InputCount()); + m_outputTensors.resize(context->OutputCount()); + + const void* executionHandle = m_provider->GetExecutionHandle(); + if (executionHandle) + { + // We assume the execution object inherits IUnknown as its first base + m_providerExecutionObject = const_cast(static_cast(executionHandle)); + m_providerExecutionObject.As(&m_winmlProvider); + + // Query the actual object to return through the ABI, based on options registered + // with the kernel + m_abiExecutionObject = m_providerExecutionObject; + if (m_winmlProvider) + { + m_winmlProvider->GetABIExecutionInterface(isInternalOperator, m_abiExecutionObject.ReleaseAndGetAddressOf()); + } + + TransitionResourcesForOperatorIfRequired(true); + } + } + + OpKernelContextWrapper::~OpKernelContextWrapper() + { + ClearTempAllocations(); + } + + void OpKernelContextWrapper::ClearTempAllocations() + { + if (m_winmlProvider) + { + m_temporaryAllocations.clear(); + m_temporaryAbiAllocations.clear(); + } + } + + void OpKernelContextWrapper::Close() + { + if (m_winmlProvider && m_winmlProvider->TransitionsRequiredForOperator(m_internalOperator)) + { + TransitionResourcesForOperatorIfRequired(false); + } + + for (auto& tensor : m_inputTensors) + { + if (tensor) + { + tensor->Close(); + } + } + + for (auto& tensor : m_outputTensors) + { + if (tensor) + { + tensor->Close(); + } + } + + ClearTempAllocations(); + + Closable::Close(); + } + + HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::GetInputTensor(uint32_t inputIndex, IMLOperatorTensor** tensor) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); + *tensor = nullptr; + + ML_CHECK_BOOL(inputIndex < m_inputTensors.size()); + + if (m_inputTensors[inputIndex]->GetInterface() == nullptr) + { + auto inputTensor = m_impl->Input(inputIndex); + + ComPtr tensorWrapper = wil::MakeOrThrow( + const_cast(inputTensor), + IsAllocationInterface(inputTensor->Location()), + m_winmlProvider.Get(), + m_internalOperator); + + const_cast(this)->m_inputTensors[inputIndex] = tensorWrapper; + } + + const_cast(this)->m_inputTensors[inputIndex].CopyTo(tensor); + + return S_OK; + } + ORT_CATCH_RETURN + } + + HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::GetOutputTensor(uint32_t outputIndex, IMLOperatorTensor** tensor) noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + *tensor = nullptr; + + ML_CHECK_BOOL(outputIndex < m_outputTensors.size()); + + // GetOutputTensor must be called unless a kernel provides shape inferencing, + // in which case m_outputShapes will be valid here. + if (!m_outputShapes) + { + return E_FAIL; + //return MLStatus::SHAPE_INFERENCE_NOT_REGISTERED; + } + + uint32_t dimensionCount = gsl::narrow_cast(m_outputShapes->GetShape(outputIndex).size()); + return GetOutputTensor(outputIndex, dimensionCount, m_outputShapes->GetShape(outputIndex).data(), tensor); + } + ORT_CATCH_RETURN + } + + HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::GetOutputTensor(uint32_t outputIndex, uint32_t dimensions, const uint32_t* dimensionSizes, IMLOperatorTensor** tensor) noexcept + { + ORT_TRY + { + VerifyNotClosed(); + *tensor = nullptr; + + ML_CHECK_BOOL(outputIndex < m_outputTensors.size()); + + // Verify that the provided shape matches the shape determined using the kernel's shape inference function. + if (m_outputTensors[outputIndex]->GetInterface() == nullptr) + { + if (m_outputShapes) + { + if ((m_outputShapes->GetShape(outputIndex).size() != dimensions || + memcmp(dimensionSizes, m_outputShapes->GetShape(outputIndex).data(), dimensions * sizeof(*dimensionSizes)))) + { + return E_INVALIDARG; + } + } + std::vector convertedSizes(dimensions); + for (size_t i = 0; i < dimensions; ++i) + { + convertedSizes[i] = dimensionSizes[i]; + } + + onnxruntime::TensorShape shape(convertedSizes.data(), dimensions); + auto outputTensor = m_impl->Output(outputIndex, shape); + if (outputTensor) + { + ComPtr tensorWrapper = wil::MakeOrThrow( + const_cast(outputTensor), + IsAllocationInterface(outputTensor->Location()), + m_winmlProvider.Get(), + m_internalOperator); + + const_cast(this)->m_outputTensors[outputIndex] = tensorWrapper; + } + } + + m_outputTensors[outputIndex].CopyTo(tensor); + + return S_OK; + } + ORT_CATCH_RETURN + } + + HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::AllocateTemporaryData(size_t size, IUnknown** abiAllocation) const + { + ORT_TRY + { + uint64_t allocId; + return AllocateTemporaryData(size, abiAllocation, &allocId); + } + ORT_CATCH_RETURN + } + + HRESULT STDMETHODCALLTYPE OpKernelContextWrapper::AllocateTemporaryData(size_t size, IUnknown** abiAllocation, uint64_t* allocId) const + { + ORT_TRY + { + VerifyNotClosed(); + + *abiAllocation = nullptr; + onnxruntime::AllocatorPtr alloc; + THROW_IF_NOT_OK(m_impl->GetTempSpaceAllocator(&alloc)); + + if (!IsAllocationInterface(alloc->Info())) + { + return E_FAIL; + } + + ComPtr allocation; + allocation.Attach(static_cast(alloc->Alloc(size))); + + *allocId = m_winmlProvider->TryGetPooledAllocationId(allocation.Get(), 0); + + TranslateAllocationDataToAbi(m_winmlProvider.Get(), m_internalOperator, alloc->Info(), allocation.Get(), abiAllocation); + + if (m_winmlProvider->TransitionsRequiredForOperator(m_internalOperator)) + { + m_winmlProvider->TransitionResourcesForOperator(true, 1, abiAllocation); + } + + // Ensure the allocation is freed and transitioned when the context destructs + m_temporaryAllocations.push_back(allocation); + m_temporaryAbiAllocations.push_back(*abiAllocation); + + return S_OK; + } + ORT_CATCH_RETURN + } + + void STDMETHODCALLTYPE OpKernelContextWrapper::GetExecutionInterface(IUnknown** executionInterface) const noexcept + { + m_abiExecutionObject.CopyTo(executionInterface); + } + + std::vector OpKernelContextWrapper::GetInputTensors() + { + std::vector ret; + ret.reserve(m_inputTensors.size()); + + for (int i = 0; i < m_impl->InputCount(); ++i) + { + ComPtr tensor; + ORT_THROW_IF_FAILED(GetInputTensor(i, tensor.GetAddressOf())); + ret.push_back(m_inputTensors[i].Get()); + } + + return ret; + } + + std::vector OpKernelContextWrapper::GetOutputTensors(const EdgeShapes& outputShapes) + { + std::vector ret; + ret.reserve(m_outputTensors.size()); + + ORT_THROW_HR_IF(E_INVALIDARG, static_cast(m_impl->OutputCount()) != outputShapes.EdgeCount()); + + for (int i = 0; i < m_impl->OutputCount(); ++i) + { + ComPtr tensor; + ORT_THROW_IF_FAILED(GetOutputTensor( + i, + static_cast(outputShapes.GetShape(i).size()), + outputShapes.GetShape(i).data(), + tensor.GetAddressOf())); + + ret.push_back(m_outputTensors[i].Get()); + } + + return ret; + } + + AbiOpKernel::AbiOpKernel( + IMLOperatorKernelFactory* operatorFactory, + const onnxruntime::OpKernelInfo& kerneInfo, + bool requiresInputShapesAtCreation, + bool requiresOutputShapesAtCreation, + bool isInternalOperator, + gsl::span requiredConstantCpuInputs, + IMLOperatorShapeInferrer* shapeInferrer, + const AttributeMap* defaultAttributes) + : OpKernel(kerneInfo), + m_requiresInputShapesAtCreation(requiresInputShapesAtCreation), + m_requiresOutputShapesAtCreation(requiresOutputShapesAtCreation), + m_shapeInferrer(shapeInferrer), + m_internalOperator(isInternalOperator), + m_defaultAttributes(defaultAttributes) + { + assert(requiresInputShapesAtCreation || !requiresOutputShapesAtCreation); + + m_requiredConstantCpuInputs.assign(requiredConstantCpuInputs.begin(), requiredConstantCpuInputs.end()); + + const void* executionHandle = kerneInfo.GetExecutionProvider()->GetExecutionHandle(); + if (executionHandle) + { + // We assume the execution object inherits IUnknown as its first base + ComPtr providerExecutionObject = const_cast(static_cast(executionHandle)); + m_abiExecutionObject = providerExecutionObject; + + // Get the WinML-specific execution provider interface from the execution object. + providerExecutionObject.As(&m_winmlProvider); + + if (m_winmlProvider) + { + // Get the particular object to return to a isInternalOperator based on the registration of that kernel. + m_winmlProvider->GetABIExecutionInterface(isInternalOperator, m_abiExecutionObject.ReleaseAndGetAddressOf()); + } + } + + bool requiredConstantCpuInputsAvailable = true; + for (uint32_t index : requiredConstantCpuInputs) + { + const onnxruntime::Tensor* tensor = nullptr; + if (!kerneInfo.TryGetConstantInput(index, &tensor) || !tensor) + { + requiredConstantCpuInputsAvailable = false; + break; + } + } + + // If input sizes are either available or not required at creation, no need to delay kernel creation. + if (requiredConstantCpuInputsAvailable && (!m_requiresInputShapesAtCreation || InputTensorShapesDefined())) + { + auto winmlProviderCapture = m_winmlProvider; + auto internalOpCapture = m_internalOperator; + + MLOperatorTensorGetter constantInputGetter = [kerneInfo, winmlProviderCapture, internalOpCapture](uint32_t index) + { + Microsoft::WRL::ComPtr tensorWrapper = nullptr; + const onnxruntime::Tensor* tensor = nullptr; + if (kerneInfo.TryGetConstantInput(index, &tensor)) + { + tensorWrapper = wil::MakeOrThrow( + const_cast(tensor), + IsAllocationInterface(tensor->Location()), + winmlProviderCapture.Get(), + internalOpCapture); + } + + return tensorWrapper; + }; + + // If the output size is not dynamic, infer it using the kernel. Then if the output size was predicted + // by schema, verify consistency. The result of inference is stored in m_inferredOutputShapes. + if (m_requiresOutputShapesAtCreation) + { + // Use the same list of required inputs for the shape inferrer and the kernel. + InferAndVerifyOutputSizes(m_requiredConstantCpuInputs, constantInputGetter, nullptr, m_inferredOutputShapes); + } + + // Create the kernel while allowing input shape and output shape queries according to options + ComPtr kernelInfoWrapper = wil::MakeOrThrow( + &kerneInfo, + m_abiExecutionObject.Get(), + nullptr, + m_requiresOutputShapesAtCreation ? &m_inferredOutputShapes : nullptr, + m_requiresInputShapesAtCreation, + m_requiresOutputShapesAtCreation, + isInternalOperator, + m_defaultAttributes, + m_requiredConstantCpuInputs, + constantInputGetter); + + ORT_THROW_IF_FAILED(operatorFactory->CreateKernel(kernelInfoWrapper.Get(), m_kernel.GetAddressOf())); + kernelInfoWrapper->Close(); + + // Ensure that scheduled work, if any, is completed before freeing the kernel if the execution + // provider requires this. + if (m_winmlProvider) + { + m_winmlProvider->QueueReference(m_kernel.Get()); + } + } + else + { + m_operatorFactory = operatorFactory; + } + } + + onnxruntime::Status AbiOpKernel::Compute(onnxruntime::OpKernelContext* context) const + { + auto winmlProviderCapture = m_winmlProvider; + auto internalOpCapture = m_internalOperator; + + MLOperatorTensorGetter constantInputGetter = [context, winmlProviderCapture, internalOpCapture](uint32_t index) + { + Microsoft::WRL::ComPtr tensorWrapper = nullptr; + const onnxruntime::Tensor* tensor = context->Input(static_cast(index)); + if (tensor != nullptr) + { + tensorWrapper = wil::MakeOrThrow( + const_cast(tensor), + tensor ? IsAllocationInterface(tensor->Location()) : false, + winmlProviderCapture.Get(), + internalOpCapture); + } + + return tensorWrapper; + }; + + auto inferShapesAndCreateKernel = [&](const EdgeShapes& inputShapes, EdgeShapes& outputShapes) -> ComPtr { + // If the output size is not dynamic, infer it using the kernel. The result of inference is stored in m_inferredOutputShapes. + if (m_requiresOutputShapesAtCreation) + { + // Use the same list of required inputs for the shape inferrer and the kernel. + InferAndVerifyOutputSizes(m_requiredConstantCpuInputs, constantInputGetter, &inputShapes, outputShapes); + } + + // Create the kernel while allowing input shape and output shape queries according to options + ComPtr kernelInfoWrapper = wil::MakeOrThrow( + &Info(), + m_abiExecutionObject.Get(), + &inputShapes, + m_requiresInputShapesAtCreation ? &outputShapes : nullptr, + m_requiresInputShapesAtCreation, + m_requiresOutputShapesAtCreation, + m_internalOperator, + m_defaultAttributes, + m_requiredConstantCpuInputs, + constantInputGetter); + + ComPtr ret; + ORT_THROW_IF_FAILED(m_operatorFactory->CreateKernel(kernelInfoWrapper.Get(), ret.GetAddressOf())); + kernelInfoWrapper->Close(); + + return ret; + }; + + // The kernel creation may have been delayed because input shapes were required but not inferred by schema. + if (RequiresLazyInitialization()) + { + std::lock_guard lock(m_mutex); + + if (RequiresLazyInitialization()) + { + m_inputShapesOfKernelInference = GetInputShapes(context); + + m_constantInputTensorContentsOfKernel.resize(context->InputCount()); + for (uint32_t index : m_requiredConstantCpuInputs) + { + const onnxruntime::Tensor* weakTensor = context->Input(static_cast(index)); + + // Skip optional constant tensors. + if (weakTensor != nullptr) + { + MLOperatorTensor tensor = MLOperatorTensor(constantInputGetter(index).Get()); + + if (index >= static_cast(context->InputCount())) + { + continue; + } + m_constantInputTensorContentsOfKernel[index].isValid = (tensor.GetInterface() != nullptr); + + if (tensor.GetInterface() != nullptr) + { + m_constantInputTensorContentsOfKernel[index].shape = tensor.GetShape(); + m_constantInputTensorContentsOfKernel[index].type = tensor.GetTensorDataType(); + m_constantInputTensorContentsOfKernel[index].data.resize(tensor.GetUnalignedTensorByteSize()); + } + m_constantInputTensorContentsOfKernel[index].data.assign( + reinterpret_cast(tensor.GetByteData()), + reinterpret_cast(tensor.GetByteData()) + tensor.GetUnalignedTensorByteSize()); + } + } + + m_kernel = inferShapesAndCreateKernel(m_inputShapesOfKernelInference, m_inferredOutputShapes); + SetLazyInitialized(); + } + } + else if (m_inputShapesOfKernelInference.EdgeCount() > 0) + { + EdgeShapes local_input_shapes = GetInputShapes(context); + + bool requiredCpuInputsChanged = false; + for (uint32_t index : m_requiredConstantCpuInputs) + { + if (index >= m_constantInputTensorContentsOfKernel.size()) + { + continue; + } + + const TensorContent& lastValue = m_constantInputTensorContentsOfKernel[index]; + MLOperatorTensor currentValue(constantInputGetter(index).Get()); + + if (lastValue.isValid != (currentValue.GetInterface() != nullptr)) + { + break; + } + + if (lastValue.isValid) + { + if (lastValue.shape != currentValue.GetShape() || + lastValue.type != currentValue.GetTensorDataType() || + currentValue.GetUnalignedTensorByteSize() != lastValue.data.size() || + (memcmp(lastValue.data.data(), currentValue.GetByteData(), lastValue.data.size()) != 0)) + { + requiredCpuInputsChanged = true; + break; + } + } + } + + // In the edge case that the input size is changing across invocations and the kernel requires + // its input size at construction, use a local instance of the kernel. + if (local_input_shapes != m_inputShapesOfKernelInference || requiredCpuInputsChanged) + { + EdgeShapes localInferredOutputShapes; + ComPtr localKernel = inferShapesAndCreateKernel(local_input_shapes, localInferredOutputShapes); + + ComPtr kernelContextWrapper = wil::MakeOrThrow( + context, + Info().GetExecutionProvider(), + m_internalOperator, + m_requiresOutputShapesAtCreation ? &localInferredOutputShapes : nullptr); + + ORT_THROW_IF_FAILED(localKernel->Compute(kernelContextWrapper.Get())); + kernelContextWrapper->Close(); + + // Ensure that scheduled work, if any, is completed before freeing the kernel if the execution + // provider requires this. + if (m_winmlProvider) + { + m_winmlProvider->QueueReference(localKernel.Get()); + } + return onnxruntime::Status(); + } + } + + ComPtr kernelContextWrapper = wil::MakeOrThrow( + context, + Info().GetExecutionProvider(), + m_internalOperator, + m_requiresOutputShapesAtCreation ? &m_inferredOutputShapes : nullptr); + + ORT_THROW_IF_FAILED(m_kernel->Compute(kernelContextWrapper.Get())); + kernelContextWrapper->Close(); + + // Ensure that scheduled work, if any, is completed before freeing the kernel if the execution + // provider requires this. + if (m_winmlProvider) + { + m_winmlProvider->QueueReference(m_kernel.Get()); + } + + return onnxruntime::Status(); + } + + bool AbiOpKernel::InputTensorShapesDefined() const + { + onnxruntime::ProtoHelperNodeContext protoContext(Node()); + onnxruntime::OpNodeProtoHelper info(&protoContext); + + return InputTensorShapesDefinedOnNode(info); + } + + EdgeShapes AbiOpKernel::GetInputShapes(onnxruntime::OpKernelContext* context) const + { + EdgeShapes ret(context->InputCount()); + + for (size_t i = 0; i < ret.EdgeCount(); ++i) + { + // The input type is null if unused + auto inputType = context->InputType(static_cast(i)); + if (inputType != nullptr && inputType->IsTensorType()) + { + const onnxruntime::Tensor* tensor = context->Input(static_cast(i)); + if (tensor) + { + ret.GetMutableShape(i).resize(tensor->Shape().GetDims().size()); + for (size_t j = 0; j < ret.GetMutableShape(i).size(); ++j) + { + ret.GetMutableShape(i)[j] = gsl::narrow_cast(tensor->Shape().GetDims()[j]); + } + } + } + } + + return ret; + } + + void AbiOpKernel::InferAndVerifyOutputSizes( + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& constantInputGetter, + const EdgeShapes* inputShapes, + EdgeShapes& outputShapes) const + { + // call the non member function (below) + Windows::AI::MachineLearning::Adapter::InferAndVerifyOutputSizes( + Node(), + m_defaultAttributes, + m_shapeInferrer.Get(), + requiredConstantCpuInputs, + constantInputGetter, + inputShapes, + outputShapes + ); + } + + void InferAndVerifyOutputSizes( + const onnxruntime::Node& node, + const AttributeMap* defaultAttributes, + IMLOperatorShapeInferrer* shapeInferrer, + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& constantInputGetter, + const EdgeShapes* inputShapes, + EdgeShapes& outputShapes) + { + onnxruntime::ProtoHelperNodeContext protoContext(node); + onnxruntime::OpNodeProtoHelper info(&protoContext); + + ComPtr inferenceContext = wil::MakeOrThrow(&info, inputShapes, outputShapes, defaultAttributes, requiredConstantCpuInputs, constantInputGetter); + + outputShapes.Reset(info.GetOutputCount()); + + ORT_THROW_IF_FAILED(shapeInferrer->InferOutputShapes(inferenceContext.Get())); + inferenceContext->Close(); + + for (size_t outputIndex = 0; outputIndex < outputShapes.EdgeCount(); ++outputIndex) + { + const onnx::TypeProto* outputProto = info.GetOutputType(outputIndex); + + // Skip this output if it is not valid. + if (outputProto == nullptr) + { + continue; + } + + if (outputProto->value_case() != onnx::TypeProto::kTensorType) + { + ML_CHECK_BOOL(outputShapes.GetShape(outputIndex).empty()); + continue; + } + + const auto& tensorType = outputProto->tensor_type(); + + if (tensorType.has_shape()) + { + const auto& shape = tensorType.shape(); + ML_CHECK_BOOL(static_cast(shape.dim_size()) == outputShapes.GetShape(outputIndex).size()); + + for (uint32_t output_dim = 0; output_dim < outputShapes.GetShape(outputIndex).size(); ++output_dim) + { + if (shape.dim(output_dim).has_dim_value()) + { + int64_t expected_size = shape.dim(output_dim).dim_value(); + int64_t actual_size = outputShapes.GetShape(outputIndex)[output_dim]; + ML_CHECK_BOOL(expected_size == actual_size); + } + } + } + } + } + + ComPtr MLSchemaInferenceContext::Create(onnxruntime::OpNodeProtoHelper* info, onnx::InferenceContext* ctx, - gsl::span requiredConstantCpuInputs) { - MLOperatorTensorGetter mlOperatorTensorGetter = MLOperatorTensorGetter([ctx](uint32_t index) { - Microsoft::WRL::ComPtr tensorWrapper = wil::MakeOrThrow( - const_cast(ctx->getInputData(index))); - return tensorWrapper; - }); - - return wil::MakeOrThrow(info, ctx, requiredConstantCpuInputs, mlOperatorTensorGetter); -} - -MLSchemaInferenceContext::MLSchemaInferenceContext( - onnxruntime::OpNodeProtoHelper* info, - onnx::InferenceContext* ctx, - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& mLOperatorTensorGetter) : OpNodeInfoWrapper(info, nullptr, nullptr, - requiredConstantCpuInputs, mLOperatorTensorGetter), - m_context(ctx) -{ -} - -HRESULT STDMETHODCALLTYPE MLSchemaInferenceContext::SetOutputTensorShape( - uint32_t outputIndex, - uint32_t dimensionCount, - const uint32_t* dimensions) noexcept -{ - ORT_TRY + gsl::span requiredConstantCpuInputs) { - VerifyNotClosed(); + MLOperatorTensorGetter mlOperatorTensorGetter = MLOperatorTensorGetter( + [ctx](uint32_t index) + { + Microsoft::WRL::ComPtr tensorWrapper = wil::MakeOrThrow( + const_cast(ctx->getInputData(index))); + return tensorWrapper; + } + ); - MLOperatorEdgeDescription edgeDesc; - ORT_THROW_IF_FAILED(GetOutputEdgeDescription(outputIndex, &edgeDesc)); - ML_CHECK_BOOL(edgeDesc.edgeType == MLOperatorEdgeType::Undefined || edgeDesc.edgeType == MLOperatorEdgeType::Tensor); - - // In the process of calling mutable_tensor_type, the type may switch from undefined to tensor. - // This is done here in case the dimension count is zero (scalar) - m_context->getOutputType(outputIndex)->mutable_tensor_type(); - - for (uint32_t i = 0; i < dimensionCount; ++i) { - auto dim = m_context->getOutputType(outputIndex)->mutable_tensor_type()->mutable_shape()->add_dim(); - dim->set_dim_value(dimensions[i]); - } - - return S_OK; + return wil::MakeOrThrow(info, ctx, requiredConstantCpuInputs, mlOperatorTensorGetter); } - ORT_CATCH_RETURN -} -HRESULT STDMETHODCALLTYPE MLSchemaInferenceContext::SetOutputEdgeDescription( - uint32_t outputIndex, - const MLOperatorEdgeDescription* edgeDesc) const noexcept -{ - ORT_TRY + MLSchemaInferenceContext::MLSchemaInferenceContext( + onnxruntime::OpNodeProtoHelper* info, + onnx::InferenceContext* ctx, + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& mLOperatorTensorGetter + ) + : OpNodeInfoWrapper(info, nullptr, nullptr, requiredConstantCpuInputs, mLOperatorTensorGetter), + m_context(ctx) { - VerifyNotClosed(); - - std::string typeStr = ToTypeString(*edgeDesc); - m_context->getOutputType(outputIndex)->CopyFrom(onnx::Utils::DataTypeUtils::ToTypeProto(&typeStr)); - return S_OK; } - ORT_CATCH_RETURN -} -HRESULT STDMETHODCALLTYPE MLKernelInferenceContext::SetOutputTensorShape( - uint32_t outputIndex, - uint32_t dimensionCount, - const uint32_t* dimensions) noexcept -{ - ORT_TRY + HRESULT STDMETHODCALLTYPE MLSchemaInferenceContext::SetOutputTensorShape( + uint32_t outputIndex, + uint32_t dimensionCount, + const uint32_t* dimensions) noexcept { - VerifyNotClosed(); + ORT_TRY + { + VerifyNotClosed(); - if (outputIndex >= m_inferredOutputShapes.EdgeCount()) { - return E_INVALIDARG; - } + MLOperatorEdgeDescription edgeDesc; + ORT_THROW_IF_FAILED(GetOutputEdgeDescription(outputIndex, &edgeDesc)); + ML_CHECK_BOOL(edgeDesc.edgeType == MLOperatorEdgeType::Undefined || edgeDesc.edgeType == MLOperatorEdgeType::Tensor); - m_inferredOutputShapes.GetMutableShape(outputIndex).assign(dimensions, dimensions + dimensionCount); + // In the process of calling mutable_tensor_type, the type may switch from undefined to tensor. + // This is done here in case the dimension count is zero (scalar) + m_context->getOutputType(outputIndex)->mutable_tensor_type(); - return S_OK; + for (uint32_t i = 0; i < dimensionCount; ++i) + { + auto dim = m_context->getOutputType(outputIndex)->mutable_tensor_type()->mutable_shape()->add_dim(); + dim->set_dim_value(dimensions[i]); + } + + return S_OK; + } + ORT_CATCH_RETURN } - ORT_CATCH_RETURN -} -ComPtr MLSupportQueryContext::Create(onnxruntime::OpNodeProtoHelper* info, - const AttributeMap* defaultAttributes) { - MLOperatorTensorGetter mLOperatorTensorGetter = MLOperatorTensorGetter(); - return wil::MakeOrThrow(info, defaultAttributes, mLOperatorTensorGetter); -} + HRESULT STDMETHODCALLTYPE MLSchemaInferenceContext::SetOutputEdgeDescription( + uint32_t outputIndex, + const MLOperatorEdgeDescription* edgeDesc) const noexcept + { + ORT_TRY + { + VerifyNotClosed(); -MLSupportQueryContext::MLSupportQueryContext( + std::string typeStr = ToTypeString(*edgeDesc); + m_context->getOutputType(outputIndex)->CopyFrom(onnx::Utils::DataTypeUtils::ToTypeProto(&typeStr)); + return S_OK; + } + ORT_CATCH_RETURN + } + + HRESULT STDMETHODCALLTYPE MLKernelInferenceContext::SetOutputTensorShape( + uint32_t outputIndex, + uint32_t dimensionCount, + const uint32_t* dimensions) noexcept + { + ORT_TRY + { + VerifyNotClosed(); + + if (outputIndex >= m_inferredOutputShapes.EdgeCount()) + { + return E_INVALIDARG; + } + + m_inferredOutputShapes.GetMutableShape(outputIndex).assign(dimensions, dimensions + dimensionCount); + + return S_OK; + } + ORT_CATCH_RETURN + } + + ComPtr MLSupportQueryContext::Create(onnxruntime::OpNodeProtoHelper* info, + const AttributeMap* defaultAttributes) + { + MLOperatorTensorGetter mLOperatorTensorGetter = MLOperatorTensorGetter(); + return wil::MakeOrThrow(info, defaultAttributes, mLOperatorTensorGetter); + } + + MLSupportQueryContext::MLSupportQueryContext( onnxruntime::OpNodeProtoHelper* info, const AttributeMap* defaultAttributes, - MLOperatorTensorGetter& mLOperatorTensorGetter) : - OpNodeInfoWrapper(info, nullptr, defaultAttributes, gsl::span(), mLOperatorTensorGetter) -{ -} - -bool TryGetStaticShapeIfTensor( - const onnx::TypeProto* inputProto, - std::vector& shapeDims) { - // Skip this input if it is not valid. - if (inputProto == nullptr) { - return true; - } - - if (inputProto->value_case() != onnx::TypeProto::kTensorType) { - return true; - } - - const auto& tensorType = inputProto->tensor_type(); - - if (!tensorType.has_shape()) { - return false; - } - - const auto& shape = tensorType.shape(); - shapeDims.resize(shape.dim_size()); - - for (uint32_t dimIndex = 0; dimIndex < static_cast(shape.dim_size()); ++dimIndex) { - if (!shape.dim(dimIndex).has_dim_value()) { - return false; + MLOperatorTensorGetter& mLOperatorTensorGetter + ) + : OpNodeInfoWrapper(info, nullptr, defaultAttributes, gsl::span(), mLOperatorTensorGetter) + { } - shapeDims[dimIndex] = gsl::narrow(shape.dim(dimIndex).dim_value()); - } + bool TryGetStaticShapeIfTensor( + const onnx::TypeProto* inputProto, + std::vector& shapeDims) + { + // Skip this input if it is not valid. + if (inputProto == nullptr) + { + return true; + } - return true; -} + if (inputProto->value_case() != onnx::TypeProto::kTensorType) + { + return true; + } -bool TryGetStaticInputShapes(const onnxruntime::Node& node, EdgeShapes& inputShapes) { - onnxruntime::ProtoHelperNodeContext protoContext(node); - onnxruntime::OpNodeProtoHelper info(&protoContext); + const auto& tensorType = inputProto->tensor_type(); - inputShapes.Reset(info.GetInputCount()); + if (!tensorType.has_shape()) + { + return false; + } - for (size_t inputIndex = 0; inputIndex < inputShapes.EdgeCount(); ++inputIndex) { - const onnx::TypeProto* inputProto = info.GetInputType(inputIndex); - if (!TryGetStaticShapeIfTensor(inputProto, inputShapes.GetMutableShape(inputIndex))) { - return false; + const auto& shape = tensorType.shape(); + shapeDims.resize(shape.dim_size()); + + for (uint32_t dimIndex = 0; dimIndex < static_cast(shape.dim_size()); ++dimIndex) + { + if (!shape.dim(dimIndex).has_dim_value()) + { + return false; + } + + shapeDims[dimIndex] = gsl::narrow(shape.dim(dimIndex).dim_value()); + } + + return true; } - } - return true; -} + bool TryGetStaticInputShapes(const onnxruntime::Node& node, EdgeShapes& inputShapes) + { + onnxruntime::ProtoHelperNodeContext protoContext(node); + onnxruntime::OpNodeProtoHelper info(&protoContext); -bool TryGetStaticOutputShapes(const onnxruntime::Node& node, EdgeShapes& outputShapes) { - onnxruntime::ProtoHelperNodeContext protoContext(node); - onnxruntime::OpNodeProtoHelper info(&protoContext); + inputShapes.Reset(info.GetInputCount()); - outputShapes.Reset(info.GetOutputCount()); + for (size_t inputIndex = 0; inputIndex < inputShapes.EdgeCount(); ++inputIndex) + { + const onnx::TypeProto* inputProto = info.GetInputType(inputIndex); + if (!TryGetStaticShapeIfTensor(inputProto, inputShapes.GetMutableShape(inputIndex))) + { + return false; + } + } - for (size_t outputIndex = 0; outputIndex < outputShapes.EdgeCount(); ++outputIndex) { - const onnx::TypeProto* outputProto = info.GetOutputType(outputIndex); - if (!TryGetStaticShapeIfTensor(outputProto, outputShapes.GetMutableShape(outputIndex))) { - return false; + return true; } - } - return true; -} + bool TryGetStaticOutputShapes(const onnxruntime::Node& node, EdgeShapes& outputShapes) + { + onnxruntime::ProtoHelperNodeContext protoContext(node); + onnxruntime::OpNodeProtoHelper info(&protoContext); -bool ContainsEmptyDimensions(const EdgeShapes& shapes, gsl::span ignoredShapeIndices) { - for (size_t i = 0; i < shapes.EdgeCount(); i++) { - const std::vector& shape = shapes.GetShape(i); + outputShapes.Reset(info.GetOutputCount()); - if (std::find(shape.begin(), shape.end(), 0u) != shape.end() && - std::find(ignoredShapeIndices.begin(), ignoredShapeIndices.end(), i) == ignoredShapeIndices.end()) { - return true; + for (size_t outputIndex = 0; outputIndex < outputShapes.EdgeCount(); ++outputIndex) + { + const onnx::TypeProto* outputProto = info.GetOutputType(outputIndex); + if (!TryGetStaticShapeIfTensor(outputProto, outputShapes.GetMutableShape(outputIndex))) + { + return false; + } + } + + return true; } - } - return false; -} + bool ContainsEmptyDimensions(const EdgeShapes& shapes, gsl::span ignoredShapeIndices) + { + for (size_t i = 0; i < shapes.EdgeCount(); i++) + { + const std::vector& shape = shapes.GetShape(i); -std::tuple, size_t> UnpackTensor(const onnx::TensorProto& initializer) { - std::unique_ptr unpackedTensor; - size_t tensorByteSize = 0; + if (std::find(shape.begin(), shape.end(), 0u) != shape.end() && + std::find(ignoredShapeIndices.begin(), ignoredShapeIndices.end(), i) == ignoredShapeIndices.end()) + { + return true; + } + } + + return false; + } + + std::tuple, size_t> UnpackTensor(const onnx::TensorProto& initializer) + { + std::unique_ptr unpackedTensor; + size_t tensorByteSize = 0; #define CASE_PROTO(X, Y, Z) \ case ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_##X: { \ size_t elementCount = initializer.##Z(); \ tensorByteSize = elementCount * sizeof(Y); \ unpackedTensor.reset(new std::byte[tensorByteSize]); \ - ORT_THROW_HR_IF(E_FAIL, !onnxruntime::utils::UnpackTensor( \ + ORT_THROW_HR_IF(E_FAIL, !onnxruntime::utils::UnpackTensor( \ initializer, \ initializer.has_raw_data() ? initializer.raw_data().data() : nullptr, \ initializer.has_raw_data() ? initializer.raw_data().size() : 0, \ @@ -2053,23 +2302,23 @@ std::tuple, size_t> UnpackTensor(const onnx::Tensor .IsOK()); \ break; \ } - switch (initializer.data_type()) { - CASE_PROTO(FLOAT, float, float_data_size); - CASE_PROTO(DOUBLE, double, double_data_size); - CASE_PROTO(BOOL, bool, int32_data_size); - CASE_PROTO(INT8, int8_t, int32_data_size); - CASE_PROTO(INT16, int16_t, int32_data_size); - CASE_PROTO(INT32, int32_t, int32_data_size); - CASE_PROTO(INT64, int64_t, int64_data_size); - CASE_PROTO(UINT8, uint8_t, int32_data_size); - CASE_PROTO(UINT16, uint16_t, int32_data_size); - CASE_PROTO(UINT32, uint32_t, uint64_data_size); - CASE_PROTO(UINT64, uint64_t, int64_data_size); - CASE_PROTO(FLOAT16, onnxruntime::MLFloat16, int32_data_size); - default: - ORT_THROW_HR(E_INVALIDARG); - } + switch (initializer.data_type()) + { + CASE_PROTO(FLOAT, float, float_data_size); + CASE_PROTO(DOUBLE, double, double_data_size); + CASE_PROTO(BOOL, bool, int32_data_size); + CASE_PROTO(INT8, int8_t, int32_data_size); + CASE_PROTO(INT16, int16_t, int32_data_size); + CASE_PROTO(INT32, int32_t, int32_data_size); + CASE_PROTO(INT64, int64_t, int64_data_size); + CASE_PROTO(UINT8, uint8_t, int32_data_size); + CASE_PROTO(UINT16, uint16_t, int32_data_size); + CASE_PROTO(UINT32, uint32_t, uint64_data_size); + CASE_PROTO(UINT64, uint64_t, int64_data_size); + CASE_PROTO(FLOAT16, onnxruntime::MLFloat16, int32_data_size); + default: ORT_THROW_HR(E_INVALIDARG); + } - return std::make_tuple(std::move(unpackedTensor), tensorByteSize); -} + return std::make_tuple(std::move(unpackedTensor), tensorByteSize); + } } // namespace winrt::Windows::AI::MachineLearning::implementation diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.h index 90309b314d..e1328ee01b 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/MLOperatorAuthorImpl.h @@ -22,10 +22,10 @@ namespace WRL } namespace Windows::AI::MachineLearning::Adapter -{ +{ using namespace Microsoft::WRL; - + // Inline method querying whether tensor shapes are defined, during wrappers // of shape inference callbacks. template @@ -122,7 +122,7 @@ public: }; // Base class for ABI objects which may be "Closed", at which point calls will predictably -// fail or return a dummy value. This is used for transient ABI context objects which +// fail or return a dummy value. This is used for transient ABI context objects which // are passed to methods on kernel or inferencers, and which wrap Lotus objects whose lifetimes // are not controlled by reference counts of the encapsulating object. class Closable @@ -158,15 +158,16 @@ class OpNodeInfoWrapper : public Base1_t, public Base2_t, public Closable OpNodeInfoWrapper() = delete; OpNodeInfoWrapper( - const onnxruntime::OpNodeProtoHelper* impl, + const onnxruntime::OpNodeProtoHelper* impl, const EdgeShapes* inputShapesOverride, const AttributeMap* defaultAttributes, gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter) : - m_impl(impl), - m_inputShapesOverride(inputShapesOverride), - m_constantInputGetter(constantInputGetter), - m_defaultAttributes(defaultAttributes) + MLOperatorTensorGetter& constantInputGetter + ) + : m_impl(impl), + m_inputShapesOverride(inputShapesOverride), + m_constantInputGetter(constantInputGetter), + m_defaultAttributes(defaultAttributes) { m_requiredConstantCpuInputs.assign(requiredConstantCpuInputs.begin(), requiredConstantCpuInputs.end()); } @@ -213,12 +214,12 @@ class OpNodeInfoWrapper : public Base1_t, public Base2_t, public Closable HRESULT STDMETHODCALLTYPE GetInputTensorDimensionCount(uint32_t inputIndex, uint32_t* dimensionCount) const noexcept; HRESULT STDMETHODCALLTYPE GetInputTensorShape(uint32_t inputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept; - + bool STDMETHODCALLTYPE IsInputValid(uint32_t inputIndex) const noexcept override; bool STDMETHODCALLTYPE IsOutputValid(uint32_t outputIndex) const noexcept override; HRESULT STDMETHODCALLTYPE GetConstantInputTensor( - uint32_t inputIndex, + uint32_t inputIndex, _Outptr_ IMLOperatorTensor** tensor ) const noexcept; @@ -239,7 +240,7 @@ class OpNodeInfoWrapper : public Base1_t, public Base2_t, public Closable // May be null const EdgeShapes* m_inputShapesOverride; - + std::vector m_requiredConstantCpuInputs; MLOperatorTensorGetter m_constantInputGetter; @@ -275,7 +276,7 @@ class TensorWrapper : public WRL::Base, public Closable private: // Lifetime is managed by the caller and guaranteed to outlive this class onnxruntime::Tensor* m_impl = nullptr; - + ComPtr m_winmlExecutionProvider; bool m_internalOperator = false; @@ -284,7 +285,7 @@ class TensorWrapper : public WRL::Base, public Closable bool m_isDataInterface = false; // The returned data may be a converted shadow copy, and the piece of it which - // is returned may vary according to kernel registration options. + // is returned may vary according to kernel registration options. ComPtr m_dataInterfaceOrShadowCopy; ComPtr m_abiDataInterface; @@ -326,7 +327,7 @@ class OnnxTensorWrapper : public WRL::Base, public Closable }; class OpKernelInfoWrapper : public OpNodeInfoWrapper< - onnxruntime::ProtoHelperNodeContext, + onnxruntime::ProtoHelperNodeContext, WRL::Base< Microsoft::WRL::ChainInterfaces, IMLOperatorTensorShapeDescription, IMLOperatorAttributes1>, @@ -334,17 +335,17 @@ class OpKernelInfoWrapper : public OpNodeInfoWrapper< { public: OpKernelInfoWrapper( - const onnxruntime::OpKernelInfo* kerneInfo, - IUnknown* abiExecutionObject, - const EdgeShapes* inputShapeOverrides, - const EdgeShapes* inferredOutputShapes, - bool allowInputShapeQuery, - bool allowOutputShapeQuery, - bool isInternalOperator, - const AttributeMap* defaultAttributes, - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter - ); + const onnxruntime::OpKernelInfo* kerneInfo, + IUnknown* abiExecutionObject, + const EdgeShapes* inputShapeOverrides, + const EdgeShapes* inferredOutputShapes, + bool allowInputShapeQuery, + bool allowOutputShapeQuery, + bool isInternalOperator, + const AttributeMap* defaultAttributes, + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& constantInputGetter + ); // HasTensorShapeDescription returns false if and only if the kernel is registered using // MLOperatorKernelOptions::AllowDynamicInputTensorSizes. If this flag is specified and upstream @@ -372,7 +373,7 @@ class OpKernelInfoWrapper : public OpNodeInfoWrapper< { return E_NOTIMPL; } - + private: // For shape info, in addition to the info const EdgeShapes* m_inferredOutputShapes = nullptr; @@ -383,15 +384,15 @@ private: ComPtr m_winmlProvider; const onnxruntime::OpKernelInfo* m_impl = nullptr; - + // The execution object returned through the ABI, which may vary according to kernel // registration options. - ComPtr m_abiExecutionObject; + ComPtr m_abiExecutionObject; }; // OpKernelInfo used for DML graph fusion. This uses the ONNX graph structures instead of ORT OpKernelInfo. class DmlGraphOpKernelInfoWrapper : public OpNodeInfoWrapper< - onnxruntime::ProtoHelperNodeContext, + onnxruntime::ProtoHelperNodeContext, WRL::Base< Microsoft::WRL::ChainInterfaces, IMLOperatorTensorShapeDescription, IMLOperatorAttributes1>, @@ -399,15 +400,15 @@ class DmlGraphOpKernelInfoWrapper : public OpNodeInfoWrapper< { public: DmlGraphOpKernelInfoWrapper( - const onnxruntime::OpNodeProtoHelper * protoHelper, - const void* executionHandle, - bool isInternalOperator, - const EdgeShapes* inferredOutputShapes, - const AttributeMap* defaultAttributes, - DmlGraphNodeCreateInfo* graphNodeCreateInfo, - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter - ); + const onnxruntime::OpNodeProtoHelper * protoHelper, + const void* executionHandle, + bool isInternalOperator, + const EdgeShapes* inferredOutputShapes, + const AttributeMap* defaultAttributes, + DmlGraphNodeCreateInfo* graphNodeCreateInfo, + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& constantInputGetter + ); // HasTensorShapeDescription returns false if and only if the kernel is registered using // MLOperatorKernelOptions::AllowDynamicInputTensorSizes. If this flag is specified and upstream @@ -421,9 +422,9 @@ class DmlGraphOpKernelInfoWrapper : public OpNodeInfoWrapper< HRESULT STDMETHODCALLTYPE GetOutputTensorDimensionCount(uint32_t inputIndex, uint32_t* dimensionCount) const noexcept override; bool STDMETHODCALLTYPE HasOutputShapeDescription() const noexcept override; HRESULT STDMETHODCALLTYPE GetOutputTensorShape(uint32_t inputIndex, uint32_t dimensionCount, uint32_t* dimensions) const noexcept override; - + bool STDMETHODCALLTYPE IsDmlGraphNode() const noexcept override; - + HRESULT STDMETHODCALLTYPE SetDmlOperator( IDMLOperator* op, _In_ const DML_OPERATOR_DESC* desc, @@ -459,7 +460,7 @@ class OpKernelContextWrapper : public WRL::Base, publi HRESULT STDMETHODCALLTYPE AllocateTemporaryData(size_t size, IUnknown** data, uint64_t* allocId) const; void STDMETHODCALLTYPE GetExecutionInterface(IUnknown** executionInterface) const noexcept override; - + void Close() override; std::vector GetInputTensors(); @@ -488,20 +489,21 @@ class OpKernelContextWrapper : public WRL::Base, publi // Compute being called on the kernel. This list is used to maintain their lifetime. mutable std::vector> m_temporaryAllocations; mutable std::vector> m_temporaryAbiAllocations; -}; +}; class AbiOpKernel : public onnxruntime::OpKernel { public: AbiOpKernel( - IMLOperatorKernelFactory* operatorFactory, - const onnxruntime::OpKernelInfo& kerneInfo, - bool requiresInputShapesAtCreation, - bool requiresOutputShapesAtCreation, - bool isInternalOperator, - gsl::span requiredConstantCpuInputs, - IMLOperatorShapeInferrer* shapeInferrer, - const AttributeMap* defaultAttributes); + IMLOperatorKernelFactory* operatorFactory, + const onnxruntime::OpKernelInfo& kerneInfo, + bool requiresInputShapesAtCreation, + bool requiresOutputShapesAtCreation, + bool isInternalOperator, + gsl::span requiredConstantCpuInputs, + IMLOperatorShapeInferrer* shapeInferrer, + const AttributeMap* defaultAttributes + ); onnxruntime::Status Compute(onnxruntime::OpKernelContext* context) const override; @@ -545,19 +547,19 @@ class AbiOpKernel : public onnxruntime::OpKernel ComPtr m_winmlProvider; bool m_internalOperator = false; std::vector m_requiredConstantCpuInputs; - + // The execution object returned through the ABI may vary according to kernel - // registration options. + // registration options. ComPtr m_providerExecutionObject; ComPtr m_abiExecutionObject; - + const AttributeMap* m_defaultAttributes = nullptr; }; class MLSchemaInferenceContext final : public OpNodeInfoWrapper< - onnx::InferenceContext, + onnx::InferenceContext, WRL::Base< - Microsoft::WRL::ChainInterfaces, + Microsoft::WRL::ChainInterfaces, IMLOperatorTypeInferenceContext, IMLOperatorAttributes, IMLOperatorAttributes1>, onnxruntime::null_type> { @@ -565,13 +567,13 @@ class MLSchemaInferenceContext final : public OpNodeInfoWrapper< MLSchemaInferenceContext() = delete; MLSchemaInferenceContext( - onnxruntime::OpNodeProtoHelper* info, + onnxruntime::OpNodeProtoHelper* info, onnx::InferenceContext* ctx, gsl::span requiredConstantCpuInputs, MLOperatorTensorGetter& mLOperatorTensorGetter ); - static ComPtr Create(onnxruntime::OpNodeProtoHelper* info, + static ComPtr Create(onnxruntime::OpNodeProtoHelper* info, onnx::InferenceContext* ctx, gsl::span requiredConstantCpuInputs); @@ -595,13 +597,14 @@ class MLKernelInferenceContext final : public OpNodeInfoWrapper< public: MLKernelInferenceContext() = delete; MLKernelInferenceContext( - onnxruntime::OpNodeProtoHelper* info, - const EdgeShapes* inputShapesOverride, - EdgeShapes& inferredOutputShapes, - const AttributeMap* defaultAttributes, - gsl::span requiredConstantCpuInputs, - MLOperatorTensorGetter& constantInputGetter) : - OpNodeInfoWrapper(info, inputShapesOverride, defaultAttributes, requiredConstantCpuInputs, constantInputGetter), + onnxruntime::OpNodeProtoHelper* info, + const EdgeShapes* inputShapesOverride, + EdgeShapes& inferredOutputShapes, + const AttributeMap* defaultAttributes, + gsl::span requiredConstantCpuInputs, + MLOperatorTensorGetter& constantInputGetter + ) + : OpNodeInfoWrapper(info, inputShapesOverride, defaultAttributes, requiredConstantCpuInputs, constantInputGetter), m_inferredOutputShapes(inferredOutputShapes) { } @@ -630,13 +633,15 @@ class MLSupportQueryContext final : public OpNodeInfoWrapper< MLSupportQueryContext() = delete; MLSupportQueryContext( - onnxruntime::OpNodeProtoHelper* info, - const AttributeMap* defaultAttributes, - MLOperatorTensorGetter& mLOperatorTensorGetter); + onnxruntime::OpNodeProtoHelper* info, + const AttributeMap* defaultAttributes, + MLOperatorTensorGetter& mLOperatorTensorGetter + ); static ComPtr Create( - onnxruntime::OpNodeProtoHelper* info, - const AttributeMap* defaultAttributes); + onnxruntime::OpNodeProtoHelper* info, + const AttributeMap* defaultAttributes + ); // TODO - ... }; diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperator.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperator.cpp index 2b44b68650..85284c6ada 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperator.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperator.cpp @@ -13,7 +13,7 @@ namespace Dml } void DmlOperator::SetDmlOperatorDesc( - const DML_OPERATOR_DESC& operatorDesc, + const DML_OPERATOR_DESC& operatorDesc, const MLOperatorKernelCreationContext& kernelInfo ) { @@ -99,7 +99,7 @@ namespace Dml } void DmlOperator::SetDmlOperatorDesc( - const DML_OPERATOR_DESC& operatorDesc, + const DML_OPERATOR_DESC& operatorDesc, const MLOperatorKernelContext& kernelInfo ) { diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp index d97dbc3d4a..e7198824bd 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp @@ -6,7 +6,7 @@ namespace Dml { -class DmlOperatorBatchNormalization : public DmlOperator +class DmlOperatorBatchNormalization : public DmlOperator, BatchNormalizationHelper { // This order matches the ONNX schema. enum OnnxInputIndex @@ -21,10 +21,13 @@ class DmlOperatorBatchNormalization : public DmlOperator public: DmlOperatorBatchNormalization(const MLOperatorKernelCreationContext& kernelCreationContext) - : DmlOperator(kernelCreationContext) + : DmlOperator(kernelCreationContext), + BatchNormalizationHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) { - std::vector> kernelInputIndices = {X, Mean, Variance, Scale, Bias}; - DmlOperator::Initialize(kernelCreationContext, kernelInputIndices); + // DML's BatchNormalization and ONNX order the input tensors differently (with DML as X, Mean, Variance, Scale, Bias), + // and normally we'd need to specify kernelInputIndices to Initialize, but we'll utilize DMLX's mapping instead. + // Passing both reordered kernelInputIndices to Initialize would otherwise confuse DMLX. + DmlOperator::Initialize(kernelCreationContext); ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs.size() == 5); ML_CHECK_VALID_ARGUMENT(m_outputTensorDescs.size() >= 1); @@ -47,19 +50,52 @@ public: std::vector inputDescs = GetDmlInputDescs(); std::vector outputDescs = GetDmlOutputDescs(); - DML_BATCH_NORMALIZATION_OPERATOR_DESC operatorDesc = {}; - operatorDesc.InputTensor = &inputDescs[X]; - operatorDesc.MeanTensor = &inputDescs[Mean]; - operatorDesc.VarianceTensor = &inputDescs[Variance]; - operatorDesc.ScaleTensor = &inputDescs[Scale]; - operatorDesc.BiasTensor = &inputDescs[Bias]; - operatorDesc.OutputTensor = &outputDescs[0]; - operatorDesc.Spatial = static_cast(spatial); - operatorDesc.Epsilon = epsilon; - operatorDesc.FusedActivation = fusedActivation ? &fusedActivationDmlDesc : nullptr; + dml::Graph graph(m_dmlDevice.Get()); + dml::TensorDesc inputTensorDesc = inputDescs[OnnxInputIndex::X]; + dml::TensorDesc scaleTensorDesc = inputDescs[OnnxInputIndex::Scale]; + dml::TensorDesc biasTensorDesc = inputDescs[OnnxInputIndex::Bias]; + dml::Expression input = dml::InputTensor(graph, OnnxInputIndex::X, inputTensorDesc); + dml::Expression scale = dml::InputTensor(graph, OnnxInputIndex::Scale, scaleTensorDesc); + dml::Expression bias = dml::InputTensor(graph, OnnxInputIndex::Bias, biasTensorDesc); + dml::Expression mean = dml::InputTensor(graph, OnnxInputIndex::Mean, inputDescs[OnnxInputIndex::Mean]); + dml::Expression variance = dml::InputTensor(graph, OnnxInputIndex::Variance, inputDescs[OnnxInputIndex::Variance]); - DML_OPERATOR_DESC opDesc = { DML_OPERATOR_BATCH_NORMALIZATION, &operatorDesc }; - SetDmlOperatorDesc(opDesc, kernelCreationContext); + // If scale and bias have different data types than input, then coerce them. + if (scaleTensorDesc.dataType != inputTensorDesc.dataType) + { + scale = dml::Cast(scale, inputTensorDesc.dataType); + } + if (biasTensorDesc.dataType != inputTensorDesc.dataType) + { + bias = dml::Cast(bias, inputTensorDesc.dataType); + } + + dml::Expression batchNormalization = dml::BatchNormalization( + input, + mean, + variance, + scale, + bias, + static_cast(spatial), + epsilon, + fusedActivation ? &fusedActivationDmlDesc : nullptr + ); + + DML_EXECUTION_FLAGS executionFlags = GetExecutionFlags(); + m_compiledOperator.Attach(graph.Compile(executionFlags, { batchNormalization }).Detach()); + } + + void Compute(const MLOperatorKernelContext& kernelContext) override + { + std::vector inputTensors = GetInputTensorsForExecute(kernelContext); + std::vector outputTensors = GetOutputTensorsForExecute(kernelContext); + + ORT_THROW_IF_FAILED(m_executionProvider->ExecuteOperator( + m_compiledOperator.Get(), + m_persistentResourceBinding ? &*m_persistentResourceBinding : nullptr, + gsl::make_span(inputTensors), + gsl::make_span(outputTensors) + )); } }; @@ -73,6 +109,7 @@ void CALLBACK QueryBatchNormalization(IMLOperatorSupportQueryContextPrivate* con } DML_OP_DEFINE_CREATION_FUNCTION(BatchNormalization, DmlOperatorBatchNormalization); +DML_OP_DEFINE_CREATION_FUNCTION(BatchNormalization15, DmlOperatorBatchNormalization); DML_OP_DEFINE_CREATION_FUNCTION(FusedBatchNormalization, DmlOperatorBatchNormalization); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp index 04aea8562e..0b22da3f07 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorSlice.cpp @@ -19,18 +19,13 @@ public: ML_CHECK_VALID_ARGUMENT(kernelInfo.GetOutputCount() == 1); std::vector> kernelInputIndices = { 0 }; // Only bind GPU to first 'data' tensor. - DmlOperator::Initialize(kernelInfo, kernelInputIndices); + DmlOperator::Initialize(kernelInfo, kernelInputIndices, std::nullopt, std::nullopt, std::nullopt, /*minimumDimensionCount*/ 1); const uint32_t inputTensorRank = m_inputTensorDescs[0].GetDimensionCount(); assert(inputTensorRank >= gsl::narrow_cast(m_offsets.size())); assert(inputTensorRank >= gsl::narrow_cast(m_sizes.size())); assert(inputTensorRank >= gsl::narrow_cast(m_strides.size())); - // Pad the parameters to respect DML's requirements - FillWithLeadingValues(/*inout*/ m_offsets, inputTensorRank, 0u); - FillWithLeadingValues(/*inout*/ m_sizes, inputTensorRank, 1u); - FillWithLeadingValues(/*inout*/ m_strides, inputTensorRank, 1); - std::vector inputDescs = GetDmlInputDescs(); std::vector outputDescs = GetDmlOutputDescs(); diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index d0198cadb6..b8249f7c04 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -95,6 +95,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(MaxRoiPool); DML_OP_EXTERN_CREATION_FUNCTION(RoiAlign10); DML_OP_EXTERN_CREATION_FUNCTION(InstanceNormalization); DML_OP_EXTERN_CREATION_FUNCTION(BatchNormalization); +DML_OP_EXTERN_CREATION_FUNCTION(BatchNormalization15); DML_OP_EXTERN_CREATION_FUNCTION(LRN); DML_OP_EXTERN_CREATION_FUNCTION(MeanVarianceNormalization); DML_OP_EXTERN_CREATION_FUNCTION(LpNormalization); @@ -396,11 +397,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, MaxRoiPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO_VER( 10, RoiAlign, typeNameListTwo, supportedTypeListRoiAlign, DmlGraphSupport::Supported)}, {REG_INFO( 7, InstanceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 7, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 9, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. - {REG_INFO( 14, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v14 adds training_mode attribute - // TODO: Add additional type constraints in BatchNormalization-15, with scale and bias (T1) being different from input X (T). - // {REG_INFO( 15, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v15 adds differing types for scale and bias vs input. + {REG_INFO( 7, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, + {REG_INFO( 9, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, // v9 just removes 'spatial' attribute. + {REG_INFO( 14, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v14 adds training_mode attribute + {REG_INFO_VER( 15, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v15 adds differing types for scale and bias vs input. {REG_INFO( 7, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 13, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, MeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp index 88070d5543..49fdac6eb3 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorUtility.cpp @@ -149,6 +149,7 @@ namespace Dml OperatorInfo{ "BatchNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_BatchNormalization }, OperatorInfo{ "BatchNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet9::sc_sinceVer_BatchNormalization }, OperatorInfo{ "BatchNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet14::sc_sinceVer_BatchNormalization }, + OperatorInfo{ "BatchNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet15::sc_sinceVer_BatchNormalization }, OperatorInfo{ "InstanceNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_InstanceNormalization }, OperatorInfo{ "MeanVarianceNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet7::sc_sinceVer_MeanVarianceNormalization }, OperatorInfo{ "MeanVarianceNormalization", onnxruntime::kOnnxDomain, OnnxOperatorSet9::sc_sinceVer_MeanVarianceNormalization }, diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h index 5de0c39def..8d9a8d10dd 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/precomp.h @@ -52,6 +52,7 @@ #include "External/DirectMLHelpers/GeneratedSchemaTypes.h" #include "External/DirectMLHelpers/SchemaHelpers.h" #include "External/DirectMLHelpers/GeneratedSchemaHelpers.h" +#include "External/DirectMLHelpers/DirectMLX.h" using Microsoft::WRL::ComPtr; diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp index 268e12706d..9169028ae1 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.cpp @@ -2253,4 +2253,27 @@ namespace OperatorHelper return { EdgeShapes(m_outputDimensions) }; } + void BatchNormalizationHelper::Initialize( + const IKernelInformationAdapter& kernelInformation, + const IShapeInformationAdapter& shapeInformation + ) + { + ML_CHECK_VALID_ARGUMENT(kernelInformation.GetInputCount() == 5); + ML_CHECK_VALID_ARGUMENT(kernelInformation.GetOutputCount() >= 1 && kernelInformation.GetOutputCount() <= 3); + } + + std::vector BatchNormalizationHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInformation) const + { + std::vector outputDimensionsList; + + outputDimensionsList.push_back(EdgeShapes(shapeInformation.GetInputTensorShape(0))); // output.shape = input.shape + int32_t trainingMode = shapeInformation.GetOptionalAttribute(AttrName::TrainingMode, 0); + if (trainingMode && shapeInformation.GetOutputCount() >= 3) + { + outputDimensionsList.push_back(EdgeShapes(shapeInformation.GetInputTensorShape(3))); // running_mean.shape = input_mean.shape + outputDimensionsList.push_back(EdgeShapes(shapeInformation.GetInputTensorShape(4))); // running_variance.shape = input_variance.shape + } + return outputDimensionsList; + } + } // namespace OperatorHelper diff --git a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h index f9c4f5eae9..0093160a6a 100644 --- a/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h +++ b/onnxruntime/core/providers/dml/OperatorAuthorHelper/OperatorHelper.h @@ -1302,6 +1302,23 @@ protected: std::vector m_outputDimensions; }; +class BatchNormalizationHelper +{ + void Initialize( + const IKernelInformationAdapter& kernelInformation, + const IShapeInformationAdapter& shapeInformation + ); + +public: + template + BatchNormalizationHelper(const Info_t& info, const Shape_t& shapeInfo) + { + Initialize(KernelInformationAdapter(info), ShapeInformationAdapter(shapeInfo)); + } + + std::vector GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const; +}; + using ShapeInferenceHelper_Conv = ConvHelper; using ShapeInferenceHelper_ConvTranspose = ConvTransposeHelper; using ShapeInferenceHelper_ConvTransposeWithDynamicPads = ConvTransposeWithDynamicPadsHelper; @@ -1318,6 +1335,7 @@ using ShapeInferenceHelper_MaxRoiPool = RoiPoolingHelper; using ShapeInferenceHelper_RoiAlign10 = VersionedOpsetHelper; using ShapeInferenceHelper_InstanceNormalization = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_BatchNormalization = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_BatchNormalization15 = BatchNormalizationHelper; using ShapeInferenceHelper_LRN = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_MeanVarianceNormalization = GetOutputShapeAsInputShapeHelper; @@ -1504,7 +1522,7 @@ using ShapeInferenceHelper_CastLike15 = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_FusedConv = ConvHelper; using ShapeInferenceHelper_FusedConvTranspose = ConvTransposeHelper; using ShapeInferenceHelper_FusedInstanceNormalization = GetOutputShapeAsInputShapeHelper; -using ShapeInferenceHelper_FusedBatchNormalization = GetOutputShapeAsInputShapeHelper; +using ShapeInferenceHelper_FusedBatchNormalization = BatchNormalizationHelper; using ShapeInferenceHelper_FusedMeanVarianceNormalization = GetOutputShapeAsInputShapeHelper; using ShapeInferenceHelper_FusedGemm = GemmHelper; using ShapeInferenceHelper_FusedMatMul = MatMulHelper; From 76024b8a6a4c45a5453df8e26fda91d190b13d9a Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Sat, 11 Jun 2022 18:51:32 -0700 Subject: [PATCH 11/19] Update DirectML.dll to 1.9.0 Preview --- .pipelines/nuget_config/x64/packages.config | 2 +- .pipelines/nuget_config/x86/packages.config | 2 +- cmake/external/dml.cmake | 2 +- packages.config | 2 +- tools/nuget/generate_nuspec_for_native_nuget.py | 4 ++-- 5 files changed, 6 insertions(+), 6 deletions(-) diff --git a/.pipelines/nuget_config/x64/packages.config b/.pipelines/nuget_config/x64/packages.config index 6fb9396686..9e318aacb0 100644 --- a/.pipelines/nuget_config/x64/packages.config +++ b/.pipelines/nuget_config/x64/packages.config @@ -1,6 +1,6 @@  - + diff --git a/.pipelines/nuget_config/x86/packages.config b/.pipelines/nuget_config/x86/packages.config index 84b98f86cd..61b20269c4 100644 --- a/.pipelines/nuget_config/x86/packages.config +++ b/.pipelines/nuget_config/x86/packages.config @@ -1,6 +1,6 @@  - + diff --git a/cmake/external/dml.cmake b/cmake/external/dml.cmake index 7beb571dcd..4b01e0d84f 100644 --- a/cmake/external/dml.cmake +++ b/cmake/external/dml.cmake @@ -40,7 +40,7 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML) set(NUGET_CONFIG ${PROJECT_SOURCE_DIR}/../NuGet.config) set(PACKAGES_CONFIG ${PROJECT_SOURCE_DIR}/../packages.config) get_filename_component(PACKAGES_DIR ${CMAKE_CURRENT_BINARY_DIR}/../packages ABSOLUTE) - set(DML_PACKAGE_DIR ${PACKAGES_DIR}/Microsoft.AI.DirectML.1.8.2) + set(DML_PACKAGE_DIR ${PACKAGES_DIR}/Microsoft.AI.DirectML.Preview.1.9.0-dev2b57b4f738b1d0dcc2dd31ecd502e36f4e3ea5a0) set(DML_SHARED_LIB DirectML.dll) # Restore nuget packages, which will pull down the DirectML redist package diff --git a/packages.config b/packages.config index 507e44f896..1df88367b9 100644 --- a/packages.config +++ b/packages.config @@ -1,6 +1,6 @@  - + diff --git a/tools/nuget/generate_nuspec_for_native_nuget.py b/tools/nuget/generate_nuspec_for_native_nuget.py index 9446edaa07..61a466748e 100644 --- a/tools/nuget/generate_nuspec_for_native_nuget.py +++ b/tools/nuget/generate_nuspec_for_native_nuget.py @@ -180,7 +180,7 @@ def generate_repo_url(list, repo_url, commit_id): def generate_dependencies(list, package_name, version): - dml_dependency = '' + dml_dependency = '' if package_name == "Microsoft.AI.MachineLearning": list.append("") @@ -200,7 +200,7 @@ def generate_dependencies(list, package_name, version): list.append("") else: - include_dml = package_name == "Microsoft.ML.OnnxRuntime.DirectML" + include_dml = package_name == "Microsoft.ML.OnnxRuntime.DirectML.Preview" list.append("") # Support .Net Core From 2bc487a816b55445d061c2e573c5f52fab5917dc Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Sat, 11 Jun 2022 19:15:19 -0700 Subject: [PATCH 12/19] Appease flaky flake tool --- tools/nuget/generate_nuspec_for_native_nuget.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tools/nuget/generate_nuspec_for_native_nuget.py b/tools/nuget/generate_nuspec_for_native_nuget.py index 1bf051332d..4408bb53a1 100644 --- a/tools/nuget/generate_nuspec_for_native_nuget.py +++ b/tools/nuget/generate_nuspec_for_native_nuget.py @@ -188,7 +188,8 @@ def generate_dependencies(xml_text, package_name, version, dependency_id, depend xml_text.append("") return - dml_dependency = '' + dml_dependency = '' if package_name == "Microsoft.AI.MachineLearning": xml_text.append("") From 04dd6639de1d2dcb8ff410994ebc68000ab448a9 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Sat, 11 Jun 2022 19:17:20 -0700 Subject: [PATCH 13/19] And appease the time wasting formatting tool now -_-... --- tools/nuget/generate_nuspec_for_native_nuget.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tools/nuget/generate_nuspec_for_native_nuget.py b/tools/nuget/generate_nuspec_for_native_nuget.py index 4408bb53a1..f1beaf6de9 100644 --- a/tools/nuget/generate_nuspec_for_native_nuget.py +++ b/tools/nuget/generate_nuspec_for_native_nuget.py @@ -188,8 +188,9 @@ def generate_dependencies(xml_text, package_name, version, dependency_id, depend xml_text.append("") return - dml_dependency = '' + dml_dependency = ( + '' + ) if package_name == "Microsoft.AI.MachineLearning": xml_text.append("") From 4c1a410d54bc63a0d18df940af45260018f65b94 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Tue, 14 Jun 2022 23:12:58 -0700 Subject: [PATCH 14/19] Unmangle DML preview package filenames --- cmake/external/dml.cmake | 61 ++++++++++++++++++++++++++++++++++++---- 1 file changed, 56 insertions(+), 5 deletions(-) diff --git a/cmake/external/dml.cmake b/cmake/external/dml.cmake index 4b01e0d84f..6362e9ea7d 100644 --- a/cmake/external/dml.cmake +++ b/cmake/external/dml.cmake @@ -43,15 +43,66 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML) set(DML_PACKAGE_DIR ${PACKAGES_DIR}/Microsoft.AI.DirectML.Preview.1.9.0-dev2b57b4f738b1d0dcc2dd31ecd502e36f4e3ea5a0) set(DML_SHARED_LIB DirectML.dll) - # Restore nuget packages, which will pull down the DirectML redist package + # If using the preview package, extract the SHA-1 from the path so we can unmangle the filenames later. + # e.g. "Microsoft.AI.DirectML.Preview.1.9.0-dev2b57b4f738b1d0dcc2dd31ecd502e36f4e3ea5a0" + if(DML_PACKAGE_DIR MATCHES ".*Preview.*-dev(.*)") + set(DML_PREVIEW_FILENAME_SUFFIX ".${CMAKE_MATCH_1}") + endif() + + # Restore nuget packages, which will pull down the DirectML redist package. add_custom_command( - OUTPUT ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib - DEPENDS ${PACKAGES_CONFIG} ${NUGET_CONFIG} + OUTPUT + ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + DEPENDS + ${PACKAGES_CONFIG} + ${NUGET_CONFIG} COMMAND ${CMAKE_CURRENT_BINARY_DIR}/nuget/src/nuget restore ${PACKAGES_CONFIG} -PackagesDirectory ${PACKAGES_DIR} -ConfigFile ${NUGET_CONFIG} - VERBATIM) + VERBATIM + ) + + # If using a preview package, unmangle the filenames from the nuget so they're useable. + # e.g. Map DirectML.2b57b4f738b1d0dcc2dd31ecd502e36f4e3ea5a0.dll -> DirectML.dll + if(DEFINED DML_PREVIEW_FILENAME_SUFFIX) + add_custom_command( + OUTPUT + ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib + DEPENDS + ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.pdb + VERBATIM + ) + endif() include_directories(BEFORE "${DML_PACKAGE_DIR}/include") - add_custom_target(RESTORE_PACKAGES ALL DEPENDS ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib) + add_custom_target( + RESTORE_PACKAGES ALL + DEPENDS + ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib + ) + add_dependencies(RESTORE_PACKAGES nuget) else() if (dml_EXTERNAL_PROJECT) From e3ec30efb65d596a6aab5dea4b8edbdb8c1a1370 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Tue, 14 Jun 2022 23:28:15 -0700 Subject: [PATCH 15/19] Add missing GELU to ApiHelpers.h --- .../src/External/DirectMLHelpers/ApiHelpers.h | 2 ++ 1 file changed, 2 insertions(+) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h index 8c85e4ec1d..76ca37bd05 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/External/DirectMLHelpers/ApiHelpers.h @@ -28,6 +28,7 @@ union ActivationOperatorDescUnion DML_ACTIVATION_TANH_OPERATOR_DESC tanh; DML_ACTIVATION_THRESHOLDED_RELU_OPERATOR_DESC thresholdedRelu; DML_ACTIVATION_SHRINK_OPERATOR_DESC shrink; + DML_ACTIVATION_GELU_OPERATOR_DESC gelu; }; struct ActivationOperatorDesc @@ -64,6 +65,7 @@ struct ActivationOperatorDesc case DML_OPERATOR_ACTIVATION_TANH: return { activationType, ¶ms.tanh }; case DML_OPERATOR_ACTIVATION_THRESHOLDED_RELU: return { activationType, ¶ms.thresholdedRelu }; case DML_OPERATOR_ACTIVATION_SHRINK: return { activationType, ¶ms.shrink }; + case DML_OPERATOR_ACTIVATION_GELU: return { activationType, ¶ms.gelu }; default: ORT_THROW_HR(E_INVALIDARG); return { activationType, ¶ms.relu }; From 508c76a246a2cbebcd2a54b2d4ff371413d8ab02 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 15 Jun 2022 00:16:10 -0700 Subject: [PATCH 16/19] Add missing DirectML.Debug.dll --- cmake/external/dml.cmake | 37 +++++++++++++++++++++++++------------ 1 file changed, 25 insertions(+), 12 deletions(-) diff --git a/cmake/external/dml.cmake b/cmake/external/dml.cmake index 6362e9ea7d..3c2bb5cbfd 100644 --- a/cmake/external/dml.cmake +++ b/cmake/external/dml.cmake @@ -77,18 +77,31 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML) ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.pdb + + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug.pdb + + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug.pdb + + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug.pdb + + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.dll + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug.pdb + VERBATIM ) endif() From ff8b173286248480b7b316b4c97b0031a4190b11 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 15 Jun 2022 00:18:40 -0700 Subject: [PATCH 17/19] Typo in DirectML.Debug.dll --- cmake/external/dml.cmake | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/cmake/external/dml.cmake b/cmake/external/dml.cmake index 3c2bb5cbfd..4f8a79a59f 100644 --- a/cmake/external/dml.cmake +++ b/cmake/external/dml.cmake @@ -81,25 +81,25 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML) COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug.pdb COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug.pdb COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug.pdb COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug.pdb + COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug.dll COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug.pdb VERBATIM From babd6e3fcdcbf5246e675e8d4dcb3e07e581e176 Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 15 Jun 2022 18:16:58 -0700 Subject: [PATCH 18/19] Update DirectML preview package with unmangled names --- .pipelines/nuget_config/x64/packages.config | 2 +- .pipelines/nuget_config/x86/packages.config | 2 +- cmake/external/dml.cmake | 59 ++----------------- packages.config | 2 +- .../nuget/generate_nuspec_for_native_nuget.py | 2 +- 5 files changed, 9 insertions(+), 58 deletions(-) diff --git a/.pipelines/nuget_config/x64/packages.config b/.pipelines/nuget_config/x64/packages.config index 9e318aacb0..0cb753af5d 100644 --- a/.pipelines/nuget_config/x64/packages.config +++ b/.pipelines/nuget_config/x64/packages.config @@ -1,6 +1,6 @@  - + diff --git a/.pipelines/nuget_config/x86/packages.config b/.pipelines/nuget_config/x86/packages.config index 61b20269c4..db298cbb38 100644 --- a/.pipelines/nuget_config/x86/packages.config +++ b/.pipelines/nuget_config/x86/packages.config @@ -1,6 +1,6 @@  - + diff --git a/cmake/external/dml.cmake b/cmake/external/dml.cmake index 4f8a79a59f..d15a110695 100644 --- a/cmake/external/dml.cmake +++ b/cmake/external/dml.cmake @@ -40,22 +40,16 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML) set(NUGET_CONFIG ${PROJECT_SOURCE_DIR}/../NuGet.config) set(PACKAGES_CONFIG ${PROJECT_SOURCE_DIR}/../packages.config) get_filename_component(PACKAGES_DIR ${CMAKE_CURRENT_BINARY_DIR}/../packages ABSOLUTE) - set(DML_PACKAGE_DIR ${PACKAGES_DIR}/Microsoft.AI.DirectML.Preview.1.9.0-dev2b57b4f738b1d0dcc2dd31ecd502e36f4e3ea5a0) + set(DML_PACKAGE_DIR ${PACKAGES_DIR}/Microsoft.AI.DirectML.Preview.1.9.0-devd10042c94985065a565c042540e15eb75b554663) set(DML_SHARED_LIB DirectML.dll) - # If using the preview package, extract the SHA-1 from the path so we can unmangle the filenames later. - # e.g. "Microsoft.AI.DirectML.Preview.1.9.0-dev2b57b4f738b1d0dcc2dd31ecd502e36f4e3ea5a0" - if(DML_PACKAGE_DIR MATCHES ".*Preview.*-dev(.*)") - set(DML_PREVIEW_FILENAME_SUFFIX ".${CMAKE_MATCH_1}") - endif() - # Restore nuget packages, which will pull down the DirectML redist package. add_custom_command( OUTPUT - ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib + ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib + ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib DEPENDS ${PACKAGES_CONFIG} ${NUGET_CONFIG} @@ -63,49 +57,6 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML) VERBATIM ) - # If using a preview package, unmangle the filenames from the nuget so they're useable. - # e.g. Map DirectML.2b57b4f738b1d0dcc2dd31ecd502e36f4e3ea5a0.dll -> DirectML.dll - if(DEFINED DML_PREVIEW_FILENAME_SUFFIX) - add_custom_command( - OUTPUT - ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib - ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib - ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib - ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib - DEPENDS - ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib - - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x64-win/DirectML.Debug.pdb - - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/x86-win/DirectML.Debug.pdb - - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm-win/DirectML.Debug.pdb - - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.lib ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.lib - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.pdb - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.dll ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug.dll - COMMAND ${CMAKE_COMMAND} -E copy_if_different ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug${DML_PREVIEW_FILENAME_SUFFIX}.pdb ${DML_PACKAGE_DIR}/bin/arm64-win/DirectML.Debug.pdb - - VERBATIM - ) - endif() - include_directories(BEFORE "${DML_PACKAGE_DIR}/include") add_custom_target( RESTORE_PACKAGES ALL diff --git a/packages.config b/packages.config index 1df88367b9..5442a72e1a 100644 --- a/packages.config +++ b/packages.config @@ -1,6 +1,6 @@  - + diff --git a/tools/nuget/generate_nuspec_for_native_nuget.py b/tools/nuget/generate_nuspec_for_native_nuget.py index f1beaf6de9..4a6b55d731 100644 --- a/tools/nuget/generate_nuspec_for_native_nuget.py +++ b/tools/nuget/generate_nuspec_for_native_nuget.py @@ -189,7 +189,7 @@ def generate_dependencies(xml_text, package_name, version, dependency_id, depend return dml_dependency = ( - '' + '' ) if package_name == "Microsoft.AI.MachineLearning": From fe7b8b80ae35c62e5bf91eed068be4279b76173d Mon Sep 17 00:00:00 2001 From: Dwayne Robinson Date: Wed, 15 Jun 2022 21:49:18 -0700 Subject: [PATCH 19/19] Revert BatchNormalization change for now, falling back to CPU on mixed types until a more advanced solution is written --- .../DmlOperatorBatchNormalization.cpp | 98 ++++++++++++++++++- .../src/Operators/OperatorRegistration.cpp | 8 +- 2 files changed, 99 insertions(+), 7 deletions(-) diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp index e7198824bd..c69e76731a 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/DmlOperatorBatchNormalization.cpp @@ -21,6 +21,67 @@ class DmlOperatorBatchNormalization : public DmlOperator, BatchNormalizationHelp public: DmlOperatorBatchNormalization(const MLOperatorKernelCreationContext& kernelCreationContext) + : DmlOperator(kernelCreationContext), + BatchNormalizationHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) + { + std::vector> kernelInputIndices = {X, Mean, Variance, Scale, Bias}; + DmlOperator::Initialize(kernelCreationContext, kernelInputIndices); + + ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs.size() == 5); + ML_CHECK_VALID_ARGUMENT(m_outputTensorDescs.size() >= 1); + + const float epsilon = kernelCreationContext.GetOptionalAttribute(AttrName::Epsilon, 0.0f); + const int spatial = kernelCreationContext.GetOptionalAttribute(AttrName::Spatial, 1); + const std::optional fusedActivation = FusionHelpers::TryGetFusedActivationDesc(kernelCreationContext); + DML_OPERATOR_DESC fusedActivationDmlDesc = fusedActivation ? fusedActivation->GetDmlDesc() : DML_OPERATOR_DESC(); + + m_inputTensorDescs[0] = CreateTensorDescFromInput(kernelCreationContext, 0, TensorAxis::DoNotCoerce, TensorAxis::N, TensorAxis::LeftAligned); + + // Massage each of these 1D tensors (of length C) into ND tensors of the form [1,C,1,1,...]. + for (uint32_t i = Scale; i < OnnxInputIndex::Count; ++i) + { + m_inputTensorDescs[i] = CreateTensorDescFromInput(kernelCreationContext, i, TensorAxis::DoNotCoerce, TensorAxis::C, TensorAxis::LeftAligned, std::nullopt, m_inputTensorDescs[0].GetDimensionCount()); + } + + m_outputTensorDescs[0] = CreateTensorDescFromOutput(kernelCreationContext, 0, TensorAxis::DoNotCoerce, TensorAxis::N, TensorAxis::LeftAligned, std::nullopt, m_inputTensorDescs[0].GetDimensionCount()); + + ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs.size() == 5); + ML_CHECK_VALID_ARGUMENT(m_outputTensorDescs.size() >= 1); + + std::vector inputDescs = GetDmlInputDescs(); + std::vector outputDescs = GetDmlOutputDescs(); + + DML_BATCH_NORMALIZATION_OPERATOR_DESC operatorDesc = {}; + operatorDesc.InputTensor = &inputDescs[X]; + operatorDesc.MeanTensor = &inputDescs[Mean]; + operatorDesc.VarianceTensor = &inputDescs[Variance]; + operatorDesc.ScaleTensor = &inputDescs[Scale]; + operatorDesc.BiasTensor = &inputDescs[Bias]; + operatorDesc.OutputTensor = &outputDescs[0]; + operatorDesc.Spatial = static_cast(spatial); + operatorDesc.Epsilon = epsilon; + operatorDesc.FusedActivation = fusedActivation ? &fusedActivationDmlDesc : nullptr; + + DML_OPERATOR_DESC opDesc = { DML_OPERATOR_BATCH_NORMALIZATION, &operatorDesc }; + SetDmlOperatorDesc(opDesc, kernelCreationContext); + } +}; + +class DmlOperatorBatchNormalization15 : public DmlOperator, BatchNormalizationHelper +{ + // This order matches the ONNX schema. + enum OnnxInputIndex + { + X, // Input + Scale, + Bias, + Mean, + Variance, + Count, + }; + +public: + DmlOperatorBatchNormalization15(const MLOperatorKernelCreationContext& kernelCreationContext) : DmlOperator(kernelCreationContext), BatchNormalizationHelper(kernelCreationContext, kernelCreationContext.GetTensorShapeDescription()) { @@ -101,15 +162,46 @@ public: void CALLBACK QueryBatchNormalization(IMLOperatorSupportQueryContextPrivate* context, /*out*/ bool* isSupported) { - // training_mode=1 is unsupported as it isn't needed for inference (https://github.com/onnx/onnx/pull/3333). + *isSupported = false; + // training_mode=1 is unsupported as it isn't needed for inference (https://github.com/onnx/onnx/pull/3333). MLOperatorAttributes attributes(context); int32_t trainingMode = attributes.GetOptionalAttribute(AttrName::TrainingMode, 0); - *isSupported = (trainingMode == 0); + if (trainingMode != 0) + { + return; + } + + if (context->GetInputCount() < 5) + { + return; + } + + // Get the data type of each tensor. + MLOperatorEdgeDescription operatorEdgeDescription[5]; + for (uint32_t i = 0; i < 5; ++i) + { + if (FAILED(context->GetInputEdgeDescription(i, &operatorEdgeDescription[i])) + || operatorEdgeDescription[i].edgeType != MLOperatorEdgeType::Tensor) + { + return; + } + } + + // Fall back if the data types of the mean/variance or scale/bias differ from the input. + MLOperatorTensorDataType inputTensorDataType = operatorEdgeDescription[0].tensorDataType; + for (uint32_t i = 1; i < 5; ++i) + { + if (operatorEdgeDescription[i].tensorDataType != inputTensorDataType) + { + return; + } + } + + *isSupported = true; } DML_OP_DEFINE_CREATION_FUNCTION(BatchNormalization, DmlOperatorBatchNormalization); -DML_OP_DEFINE_CREATION_FUNCTION(BatchNormalization15, DmlOperatorBatchNormalization); DML_OP_DEFINE_CREATION_FUNCTION(FusedBatchNormalization, DmlOperatorBatchNormalization); } // namespace Dml diff --git a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp index b8249f7c04..3ad7c4606b 100644 --- a/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp +++ b/onnxruntime/core/providers/dml/DmlExecutionProvider/src/Operators/OperatorRegistration.cpp @@ -397,10 +397,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation {REG_INFO( 7, MaxRoiPool, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO_VER( 10, RoiAlign, typeNameListTwo, supportedTypeListRoiAlign, DmlGraphSupport::Supported)}, {REG_INFO( 7, InstanceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, - {REG_INFO( 7, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, - {REG_INFO( 9, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported)}, // v9 just removes 'spatial' attribute. - {REG_INFO( 14, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v14 adds training_mode attribute - {REG_INFO_VER( 15, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::NotSupported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v15 adds differing types for scale and bias vs input. + {REG_INFO( 7, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, + {REG_INFO( 9, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, // v9 just removes 'spatial' attribute. + {REG_INFO( 14, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v14 adds training_mode attribute + {REG_INFO( 15, BatchNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryBatchNormalization)}, // v15 adds differing types for scale and bias vs input. {REG_INFO( 7, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 13, LRN, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)}, {REG_INFO( 7, MeanVarianceNormalization, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},